Problem 791: Average and Variance

View on Project Euler

Project Euler Problem 791 Solution

EulerSolve provides an optimized solution for Project Euler Problem 791, Average and Variance, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary Problem 791 studies nondecreasing positive integer quadruples \((a,b,c,d)\) whose average equals their variance. Writing $$\mu=\frac{a+b+c+d}{4},\qquad V=\frac{(a-\mu)^2+(b-\mu)^2+(c-\mu)^2+(d-\mu)^2}{4},$$ we seek all quadruples with \(\mu=V\) and \(\mu\le n\), and \(S(n)\) is the sum of the corresponding totals \(a+b+c+d=4\mu\). The direct four-variable search is far too large for \(n=10^8\), so the implementations reduce the problem to a three-dimensional lattice search with closed-form summation over the innermost coordinate. Mathematical Approach The important fact encoded by the implementations is that the original average-equals-variance constraint can be rewritten in terms of three odd integers. Once that reduction is done, the remaining work is purely arithmetic. Step 1: Pass to the Reduced Odd Coordinates After ordering the quadruple, eliminating the redundant degree of freedom, and applying the linear change of variables used by the implementations, every admissible configuration is represented by odd integers \((u,v,w)\) satisfying $$u\equiv v\equiv w\equiv 1\pmod 2.$$ Its contribution to \(S(n)\) becomes $$T(u,v,w)=\frac{u^2+v^2+w^2+2u+2v+2w+3}{2}.$$ This is the quantity accumulated by the implementations. The original quadruple is no longer enumerated explicitly; the whole computation is carried out in the reduced \((u,v,w)\)-space....

Detailed mathematical approach

Problem Summary

Problem 791 studies nondecreasing positive integer quadruples \((a,b,c,d)\) whose average equals their variance. Writing

$$\mu=\frac{a+b+c+d}{4},\qquad V=\frac{(a-\mu)^2+(b-\mu)^2+(c-\mu)^2+(d-\mu)^2}{4},$$

we seek all quadruples with \(\mu=V\) and \(\mu\le n\), and \(S(n)\) is the sum of the corresponding totals \(a+b+c+d=4\mu\). The direct four-variable search is far too large for \(n=10^8\), so the implementations reduce the problem to a three-dimensional lattice search with closed-form summation over the innermost coordinate.

Mathematical Approach

The important fact encoded by the implementations is that the original average-equals-variance constraint can be rewritten in terms of three odd integers. Once that reduction is done, the remaining work is purely arithmetic.

Step 1: Pass to the Reduced Odd Coordinates

After ordering the quadruple, eliminating the redundant degree of freedom, and applying the linear change of variables used by the implementations, every admissible configuration is represented by odd integers \((u,v,w)\) satisfying

$$u\equiv v\equiv w\equiv 1\pmod 2.$$

Its contribution to \(S(n)\) becomes

$$T(u,v,w)=\frac{u^2+v^2+w^2+2u+2v+2w+3}{2}.$$

This is the quantity accumulated by the implementations. The original quadruple is no longer enumerated explicitly; the whole computation is carried out in the reduced \((u,v,w)\)-space.

Step 2: Determine the Outer Feasible Region

Set

$$R=8n.$$

The outer coordinate is restricted to odd values in

$$1\le u\le \left\lfloor\sqrt{R+3}\right\rfloor-2,\qquad u\equiv 1\pmod 2.$$

For each fixed \(u\), the second coordinate is restricted to

$$-1\le v\le \min\!\left(u,\left\lfloor\sqrt{R-u^2-4u-1}\right\rfloor-2\right),\qquad v\equiv 1\pmod 2.$$

These bounds already remove most of the search space. Instead of a four-dimensional brute-force enumeration, the solver now scans only the feasible odd lattice points in a curved two-dimensional outer region.

Step 3: Compute the Inner Interval and the Central Hole

For each admissible pair \((u,v)\), define

$$w_{\max}=\left\lfloor\sqrt{R-u^2-v^2-4u-4v-5}\right\rfloor.$$

Then \(w\) must lie in the odd interval

$$\max(-w_{\max},-v-2)\le w\le \min(w_{\max},v),\qquad w\equiv 1\pmod 2.$$

