Problem 621: Expressing an Integer as the Sum of Triangular Numbers

View on Project Euler

Project Euler Problem 621 Solution

EulerSolve provides an optimized solution for Project Euler Problem 621, Expressing an Integer as the Sum of Triangular Numbers, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary Let $$T_k=\frac{k(k+1)}{2}\qquad (k\ge 0)$$ be the \(k\)-th triangular number. The function \(G(n)\) counts ordered triples \((a,b,c)\) of nonnegative integers such that $$T_a+T_b+T_c=n.$$ The task is to evaluate \(G(17526\cdot 10^9)\). Directly enumerating triangular triples is hopeless at that scale, so the solution converts the problem into counting representations by three squares and then uses a class-number formula. Mathematical Approach The implementation turns the additive problem for triangular numbers into an arithmetic problem for quadratic forms. Step 1: Transform Triangular Numbers into Odd Squares The identity $$8T_k+1=4k(k+1)+1=(2k+1)^2$$ shows that every triangular number corresponds to an odd square. Therefore $$n=T_a+T_b+T_c$$ is equivalent to $$8n+3=(2a+1)^2+(2b+1)^2+(2c+1)^2.$$ If we write $$N=8n+3,$$ then \(G(n)\) is the number of ordered representations of \(N\) as a sum of three positive odd squares. Step 2: Relate \(G(n)\) to the Classical Three-Squares Counting Function Define $$r_3(N)=\#\left\{(x,y,z)\in\mathbb{Z}^3:x^2+y^2+z^2=N\right\},$$ where order and signs are both counted....

Detailed mathematical approach

Problem Summary

Let

$$T_k=\frac{k(k+1)}{2}\qquad (k\ge 0)$$

be the \(k\)-th triangular number. The function \(G(n)\) counts ordered triples \((a,b,c)\) of nonnegative integers such that

$$T_a+T_b+T_c=n.$$

The task is to evaluate \(G(17526\cdot 10^9)\). Directly enumerating triangular triples is hopeless at that scale, so the solution converts the problem into counting representations by three squares and then uses a class-number formula.

Mathematical Approach

The implementation turns the additive problem for triangular numbers into an arithmetic problem for quadratic forms.

Step 1: Transform Triangular Numbers into Odd Squares

The identity

$$8T_k+1=4k(k+1)+1=(2k+1)^2$$

shows that every triangular number corresponds to an odd square. Therefore

$$n=T_a+T_b+T_c$$

is equivalent to

$$8n+3=(2a+1)^2+(2b+1)^2+(2c+1)^2.$$

If we write

$$N=8n+3,$$

then \(G(n)\) is the number of ordered representations of \(N\) as a sum of three positive odd squares.

Step 2: Relate \(G(n)\) to the Classical Three-Squares Counting Function

Define

$$r_3(N)=\#\left\{(x,y,z)\in\mathbb{Z}^3:x^2+y^2+z^2=N\right\},$$

where order and signs are both counted. Since \(N\equiv 3 \pmod{8}\), each square in any representation must be congruent to \(1 \pmod{8}\), because the only way to obtain \(3\) from three quadratic residues modulo \(8\) is

$$1+1+1\equiv 3 \pmod{8}.$$

So every valid \(x,y,z\) is odd, and none of them can be zero. Each positive odd triple corresponds to exactly \(2^3=8\) signed triples in \(r_3(N)\). Hence

$$G(n)=\frac{r_3(8n+3)}{8}.$$

Step 3: Use the Arithmetic Formula for \(r_3(N)\)

Factor \(N\) as

$$N=f^2m,$$

where \(m\) is squarefree. Because \(N\equiv 3\pmod{8}\), the squarefree part \(m\) is odd and satisfies \(m\equiv 3\pmod{4}\). Set

$$D=-m.$$

Then \(D\) is a negative fundamental discriminant, and the implementation evaluates

$$r_3(N)=\frac{12\cdot 2h(D)\left(1-\left(\frac{D}{2}\right)\right)}{w(D)}\sum_{d\mid f}\mu(d)\left(\frac{D}{d}\right)\sigma\left(\frac{f}{d}\right).$$

Here \(h(D)\) is the class number of discriminant \(D\), \(\mu\) is the Möbius function, \(\sigma\) is the sum-of-divisors function, \(\left(\frac{D}{d}\right)\) is the Jacobi or Kronecker character, and \(w(D)\) is the number of units in the quadratic order:

$$w(D)=\begin{cases} 6,& D=-3,\\ 4,& D=-4,\\ 2,& \text{otherwise.} \end{cases}$$

The sum is multiplicative in the square part \(f\), which is why the factorization of \(N\) is the first major step of the program.

Step 4: Compute \(h(D)\) by Counting Reduced Binary Quadratic Forms

For a negative fundamental discriminant \(D\), the class number can be obtained by counting reduced positive definite binary quadratic forms

$$ax^2+bxy+cy^2,\qquad b^2-4ac=D.$$

