Problem 621: Expressing an Integer as the Sum of Triangular Numbers
View on Project EulerProject 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
- Problem page: https://projecteuler.net/problem=621
- Triangular numbers: Wikipedia — Triangular number
- Three-square theorem: Wikipedia — Legendre's three-square theorem
- Binary quadratic forms: Wikipedia — Binary quadratic form
- Class numbers: Wikipedia — Class number
- 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());
}
}