There is one more condition. If

$$11-u^2-v^2>0,$$

then the middle of the interval is forbidden and we must keep only

$$|w|\ge \left\lceil\sqrt{11-u^2-v^2}\right\rceil.$$

So for a fixed \((u,v)\), the innermost search is never an arbitrary set: it is either one odd interval or two odd intervals separated by a central gap. That structural fact is exactly what makes the constant-time interval summation possible.

Step 4: Sum an Entire Odd Interval in Closed Form

Suppose one admissible \(w\)-interval is

$$w=L,L+2,\dots,H,$$

with \(L\) and \(H\) odd. Let

$$m=\frac{H-L}{2}+1,\qquad w_j=L+2j\quad (0\le j\le m-1).$$

Then

$$\sum_{j=0}^{m-1} j=\frac{m(m-1)}{2},\qquad \sum_{j=0}^{m-1} j^2=\frac{(m-1)m(2m-1)}{6},$$

and therefore

$$\sum_{j=0}^{m-1} w_j=mL+2\sum_{j=0}^{m-1}j,$$

$$\sum_{j=0}^{m-1} w_j^2=mL^2+4L\sum_{j=0}^{m-1}j+4\sum_{j=0}^{m-1}j^2.$$

If we abbreviate the part independent of \(w\) by

$$K(u,v)=u^2+v^2+2u+2v+3,$$

then the whole interval contributes

$$\sum_{j=0}^{m-1} T(u,v,w_j)=\frac{mK(u,v)+\sum w_j^2+2\sum w_j}{2}.$$

This replaces a potentially long inner loop by a handful of arithmetic operations.

Step 5: Split Only When the Hole Is Active

If the central exclusion is absent, one closed-form evaluation is enough. If the exclusion is present, the allowed odd values split into a negative tail and a positive tail:

$$[L,H]\cap\{w:w\equiv 1\pmod 2\}=[L,-c]\cup[c,H]$$

for the smallest admissible odd threshold \(c\). Each tail is summed by the same interval formula, and the two results are added modulo \(433494437\).

So the algorithm never iterates over admissible \(w\)-values one by one. That is the decisive optimization shared by the C++, Python, and Java implementations.

Worked Example: \(S(5)=48\)

When \(n=5\), we have \(R=40\). The reduced search produces exactly the following admissible odd triples:

$$\begin{aligned} (1,1,-3)&\mapsto T=6,\\ (3,-1,-1)&\mapsto T=8,\\ (3,1,-3)&\mapsto T=12,\\ (3,1,-1)&\mapsto T=10,\\ (3,1,1)&\mapsto T=12. \end{aligned}$$

Therefore

$$S(5)=6+8+12+10+12=48,$$

which matches the checkpoint built into the overall solution strategy.

How the Code Works

The C++, Python, and Java implementations all follow the same mathematical pipeline. First they compute \(R=8n\), determine the largest admissible odd outer coordinate, and loop over the feasible odd pairs \((u,v)\) using integer square roots to obtain sharp bounds instead of trial-and-error scanning.

For each outer pair, the implementation constructs the admissible odd range for the innermost coordinate. If the central hole is inactive, it evaluates one interval. If the hole is active, it evaluates a negative interval and a positive interval. In both cases it uses the closed forms for \(\sum w\) and \(\sum w^2\), so the innermost work is constant time.

The arithmetic is always reduced modulo \(433494437\). The C++ implementation parallelizes the outer loop because different odd \(u\)-slices are independent; the Python and Java implementations use the same bounds and formulas in sequential form. The C++ version also checks the intermediate values \(S(5)=48\) and \(S(10^3)=37048340\) before computing the final target.

Complexity Analysis

The outer coordinate ranges over \(O(\sqrt{n})\) odd values, and for each such value the second coordinate also spans at most \(O(\sqrt{n})\) odd values. Hence the number of feasible outer pairs is \(O(n)\). Since each pair is handled by at most two constant-time interval evaluations, the total running time is \(O(n)\). The sequential versions use \(O(1)\) auxiliary space, while the parallel C++ version uses \(O(P)\) partial storage for \(P\) worker threads.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=791
  2. Arithmetic mean: Wikipedia — Arithmetic mean
  3. Variance: Wikipedia — Variance
  4. Arithmetic progression: Wikipedia — Arithmetic progression
  5. Sum of squares: Wikipedia — Square pyramidal number