The reduced conditions used by the implementation are

$$|b|\le a\le c,$$

$$b\ge 0\quad\text{when}\quad |b|=a\ \text{or}\ a=c.$$

Since \(D<0\), one only needs to scan

$$1\le a\le \sqrt{\frac{|D|}{3}}.$$

For each odd \(a\), the congruence condition coming from the discriminant is

$$b^2\equiv D \pmod{4a}.$$

The implementation factors \(a\), finds square roots of \(D\) modulo each prime power, lifts them to higher powers, combines the local roots with the Chinese remainder theorem, and then reconstructs the admissible values of \(b\) and \(c\). Every valid reduced form contributes exactly one class.

Step 5: Worked Example \(n=9\)

For \(n=9\), we have

$$N=8\cdot 9+3=75.$$

The positive odd-square representations of \(75\) are

$$75=1^2+5^2+7^2=5^2+5^2+5^2.$$

The first pattern has \(3!=6\) orderings, and each ordering has \(8\) sign choices in \(r_3(75)\). The second pattern contributes \(1\cdot 8\). Therefore

$$r_3(75)=6\cdot 8+1\cdot 8=56,$$

so

$$G(9)=\frac{56}{8}=7.$$

In triangular-number language, those seven ordered representations are the six permutations of \(0+3+6\) and the single representation \(3+3+3\).

Step 6: Final Evaluation Strategy

For the target input, the program forms

$$N=8\cdot (17526\cdot 10^9)+3=140208000000003,$$

factors \(N\), extracts \(f\) and \(m\), computes \(h(-m)\), evaluates the divisor sum in the formula for \(r_3(N)\), and finally divides by \(8\) to obtain \(G(n)\). This avoids any search over triangular numbers themselves.

How the Code Works

The C++, Python, and Java implementations follow the same pipeline. First they factor \(N=8n+3\) with a 64-bit primality test and Pollard-Rho splitting. From that factorization they separate the squarefree part \(m\) from the square part \(f^2\), which gives the discriminant \(D=-m\).

Next they compute the class number \(h(D)\). To do that efficiently, the implementation enumerates candidate values of \(a\), factors each \(a\) with a small-prime sieve, solves the modular square-root conditions prime-power by prime-power, merges the local solutions with the Chinese remainder theorem, and counts the reduced forms that satisfy the discriminant equation.

Once \(h(D)\) is known, the code evaluates the divisor sum

$$\sum_{d\mid f}\mu(d)\left(\frac{D}{d}\right)\sigma\left(\frac{f}{d}\right)$$

by iterating over the squarefree divisors encoded by the prime factors of \(f\). That produces \(r_3(N)\), and the final answer is \(r_3(N)/8\). The implementations also include checkpoint values such as \(G(9)=7\), \(G(1000)=78\), and \(G(10^6)=2106\).

Complexity Analysis

The brute-force search space for three triangular numbers grows far too quickly, so the arithmetic method is essential. Integer factorization is handled by Miller-Rabin and Pollard-Rho, which are very fast in practice for a single 64-bit input. After factorization, the dominant deterministic work is the class-number computation.

If

$$A=\left\lfloor\sqrt{\frac{|D|}{3}}\right\rfloor,$$

then the sieve used inside the class-number computation requires \(O(A)\) memory and \(O(A\log\log A)\) preprocessing time. The subsequent enumeration scans odd \(a\le A\); each step performs a small factorization of \(a\), a few modular square-root lifts, and a small CRT merge. In practice this is easily manageable for a single target value and is incomparably faster than enumerating triangular triples directly.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=621
  2. Triangular numbers: Wikipedia — Triangular number
  3. Three-square theorem: Wikipedia — Legendre's three-square theorem
  4. Binary quadratic forms: Wikipedia — Binary quadratic form
  5. Class numbers: Wikipedia — Class number
  6. Jacobi symbol: Wikipedia — Jacobi symbol

Problem 621 source code

C++

#include <algorithm>
#include <cassert>
#include <cstdint>
#include <cstdlib>
#include <iostream>
#include <map>
#include <numeric>
#include <random>
#include <utility>
#include <vector>

using i64 = long long;
using u64 = unsigned long long;
using i128 = __int128_t;
using u128 = __uint128_t;

static u64 mod_mul(u64 a, u64 b, u64 mod) { return (u128)a * b % mod; }

static u64 mod_pow(u64 a, u64 e, u64 mod) {
    u64 r = 1 % mod;
    while (e) {
        if (e & 1) r = mod_mul(r, a, mod);
        a = mod_mul(a, a, mod);
        e >>= 1;
    }
    return r;
}

static bool is_prime(u64 n) {
    if (n < 2) return false;
    for (u64 p : {2ULL, 3ULL, 5ULL, 7ULL, 11ULL, 13ULL, 17ULL, 19ULL, 23ULL, 29ULL, 31ULL, 37ULL}) {
        if (n % p == 0) return n == p;
    }

    u64 d = n - 1;
    int s = 0;
    while ((d & 1) == 0) {
        d >>= 1;
        ++s;
    }

    auto witness = [&](u64 a) -> bool {
        if (a % n == 0) return false;
        u64 x = mod_pow(a, d, n);
        if (x == 1 || x == n - 1) return false;
        for (int i = 1; i < s; ++i) {
            x = mod_mul(x, x, n);
            if (x == n - 1) return false;
        }
        return true;
    };

    for (u64 a : {2ULL, 325ULL, 9375ULL, 28178ULL, 450775ULL, 9780504ULL, 1795265022ULL}) {
        if (witness(a)) return false;
    }
    return true;
}

static u64 pollard_rho(u64 n) {
    if ((n & 1ULL) == 0) return 2;
    static std::mt19937_64 rng(1234567);
    std::uniform_int_distribution<u64> dist(0, n - 1);

    while (true) {
        u64 c = dist(rng) % n;
        u64 x = dist(rng) % n;
        u64 y = x;
        u64 d = 1;

        auto f = [&](u64 v) { return (mod_mul(v, v, n) + c) % n; };

        while (d == 1) {
            x = f(x);
            y = f(f(y));
            u64 diff = (x > y) ? (x - y) : (y - x);
            d = std::gcd(diff, n);
        }
        if (d != n) return d;
    }
}

static void factor_rec(u64 n, std::map<u64, int> &out) {
    if (n == 1) return;
    if (is_prime(n)) {
        out[n]++;
        return;
    }
    u64 d = pollard_rho(n);
    factor_rec(d, out);
    factor_rec(n / d, out);
}

static u64 tonelli_shanks(u64 n, u64 p) {
    if (n == 0) return 0;
    if (p == 2) return n;

    if (mod_pow(n, (p - 1) / 2, p) != 1) return 0;
    if (p % 4 == 3) return mod_pow(n, (p + 1) / 4, p);

    u64 q = p - 1;
    int s = 0;
    while ((q & 1) == 0) {
        q >>= 1;
        ++s;
    }

    u64 z = 2;
    while (mod_pow(z, (p - 1) / 2, p) != p - 1) ++z;

    u64 c = mod_pow(z, q, p);
    u64 r = mod_pow(n, (q + 1) / 2, p);
    u64 t = mod_pow(n, q, p);
    int m = s;

    while (t != 1) {
        int i = 1;
        u64 tt = mod_mul(t, t, p);
        while (i < m && tt != 1) {
            tt = mod_mul(tt, tt, p);
            ++i;
        }
        u64 b = mod_pow(c, 1ULL << (m - i - 1), p);
        r = mod_mul(r, b, p);
        u64 b2 = mod_mul(b, b, p);
        t = mod_mul(t, b2, p);
        c = b2;
        m = i;
    }
    return r;
}

static u64 mod_norm_i64(i64 a, u64 mod) {
    i128 r = (i128)a % (i128)mod;
    if (r < 0) r += (i128)mod;
    return (u64)r;
}

static u64 inv_mod_u64(u64 a, u64 mod) {
    i128 t = 0, newt = 1;
    i128 r = (i128)mod, newr = (i128)a;
    while (newr != 0) {
        u64 q = (u64)(r / newr);
        i128 tmp = t - (i128)q * newt;
        t = newt;
        newt = tmp;
        tmp = r - (i128)q * newr;
        r = newr;
        newr = tmp;
    }
    assert(r == 1);
    if (t < 0) t += (i128)mod;
    return (u64)t;
}

static std::vector<u64> roots_mod_prime_power(i64 D, u64 p, int e) {
    u64 pe = 1;
    for (int i = 0; i < e; ++i) pe *= p;

    if (mod_norm_i64(D, p) == 0) {
        if (e == 1) return {0};
        return {};
    }

    u64 n = mod_norm_i64(D, p);
    u64 r = tonelli_shanks(n, p);
    if (r == 0) return {};

    u64 pk = p;
    for (int k = 1; k < e; ++k) {
        u64 pk1 = pk * p;
        u64 Dmod = mod_norm_i64(D, pk1);
        u64 r2 = (u128)r * r % pk1;
        u64 diff = (Dmod + pk1 - r2) % pk1;
        u64 delta = diff / pk;

        u64 inv = inv_mod_u64((2 * (r % p)) % p, p);
        u64 t = (u128)delta * inv % p;
        r += t * pk;
        pk = pk1;
    }

    r %= pe;
    u64 r2 = (pe - r) % pe;
    if (r2 == r) return {r};
    return {r, r2};
}

static std::vector<int> sieve_spf(int n) {
    std::vector<int> spf(n + 1, 0), primes;
    primes.reserve((size_t)n / 10);
    for (int i = 2; i <= n; ++i) {
        if (spf[i] == 0) {
            spf[i] = i;
            primes.push_back(i);
        }
        for (int p : primes) {
            long long v = 1LL * p * i;
            if (v > n) break;
            spf[(int)v] = p;
            if (p == spf[i]) break;
        }
    }
    return spf;
}