Problem 791 source code

C++

#include <atomic>
#include <cmath>
#include <cstdint>
#include <iostream>
#include <thread>
#include <vector>

using namespace std;

static constexpr long long MOD = 433494437LL;

static inline long long isqrt_ll(long long x) {
    if (x <= 0) return 0;
    long long r = (long long)std::sqrt((long double)x);
    while ((r + 1) * (r + 1) <= x) ++r;
    while (r * r > x) --r;
    return r;
}

static inline long long ceil_sqrt_ll(long long x) {
    if (x <= 0) return 0;
    long long r = isqrt_ll(x);
    if (r * r < x) ++r;
    return r;
}

static inline long long sum_interval(long long L, long long H, long long N0) {
    if ((L & 1LL) == 0) ++L;
    if ((H & 1LL) == 0) --H;
    if (L > H) return 0;

    long long cnt = (H - L) / 2 + 1;
    __int128 t = cnt - 1;
    __int128 sum_i = (__int128)cnt * t / 2;
    __int128 sum_i2 = t * (__int128)cnt * (2 * t + 1) / 6;

    __int128 sum_C = (__int128)cnt * L + 2 * sum_i;
    __int128 sum_C2 = (__int128)cnt * L * L + 4 * (__int128)L * sum_i + 4 * sum_i2;

    __int128 sum_num = (__int128)cnt * N0 + sum_C2 + 2 * sum_C;
    __int128 sum_s = sum_num / 2;
    long long mod = (long long)(sum_s % MOD);
    if (mod < 0) mod += MOD;
    return mod;
}

static long long S_value(long long n) {
    if (n <= 0) return 0;
    long long R = 8 * n;

    long long A_max = isqrt_ll(R + 3) - 2;
    if (A_max < 1) return 0;
    if ((A_max & 1LL) == 0) --A_max;

    long long countA = (A_max - 1) / 2 + 1;
    unsigned int T = thread::hardware_concurrency();
    if (T == 0) T = 4;
    if ((long long)T > countA) T = (unsigned int)countA;
    if (T == 0) T = 1;

    atomic<long long> nextA{1};
    vector<long long> partial(T, 0);
    vector<thread> threads;
    threads.reserve(T);

    for (unsigned int t = 0; t < T; ++t) {
        threads.emplace_back([&, t]() {
            long long local = 0;
            for (;;) {
                long long A = nextA.fetch_add(2, memory_order_relaxed);
                if (A > A_max) break;
                long long A2 = A * A;
                long long D = R - A2 - 4 * A - 1;
                if (D < 0) continue;
                long long B_max = isqrt_ll(D) - 2;
                if (B_max < -1) continue;
                if (B_max > A) B_max = A;

                for (long long B = -1; B <= B_max; B += 2) {
                    long long B2 = B * B;
                    long long L = R - A2 - B2 - 4 * A - 4 * B - 5;
                    if (L < 0) continue;
                    long long C_max = isqrt_ll(L);

                    long long C_low = -C_max;
                    long long low2 = -B - 2;
                    if (low2 > C_low) C_low = low2;
                    long long C_high = C_max;
                    if (B < C_high) C_high = B;
                    if (C_low > C_high) continue;

                    long long Lmin = 11 - A2 - B2;
                    long long N0 = A2 + B2 + 2 * A + 2 * B + 3;

                    if (Lmin <= 0) {
                        long long part = sum_interval(C_low, C_high, N0);
                        local += part;
                        if (local >= MOD) local -= MOD;
                    } else {
                        long long cmin = ceil_sqrt_ll(Lmin);
                        if ((cmin & 1LL) == 0) ++cmin;

                        long long neg_high = C_high < -cmin ? C_high : -cmin;
                        if (C_low <= neg_high) {
                            long long part = sum_interval(C_low, neg_high, N0);
                            local += part;
                            if (local >= MOD) local -= MOD;
                        }

                        long long pos_low = C_low > cmin ? C_low : cmin;
                        if (pos_low <= C_high) {
                            long long part = sum_interval(pos_low, C_high, N0);
                            local += part;
                            if (local >= MOD) local -= MOD;
                        }
                    }
                }
            }
            partial[t] = local;
        });
    }

    for (auto& th : threads) th.join();

    long long ans = 0;
    for (long long v : partial) {
        ans += v;
        if (ans >= MOD) ans -= MOD;
    }
    return ans;
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    auto check = [&](const string& name, long long got, long long expected) {
        if (got != expected) {
            cerr << "Validation failed: " << name << " got " << got
                 << " expected " << expected << "\n";
            exit(1);
        }
    };

    check("S(5)", S_value(5), 48);
    check("S(10^3)", S_value(1000), 37048340);

    long long ans = S_value(100000000LL);
    cout << ans % MOD << "\n";
    return 0;
}