static std::vector<std::pair<int, int>> factorize_small(int x, const std::vector<int> &spf) {
    std::vector<std::pair<int, int>> f;
    while (x > 1) {
        int p = spf[x];
        int e = 0;
        while (x % p == 0) {
            x /= p;
            ++e;
        }
        f.push_back({p, e});
    }
    return f;
}

static i64 class_number_odd_fundamental(i64 D) {
    assert(D < 0);
    assert((D & 3LL) == 1); // D ≡ 1 (mod 4)
    const u64 absD = (u64)(-D);

    i64 maxA = (i64)std::sqrt((long double)absD / 3.0L);
    while (3 * (i128)(maxA + 1) * (maxA + 1) <= (i128)absD) ++maxA;
    while (3 * (i128)maxA * maxA > (i128)absD) --maxA;

    const int lim = (int)maxA;
    const std::vector<int> spf = sieve_spf(lim);

    i64 h = 0;
    for (int a = 1; a <= lim; a += 2) {
        const auto fac = factorize_small(a, spf);
        u64 mod = 1;
        std::vector<u64> roots = {0};
        bool ok = true;

        for (auto [p_i, e] : fac) {
            u64 p = (u64)p_i;
            u64 pe = 1;
            for (int i = 0; i < e; ++i) pe *= p;

            std::vector<u64> rpe = roots_mod_prime_power(D, p, e);
            if (rpe.empty()) {
                ok = false;
                break;
            }

            u64 inv = inv_mod_u64(mod % pe, pe);
            std::vector<u64> next;
            next.reserve(roots.size() * rpe.size());

            for (u64 x : roots) {
                u64 xpe = x % pe;
                for (u64 y : rpe) {
                    u64 t = (y + pe - xpe) % pe;
                    t = (u128)t * inv % pe;
                    next.push_back(x + mod * t);
                }
            }

            mod *= pe;
            roots.swap(next);
        }

        if (!ok) continue;
        assert(mod == (u64)a);

        for (u64 r : roots) {
            u64 bmod = r;
            if ((bmod & 1ULL) == 0) bmod += (u64)a;
            bmod %= 2ULL * (u64)a;

            i64 b = (i64)bmod;
            if (b > a) b -= 2LL * a;
            if (b == -a) b = a;

            const i128 num = (i128)b * b - (i128)D;
            const i128 den = 4 * (i128)a;
            if (num % den != 0) continue;
            const i64 c = (i64)(num / den);
            if (a > c) continue;
            if ((a == c || std::llabs(b) == a) && b < 0) continue;

            ++h;
        }
    }
    return h;
}

static int jacobi_symbol(i64 a, i64 n) {
    assert(n > 0 && (n & 1LL) == 1);
    a %= n;
    if (a < 0) a += n;
    int s = 1;
    while (a != 0) {
        while ((a & 1LL) == 0) {
            a >>= 1;
            i64 r = n & 7LL;
            if (r == 3 || r == 5) s = -s;
        }
        std::swap(a, n);
        if ((a & 3LL) == 3 && (n & 3LL) == 3) s = -s;
        a %= n;
    }
    return (n == 1) ? s : 0;
}

static int kronecker_D_over_2(i64 D) {
    assert(D & 1LL);
    i64 r = D & 7LL;
    if (r == 1 || r == 7) return 1;
    if (r == 3 || r == 5) return -1;
    assert(false);
    return 0;
}

static u64 sigma_prime_power(u64 p, int e) {
    u128 pe1 = 1;
    for (int i = 0; i < e + 1; ++i) pe1 *= p;
    return (u64)((pe1 - 1) / (p - 1));
}

static u64 r3_sum_of_three_squares(u64 N) {
    std::map<u64, int> pf;
    factor_rec(N, pf);

    u64 m = 1;
    std::map<u64, int> f_pf;
    for (auto [p, e] : pf) {
        if (e & 1) m *= p;
        int k = e / 2;
        if (k > 0) f_pf[p] = k;
    }

    const i64 D = -(i64)m;
    const i64 h = class_number_odd_fundamental(D);
    const int w = (D == -3) ? 6 : ((D == -4) ? 4 : 2);
    const int one_minus = 1 - kronecker_D_over_2(D);

    std::vector<u64> primes;
    primes.reserve(f_pf.size());
    std::vector<int> exps;
    exps.reserve(f_pf.size());
    for (auto [p, e] : f_pf) {
        primes.push_back(p);
        exps.push_back(e);
    }

    i128 S = 0;
    const int k = (int)primes.size();
    for (int mask = 0; mask < (1 << k); ++mask) {
        u64 d = 1;
        int mu = 1;
        u128 sigma = 1;
        for (int i = 0; i < k; ++i) {
            int e = exps[i];
            if (mask & (1 << i)) {
                d *= primes[i];
                mu = -mu;
                --e;
            }
            sigma *= sigma_prime_power(primes[i], e);
        }
        int chi = jacobi_symbol(D, (i64)d);
        if (chi == 0) continue;
        S += (i128)mu * (i128)chi * (i128)sigma;
    }

    const i128 t = (i128)12 * (i128)(2 * h) * (i128)one_minus * S;
    assert(t % w == 0);
    const i128 r3 = t / w;
    assert(r3 >= 0);
    return (u64)r3;
}

static i64 brute_r3(i64 N) {
    int lim = (int)std::sqrt((long double)N);
    std::vector<unsigned char> is_sq((size_t)N + 1, 0);
    for (int i = 0; i <= lim; ++i) is_sq[(size_t)i * i] = 1;

    i64 cnt = 0;
    for (int x = -lim; x <= lim; ++x) {
        i64 x2 = (i64)x * x;
        for (int y = -lim; y <= lim; ++y) {
            i64 rem = N - x2 - (i64)y * y;
            if (rem < 0) continue;
            if (!is_sq[(size_t)rem]) continue;
            cnt += (rem == 0) ? 1 : 2;
        }
    }
    return cnt;
}

static u64 G(u64 n) {
    const u64 N = 8 * n + 3;
    const u64 r3 = r3_sum_of_three_squares(N);
    assert(r3 % 8 == 0);
    return r3 / 8;
}

int main() {
    assert(r3_sum_of_three_squares(75) == (u64)brute_r3(75));
    assert(r3_sum_of_three_squares(8003) == (u64)brute_r3(8003));

    assert(G(9) == 7);
    assert(G(1000) == 78);
    assert(G(1000000ULL) == 2106);

    const u64 n = 17526ULL * 1000000000ULL;
    std::cout << G(n) << "\n";
    return 0;
}

Python

import math
import random

def mod_pow(a, e, mod):
    return pow(a, e, mod)

def is_prime(n):
    if n < 2: return False
    for p in [2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37]:
        if n % p == 0: return n == p
    d = n - 1
    s = 0
    while d % 2 == 0:
        d //= 2
        s += 1
    def witness(a):
        if a % n == 0: return False
        x = pow(a, d, n)
        if x == 1 or x == n - 1: return False
        for _ in range(1, s):
            x = pow(x, 2, n)
            if x == n - 1: return False
        return True
    for a in [2, 325, 9375, 28178, 450775, 9780504, 1795265022]:
        if witness(a): return False
    return True

def pollard_rho(n):
    if n % 2 == 0: return 2
    while True:
        c = random.randint(0, n - 1)
        x = random.randint(0, n - 1)
        y = x
        d = 1
        f = lambda v: (v * v + c) % n
        while d == 1:
            x = f(x)
            y = f(f(y))
            diff = x - y if x > y else y - x
            d = math.gcd(diff, n)
        if d != n: return d