Python

import math

def solve():
    MOD = 433494437
    n = 100000000
    R = 8 * n

    def isqrt(x):
        if x <= 0: return 0
        r = int(math.isqrt(x))
        while (r+1)*(r+1) <= x: r += 1
        while r*r > x: r -= 1
        return r

    def ceil_sqrt(x):
        if x <= 0: return 0
        r = isqrt(x)
        return r if r*r >= x else r+1

    def sum_interval(L, H, N0):
        if L % 2 == 0: L += 1
        if H % 2 == 0: H -= 1
        if L > H: return 0
        cnt = (H - L) // 2 + 1
        t = cnt - 1
        si = cnt * t // 2
        si2 = t * cnt * (2*t+1) // 6
        sC = cnt * L + 2 * si
        sC2 = cnt * L * L + 4 * L * si + 4 * si2
        sn = cnt * N0 + sC2 + 2 * sC
        ss = sn // 2
        return ss % MOD

    A_max = isqrt(R + 3) - 2
    if A_max < 1: return '0'
    if A_max % 2 == 0: A_max -= 1

    ans = 0
    for A in range(1, A_max + 1, 2):
        A2 = A * A
        D = R - A2 - 4*A - 1
        if D < 0: continue
        B_max = isqrt(D) - 2
        if B_max < -1: continue
        if B_max > A: B_max = A

        for B in range(-1, B_max + 1, 2):
            B2 = B * B
            L = R - A2 - B2 - 4*A - 4*B - 5
            if L < 0: continue
            C_max = isqrt(L)
            C_low = max(-C_max, -B - 2)
            C_high = min(C_max, B)
            if C_low > C_high: continue

            Lmin = 11 - A2 - B2
            N0 = A2 + B2 + 2*A + 2*B + 3

            if Lmin <= 0:
                ans = (ans + sum_interval(C_low, C_high, N0)) % MOD
            else:
                cmin = ceil_sqrt(Lmin)
                if cmin % 2 == 0: cmin += 1
                nh = min(C_high, -cmin)
                if C_low <= nh:
                    ans = (ans + sum_interval(C_low, nh, N0)) % MOD
                pl = max(C_low, cmin)
                if pl <= C_high:
                    ans = (ans + sum_interval(pl, C_high, N0)) % MOD

    return str(ans)

if __name__ == '__main__':
    print(solve())

Java

import java.math.BigInteger;
import java.util.concurrent.atomic.AtomicLong;

public class Euler791 {

    static final long MOD = 433494437L;

    static long isqrtLl(long x) {
        if (x <= 0)
            return 0;
        long r = (long) Math.sqrt(x);
        while ((r + 1) * (r + 1) <= x)
            r++;
        while (r * r > x)
            r--;
        return r;
    }

    static long ceilSqrtLl(long x) {
        if (x <= 0)
            return 0;
        long r = isqrtLl(x);
        if (r * r < x)
            r++;
        return r;
    }