def factor_rec(n, out):
    if n == 1: return
    if is_prime(n):
        out[n] = out.get(n, 0) + 1
        return
    d = pollard_rho(n)
    factor_rec(d, out)
    factor_rec(n // d, out)

def tonelli_shanks(n, p):
    if n == 0: return 0
    if p == 2: return n
    if pow(n, (p - 1) // 2, p) != 1: return 0
    if p % 4 == 3: return pow(n, (p + 1) // 4, p)
    q = p - 1
    s = 0
    while q % 2 == 0:
        q //= 2
        s += 1
    z = 2
    while pow(z, (p - 1) // 2, p) != p - 1: z += 1
    c = pow(z, q, p)
    r = pow(n, (q + 1) // 2, p)
    t = pow(n, q, p)
    m = s
    while t != 1:
        i = 1
        tt = (t * t) % p
        while i < m and tt != 1:
            tt = (tt * tt) % p
            i += 1
        b = pow(c, 1 << (m - i - 1), p)
        r = (r * b) % p
        b2 = (b * b) % p
        t = (t * b2) % p
        c = b2
        m = i
    return r

def roots_mod_prime_power(D, p, e):
    pe = p ** e
    Dmod = D % p
    if Dmod == 0:
        if e == 1: return [0]
        return []
    n = Dmod
    r = tonelli_shanks(n, p)
    if r == 0: return []
    pk = p
    for k in range(1, e):
        pk1 = pk * p
        Dmod_k1 = D % pk1
        r2 = (r * r) % pk1
        diff = (Dmod_k1 - r2) % pk1
        delta = diff // pk
        inv = pow((2 * r) % p, -1, p)
        t = (delta * inv) % p
        r += t * pk
        pk = pk1
    r %= pe
    r2 = (pe - r) % pe
    if r2 == r: return [r]
    return [r, r2]

def sieve_spf(n):
    spf = [0] * (n + 1)
    primes = []
    for i in range(2, n + 1):
        if spf[i] == 0:
            spf[i] = i
            primes.append(i)
        for p in primes:
            v = p * i
            if v > n: break
            spf[v] = p
            if p == spf[i]: break
    return spf

def factorize_small(x, spf):
    f = []
    while x > 1:
        p = spf[x]
        e = 0
        while x % p == 0:
            x //= p
            e += 1
        f.append((p, e))
    return f

def class_number_odd_fundamental(D):
    absD = -D
    maxA = int(math.sqrt(absD / 3.0))
    while 3 * (maxA + 1)**2 <= absD: maxA += 1
    while 3 * maxA**2 > absD: maxA -= 1
    lim = maxA
    spf = sieve_spf(lim)
    h = 0
    for a in range(1, lim + 1, 2):
        fac = factorize_small(a, spf)
        mod = 1
        roots = [0]
        ok = True
        for p, e in fac:
            pe = p ** e
            rpe = roots_mod_prime_power(D, p, e)
            if not rpe:
                ok = False
                break
            inv = pow(mod % pe, -1, pe)
            nxt = []
            for x in roots:
                xpe = x % pe
                for y in rpe:
                    t = (y - xpe) % pe
                    t = (t * inv) % pe
                    nxt.append(x + mod * t)
            mod *= pe
            roots = nxt
        if not ok: continue
        for r in roots:
            bmod = r
            if bmod % 2 == 0: bmod += a
            bmod %= (2 * a)
            b = bmod
            if b > a: b -= 2 * a
            if b == -a: b = a
            
            num = b * b - D
            den = 4 * a
            if num % den != 0: continue
            c = num // den
            if a > c: continue
            if (a == c or abs(b) == a) and b < 0: continue
            h += 1
    return h

def jacobi_symbol(a, n):
    a %= n
    if a < 0: a += n
    s = 1
    while a != 0:
        while a % 2 == 0:
            a //= 2
            r = n % 8
            if r == 3 or r == 5: s = -s
        a, n = n, a
        if a % 4 == 3 and n % 4 == 3: s = -s
        a %= n
    return s if n == 1 else 0

def kronecker_D_over_2(D):
    r = D % 8
    if r in (1, 7): return 1
    if r in (3, 5): return -1
    return 0

def sigma_prime_power(p, e):
    pe1 = p ** (e + 1)
    return (pe1 - 1) // (p - 1)

def r3_sum_of_three_squares(N):
    pf = {}
    factor_rec(N, pf)
    m = 1
    f_pf = {}
    for p, e in pf.items():
        if e % 2 == 1: m *= p
        k = e // 2
        if k > 0: f_pf[p] = k
    
    D = -m
    h = class_number_odd_fundamental(D)
    w = 6 if D == -3 else (4 if D == -4 else 2)
    one_minus = 1 - kronecker_D_over_2(D)
    
    primes = list(f_pf.keys())
    exps = list(f_pf.values())
    k_len = len(primes)
    
    S = 0
    for mask in range(1 << k_len):
        d = 1
        mu = 1
        sigma = 1
        for i in range(k_len):
            e = exps[i]
            if (mask & (1 << i)):
                d *= primes[i]
                mu = -mu
                e -= 1
            sigma *= sigma_prime_power(primes[i], e)
        chi = jacobi_symbol(D, d)
        if chi == 0: continue
        S += mu * chi * sigma
        
    t = 12 * (2 * h) * one_minus * S
    r3 = t // w
    return r3

def G(n):
    N = 8 * n + 3
    r3 = r3_sum_of_three_squares(N)
    return r3 // 8

def solve():
    n = 17526 * 1000000000
    r = G(n)
    return str(r)

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

Java

import java.math.BigInteger;
import java.util.*;

public class Euler621 {
    static final BigInteger TWO = BigInteger.valueOf(2);
    static final BigInteger THREE = BigInteger.valueOf(3);

    static boolean isPrime(long n) {
        if (n < 2)
            return false;
        long[] bases = { 2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37 };
        for (long p : bases) {
            if (n % p == 0)
                return n == p;
        }
        BigInteger bn = BigInteger.valueOf(n);
        return bn.isProbablePrime(10);
    }

    static long pollardRho(long n) {
        if ((n & 1) == 0)
            return 2;
        BigInteger bn = BigInteger.valueOf(n);
        Random rng = new Random();
        BigInteger c = new BigInteger(bn.bitLength(), rng).mod(bn);
        BigInteger x = new BigInteger(bn.bitLength(), rng).mod(bn);
        BigInteger y = x;
        BigInteger d = BigInteger.ONE;
        while (d.equals(BigInteger.ONE)) {
            x = x.multiply(x).add(c).mod(bn);
            y = y.multiply(y).add(c).mod(bn);
            y = y.multiply(y).add(c).mod(bn);
            d = x.subtract(y).abs().gcd(bn);
        }
        if (!d.equals(bn))
            return d.longValue();
        return pollardRho(n);
    }

    static void factorRec(long n, Map<Long, Integer> out) {
        if (n == 1)
            return;
        if (isPrime(n)) {
            out.put(n, out.getOrDefault(n, 0) + 1);
            return;
        }
        long d = pollardRho(n);
        factorRec(d, out);
        factorRec(n / d, out);
    }

    static BigInteger tonelliShanks(BigInteger n, BigInteger p) {
        if (n.equals(BigInteger.ZERO))
            return BigInteger.ZERO;
        if (p.equals(TWO))
            return n;
        if (!n.modPow(p.subtract(BigInteger.ONE).divide(TWO), p).equals(BigInteger.ONE))
            return BigInteger.ZERO;
        if (p.mod(BigInteger.valueOf(4)).equals(THREE)) {
            return n.modPow(p.add(BigInteger.ONE).divide(BigInteger.valueOf(4)), p);
        }
        BigInteger q = p.subtract(BigInteger.ONE);
        int s = 0;
        while (q.mod(TWO).equals(BigInteger.ZERO)) {
            q = q.divide(TWO);
            s++;
        }
        BigInteger z = TWO;
        while (!z.modPow(p.subtract(BigInteger.ONE).divide(TWO), p).equals(p.subtract(BigInteger.ONE))) {
            z = z.add(BigInteger.ONE);
        }
        BigInteger c = z.modPow(q, p);
        BigInteger r = n.modPow(q.add(BigInteger.ONE).divide(TWO), p);
        BigInteger t = n.modPow(q, p);
        int m = s;

        while (!t.equals(BigInteger.ONE)) {
            int i = 1;
            BigInteger tt = t.multiply(t).mod(p);
            while (i < m && !tt.equals(BigInteger.ONE)) {
                tt = tt.multiply(tt).mod(p);
                i++;
            }
            BigInteger b = c.modPow(TWO.pow(m - i - 1), p);
            r = r.multiply(b).mod(p);
            BigInteger b2 = b.multiply(b).mod(p);
            t = t.multiply(b2).mod(p);
            c = b2;
            m = i;
        }
        return r;
    }

    static ArrayList<Long> rootsModPrimePower(long D, long p, int e) {
        BigInteger bp = BigInteger.valueOf(p);
        BigInteger pe = bp.pow(e);
        BigInteger bD = BigInteger.valueOf(D);
        BigInteger Dmod = bD.mod(bp);
        if (Dmod.signum() < 0)
            Dmod = Dmod.add(bp);

        ArrayList<Long> res = new ArrayList<>();
        if (Dmod.equals(BigInteger.ZERO)) {
            if (e == 1) {
                res.add(0L);
                return res;
            }
            return res;
        }

        BigInteger r = tonelliShanks(Dmod, bp);
        if (r.equals(BigInteger.ZERO))
            return res;

        BigInteger pk = bp;
        for (int k = 1; k < e; k++) {
            BigInteger pk1 = pk.multiply(bp);
            BigInteger Dmod_k1 = bD.mod(pk1);
            if (Dmod_k1.signum() < 0)
                Dmod_k1 = Dmod_k1.add(pk1);
            BigInteger r2 = r.multiply(r).mod(pk1);
            BigInteger diff = Dmod_k1.subtract(r2).mod(pk1);
            if (diff.signum() < 0)
                diff = diff.add(pk1);
            BigInteger delta = diff.divide(pk);
            BigInteger inv = r.multiply(TWO).modInverse(bp);
            BigInteger t = delta.multiply(inv).mod(bp);
            r = r.add(t.multiply(pk));
            pk = pk1;
        }
        r = r.mod(pe);
        BigInteger r2 = pe.subtract(r).mod(pe);
        res.add(r.longValue());
        if (!r2.equals(r))
            res.add(r2.longValue());
        return res;
    }

    static int[] sieveSpf(int n) {
        int[] spf = new int[n + 1];
        ArrayList<Integer> primes = new ArrayList<>();
        for (int i = 2; i <= n; i++) {
            if (spf[i] == 0) {
                spf[i] = i;
                primes.add(i);
            }
            for (int p : primes) {
                if ((long) p * i > n)
                    break;
                spf[p * i] = p;
                if (p == spf[i])
                    break;
            }
        }
        return spf;
    }

    static ArrayList<int[]> factorizeSmall(int x, int[] spf) {
        ArrayList<int[]> f = new ArrayList<>();
        while (x > 1) {
            int p = spf[x];
            int e = 0;
            while (x % p == 0) {
                x /= p;
                e++;
            }
            f.add(new int[] { p, e });
        }
        return f;
    }

    static long classNumberOddFundamental(long D) {
        long absD = -D;
        long maxA = (long) Math.sqrt(absD / 3.0);
        while (3 * (maxA + 1) * (maxA + 1) <= absD)
            maxA++;
        while (3 * maxA * maxA > absD)
            maxA--;

        int lim = (int) maxA;
        int[] spf = sieveSpf(lim);
        long h = 0;

        for (int a = 1; a <= lim; a += 2) {
            ArrayList<int[]> fac = factorizeSmall(a, spf);
            BigInteger mod = BigInteger.ONE;
            ArrayList<Long> roots = new ArrayList<>();
            roots.add(0L);
            boolean ok = true;

            for (int[] pe_arr : fac) {
                long p = pe_arr[0];
                int e = pe_arr[1];
                BigInteger pe = BigInteger.valueOf(p).pow(e);
                ArrayList<Long> rpe = rootsModPrimePower(D, p, e);
                if (rpe.isEmpty()) {
                    ok = false;
                    break;
                }
                BigInteger inv = mod.mod(pe).modInverse(pe);
                ArrayList<Long> nxt = new ArrayList<>();
                for (long x : roots) {
                    BigInteger bx = BigInteger.valueOf(x);
                    BigInteger xpe = bx.mod(pe);
                    for (long y : rpe) {
                        BigInteger t = BigInteger.valueOf(y).subtract(xpe).mod(pe);
                        if (t.signum() < 0)
                            t = t.add(pe);
                        t = t.multiply(inv).mod(pe);
                        nxt.add(bx.add(mod.multiply(t)).longValue());
                    }
                }
                mod = mod.multiply(pe);
                roots = nxt;
            }
            if (!ok)
                continue;

            for (long r : roots) {
                long bmod = r;
                if (bmod % 2 == 0)
                    bmod += a;
                bmod %= (2L * a);
                long b = bmod;
                if (b > a)
                    b -= 2L * a;
                if (b == -a)
                    b = a;

                BigInteger bb = BigInteger.valueOf(b);
                BigInteger num = bb.multiply(bb).subtract(BigInteger.valueOf(D));
                long den = 4L * a;
                if (!num.mod(BigInteger.valueOf(den)).equals(BigInteger.ZERO))
                    continue;
                long c = num.divide(BigInteger.valueOf(den)).longValue();
                if (a > c)
                    continue;
                if ((a == c || Math.abs(b) == a) && b < 0)
                    continue;
                h++;
            }
        }
        return h;
    }

    static int jacobiSymbol(long a, long n) {
        a %= n;
        if (a < 0)
            a += n;
        int s = 1;
        while (a != 0) {
            while (a % 2 == 0) {
                a /= 2;
                long r = n % 8;
                if (r == 3 || r == 5)
                    s = -s;
            }
            long temp = a;
            a = n;
            n = temp;
            if (a % 4 == 3 && n % 4 == 3)
                s = -s;
            a %= n;
        }
        return n == 1 ? s : 0;
    }

    static int kroneckerDOver2(long D) {
        long r = D % 8;
        if (r < 0)
            r += 8;
        if (r == 1 || r == 7)
            return 1;
        if (r == 3 || r == 5)
            return -1;
        return 0;
    }

    static BigInteger sigmaPrimePower(long p, int e) {
        BigInteger bp = BigInteger.valueOf(p);
        BigInteger pe1 = bp.pow(e + 1);
        return pe1.subtract(BigInteger.ONE).divide(bp.subtract(BigInteger.ONE));
    }

    static long r3SumOfThreeSquares(long N) {
        Map<Long, Integer> pf = new HashMap<>();
        factorRec(N, pf);
        long m = 1;
        ArrayList<Long> primes = new ArrayList<>();
        ArrayList<Integer> exps = new ArrayList<>();
        for (Map.Entry<Long, Integer> entry : pf.entrySet()) {
            long p = entry.getKey();
            int e = entry.getValue();
            if (e % 2 == 1)
                m *= p;
            int k = e / 2;
            if (k > 0) {
                primes.add(p);
                exps.add(k);
            }
        }

        long D = -m;
        long h = classNumberOddFundamental(D);
        int w = (D == -3) ? 6 : ((D == -4) ? 4 : 2);
        int oneMinus = 1 - kroneckerDOver2(D);

        int kLen = primes.size();
        BigInteger S = BigInteger.ZERO;

        for (int mask = 0; mask < (1 << kLen); mask++) {
            long d = 1;
            int mu = 1;
            BigInteger sigma = BigInteger.ONE;
            for (int i = 0; i < kLen; i++) {
                int e = exps.get(i);
                if ((mask & (1 << i)) != 0) {
                    d *= primes.get(i);
                    mu = -mu;
                    e--;
                }
                sigma = sigma.multiply(sigmaPrimePower(primes.get(i), e));
            }
            int chi = jacobiSymbol(D, d);
            if (chi == 0)
                continue;
            S = S.add(sigma.multiply(BigInteger.valueOf(mu)).multiply(BigInteger.valueOf(chi)));
        }

        BigInteger t = S.multiply(BigInteger.valueOf(12L * (2 * h) * oneMinus));
        return t.divide(BigInteger.valueOf(w)).longValue();
    }

    public static String solve() {
        long n = 17526L * 1000000000L;
        long N = 8 * n + 3;
        long r3 = r3SumOfThreeSquares(N);
        return Long.toString(r3 / 8);
    }

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