    static long sumInterval(long L, long H, long N0) {
        if ((L & 1) == 0)
            L++;
        if ((H & 1) == 0)
            H--;
        if (L > H)
            return 0;

        long cnt = (H - L) / 2 + 1;
        BigInteger biCnt = BigInteger.valueOf(cnt);
        BigInteger biT = BigInteger.valueOf(cnt - 1);
        BigInteger biL = BigInteger.valueOf(L);
        BigInteger biN0 = BigInteger.valueOf(N0);

        BigInteger sumI = biCnt.multiply(biT).divide(BigInteger.TWO);
        BigInteger sumI2 = biT.multiply(biCnt).multiply(biT.multiply(BigInteger.TWO).add(BigInteger.ONE))
                .divide(BigInteger.valueOf(6));

        BigInteger sumC = biCnt.multiply(biL).add(sumI.multiply(BigInteger.TWO));
        BigInteger sumC2 = biCnt.multiply(biL).multiply(biL)
                .add(BigInteger.valueOf(4).multiply(biL).multiply(sumI))
                .add(BigInteger.valueOf(4).multiply(sumI2));

        BigInteger sumNum = biCnt.multiply(biN0).add(sumC2).add(sumC.multiply(BigInteger.TWO));
        BigInteger sumS = sumNum.divide(BigInteger.TWO);

        long mod = sumS.remainder(BigInteger.valueOf(MOD)).longValue();
        if (mod < 0)
            mod += MOD;
        return mod;
    }

    static long sValue(long n) {
        if (n <= 0)
            return 0;
        long R = 8 * n;

        long aMax = isqrtLl(R + 3) - 2;
        if (aMax < 1)
            return 0;
        if ((aMax & 1) == 0)
            aMax--;

        int numThreads = Runtime.getRuntime().availableProcessors();
        if (numThreads <= 0)
            numThreads = 1;

        AtomicLong nextA = new AtomicLong(1);
        long[] partials = new long[numThreads];
        Thread[] threads = new Thread[numThreads];

        final long aMaxFinal = aMax;

        for (int t = 0; t < numThreads; t++) {
            final int tid = t;
            threads[t] = new Thread(() -> {
                long local = 0;
                while (true) {
                    long A = nextA.getAndAdd(2);
                    if (A > aMaxFinal)
                        break;

                    long A2 = A * A;
                    long D = R - A2 - 4 * A - 1;
                    if (D < 0)
                        continue;

                    long B_max = isqrtLl(D) - 2;
                    if (B_max < -1)
                        continue;
                    if (B_max > A)
                        B_max = A;

                    for (long B = -1; B <= B_max; B += 2) {
                        long B2 = B * B;
                        long L = R - A2 - B2 - 4 * A - 4 * B - 5;
                        if (L < 0)
                            continue;
                        long C_max = isqrtLl(L);

                        long C_low = Math.max(-C_max, -B - 2);
                        long C_high = Math.min(C_max, B);
                        if (C_low > C_high)
                            continue;

                        long Lmin = 11 - A2 - B2;
                        long N0 = A2 + B2 + 2 * A + 2 * B + 3;

                        if (Lmin <= 0) {
                            local += sumInterval(C_low, C_high, N0);
                            if (local >= MOD)
                                local -= MOD;
                        } else {
                            long cmin = ceilSqrtLl(Lmin);
                            if ((cmin & 1) == 0)
                                cmin++;

                            long neg_high = Math.min(C_high, -cmin);
                            if (C_low <= neg_high) {
                                local += sumInterval(C_low, neg_high, N0);
                                if (local >= MOD)
                                    local -= MOD;
                            }

                            long pos_low = Math.max(C_low, cmin);
                            if (pos_low <= C_high) {
                                local += sumInterval(pos_low, C_high, N0);
                                if (local >= MOD)
                                    local -= MOD;
                            }
                        }
                    }
                }
                partials[tid] = local;
            });
            threads[t].start();
        }

        long ans = 0;
        for (int t = 0; t < numThreads; t++) {
            try {
                threads[t].join();
                ans += partials[t];
                if (ans >= MOD)
                    ans -= MOD;
            } catch (InterruptedException e) {
                e.printStackTrace();
            }
        }
        return ans;
    }

    public static String solve() {
        return Long.toString(sValue(100000000L));
    }

    public static void main(String[] args) {
        System.out.println(solve());
    }
}