Problem 596: Number of Lattice Points in a Hyperball
View on Project EulerProject Euler Problem 596 Solution
EulerSolve provides an optimized solution for Project Euler Problem 596, Number of Lattice Points in a Hyperball, with C++, Python, Java, and a step-by-step mathematical explanation.
Problem Summary Let $$T(r)=\#\left\{(x_1,x_2,x_3,x_4)\in\mathbb Z^4: x_1^2+x_2^2+x_3^2+x_4^2\le r^2\right\}.$$ The problem asks for this lattice-point count in four dimensions when \(r=10^8\), with the final value reduced modulo \(10^9+7\). A direct enumeration of all integer quadruples is hopeless, so the solution converts the geometry into a divisor sum that can be evaluated in about \(O(\sqrt{r^2})\) time. Mathematical Approach The key observation is that the hyperball can be counted shell by shell according to the squared distance from the origin. Step 1: Turn the Hyperball into a Summatory Function Set \(N=r^2\). For each integer \(n\ge 0\), let \(r_4(n)\) be the number of integer quadruples satisfying $$x_1^2+x_2^2+x_3^2+x_4^2=n.$$ Every lattice point inside the ball has some squared radius \(n\) with \(0\le n\le N\), so $$T(r)=\sum_{n=0}^{N} r_4(n).$$ This removes the geometric viewpoint and replaces it with a summation problem. Step 2: Apply Jacobi's Four-Square Theorem For \(n>0\), Jacobi's theorem gives $$r_4(n)=8\sum_{\substack{d\mid n\\4\nmid d}} d,$$ and of course \(r_4(0)=1\). Substituting this into the previous identity yields $$T(r)=1+8\sum_{n=1}^{N}\sum_{\substack{d\mid n\\4\nmid d}} d.$$ Now interchange the order of summation....
Detailed mathematical approach
Problem Summary
Let
$$T(r)=\#\left\{(x_1,x_2,x_3,x_4)\in\mathbb Z^4: x_1^2+x_2^2+x_3^2+x_4^2\le r^2\right\}.$$
The problem asks for this lattice-point count in four dimensions when \(r=10^8\), with the final value reduced modulo \(10^9+7\). A direct enumeration of all integer quadruples is hopeless, so the solution converts the geometry into a divisor sum that can be evaluated in about \(O(\sqrt{r^2})\) time.
Mathematical Approach
The key observation is that the hyperball can be counted shell by shell according to the squared distance from the origin.
Step 1: Turn the Hyperball into a Summatory Function
Set \(N=r^2\). For each integer \(n\ge 0\), let \(r_4(n)\) be the number of integer quadruples satisfying
$$x_1^2+x_2^2+x_3^2+x_4^2=n.$$
Every lattice point inside the ball has some squared radius \(n\) with \(0\le n\le N\), so
$$T(r)=\sum_{n=0}^{N} r_4(n).$$
This removes the geometric viewpoint and replaces it with a summation problem.
Step 2: Apply Jacobi's Four-Square Theorem
For \(n>0\), Jacobi's theorem gives
$$r_4(n)=8\sum_{\substack{d\mid n\\4\nmid d}} d,$$
and of course \(r_4(0)=1\). Substituting this into the previous identity yields
$$T(r)=1+8\sum_{n=1}^{N}\sum_{\substack{d\mid n\\4\nmid d}} d.$$
Now interchange the order of summation. A fixed divisor \(d\) contributes once for each multiple of \(d\) up to \(N\), so it appears exactly \(\left\lfloor N/d\right\rfloor\) times. Therefore
$$T(r)=1+8\sum_{\substack{d\le N\\4\nmid d}} d\left\lfloor\frac{N}{d}\right\rfloor.$$
Step 3: Remove the Multiples of 4 Cleanly
Introduce the summatory function
$$S(N)=\sum_{d=1}^{N} d\left\lfloor\frac{N}{d}\right\rfloor.$$
Then the terms with \(4\nmid d\) are obtained by subtracting the terms with \(d=4k\). Those discarded terms are
$$\sum_{k=1}^{\lfloor N/4\rfloor} 4k\left\lfloor\frac{N}{4k}\right\rfloor.$$
Since \(\left\lfloor N/(4k)\right\rfloor=\left\lfloor \left\lfloor N/4\right\rfloor /k\right\rfloor\), this becomes
$$4S\!\left(\left\lfloor\frac{N}{4}\right\rfloor\right).$$
Hence the whole problem collapses to the closed form
$$\boxed{T(r)=1+8\left(S(N)-4S\!\left(\left\lfloor\frac{N}{4}\right\rfloor\right)\right),\qquad N=r^2.}$$
Step 4: Evaluate \(S(N)\) in \(O(\sqrt N)\)
A naive evaluation of \(S(N)\) would still be linear in \(N\). The implementation avoids that by exploiting repeated values of \(\left\lfloor N/d\right\rfloor\).
Let
$$M=\left\lfloor\sqrt N\right\rfloor.$$
For the small divisors \(1\le d\le M\), compute the terms directly:
$$\sum_{d=1}^{M} d\left\lfloor\frac{N}{d}\right\rfloor.$$
For large divisors \(d>M\), group them by the common quotient
$$q=\left\lfloor\frac{N}{d}\right\rfloor.$$
All \(d\) producing the same \(q\) lie in the interval
$$L_q=\left\lfloor\frac{N}{q+1}\right\rfloor+1,\qquad R_q=\left\lfloor\frac{N}{q}\right\rfloor.$$
Only the portion above \(M\) belongs to the second half, so the effective lower bound is \(\max(M+1,L_q)\). The grouped contribution is then
$$q\sum_{d=\max(M+1,L_q)}^{R_q} d,$$
and the inner sum is an arithmetic progression:
$$\sum_{d=l}^{r} d=\frac{(l+r)(r-l+1)}{2}.$$
This split-hyperbola decomposition visits only about \(2\sqrt N\) meaningful ranges.
Step 5: Work Modulo \(10^9+7\)
The target answer is needed modulo
$$P=10^9+7.$$
So every addition and multiplication is reduced modulo \(P\). The division by \(2\) in the arithmetic-series formula is replaced by the modular inverse
$$2^{-1}\equiv 500000004 \pmod{P}.$$
Because the arithmetic-series sum is an exact integer before reduction, performing the reduction at every step preserves the correct final residue.
Worked Example: \(r=2\)
Here \(N=r^2=4\). First compute
$$S(4)=1\cdot 4+2\cdot 2+3\cdot 1+4\cdot 1=15.$$
Also, \(\left\lfloor N/4\right\rfloor=1\), so
$$S(1)=1.$$
Substituting into the closed form gives
$$T(2)=1+8(15-4\cdot 1)=1+8\cdot 11=89.$$
This matches the standard checkpoint for the problem and confirms that the divisor-sum reformulation is correct.
How the Code Works
The C++, Python, and Java implementations all evaluate the same formula
$$T(r)=1+8\left(S(r^2)-4S\!\left(\left\lfloor\frac{r^2}{4}\right\rfloor\right)\right)\pmod{10^9+7}.$$
Each implementation first computes an integer square root to split the summation for \(S(N)\) into a direct part and a grouped part. The direct part handles small divisors one by one. The grouped part iterates over quotient values \(q\), reconstructs the divisor interval that shares that quotient, clips the interval so that it stays strictly above \(\sqrt N\), and then adds the interval using the arithmetic-series formula.
The implementations keep every intermediate result modulo \(10^9+7\). The Python and Java versions partition the two summation phases across available worker tasks, while the C++ version performs the same mathematics in a single-threaded pass with careful wide-integer multiplication for safety. Despite those engineering differences, all three implementations are computing exactly the same divisor decomposition.
Complexity Analysis
For one evaluation of \(S(N)\), the direct loop has length \(\lfloor\sqrt N\rfloor\), and the grouped loop also runs over only \(O(\sqrt N)\) distinct quotient values. Therefore
$$S(N)\text{ can be evaluated in }O(\sqrt N)\text{ time}.$$
The final answer requires two such evaluations, namely at \(N=r^2\) and at \(\left\lfloor N/4\right\rfloor\), so the asymptotic running time remains \(O(\sqrt N)\). The serial method uses \(O(1)\) auxiliary memory, and the parallel variants add only modest task-management overhead on top of the same \(O(\sqrt N)\) total work.
Footnotes and References
- Problem page: https://projecteuler.net/problem=596
- Jacobi's four-square theorem: Wikipedia - Jacobi's four-square theorem
- Dirichlet hyperbola method: Wikipedia - Dirichlet hyperbola method
- Sum of squares function: Wikipedia - Sum of squares function
- Lattice point: Wikipedia - Lattice point
Problem 596 source code
C++
#include <cmath>
#include <cstdint>
#include <iostream>
// Project Euler 596: Number of Lattice Points in a Hyperball
//
// Let N = r^2. Then
// T(r) = sum_{n=0}^{N} r_4(n),
// where r_4(n) is the number of representations of n as a sum of four squares.
// For n>0, Jacobi's four-square theorem gives:
// r_4(n) = 8 * sum_{d|n, 4 \nmid d} d.
// With r_4(0)=1, we obtain
// T(r) = 1 + 8 * sum_{n=1}^{N} sum_{d|n, 4\nmid d} d
// = 1 + 8 * sum_{d<=N, 4\nmid d} d * floor(N/d).
// Let
// S(N) = sum_{d=1}^{N} d * floor(N/d).
// Removing d divisible by 4 yields
// sum_{d<=N,4\nmid d} d*floor(N/d) = S(N) - sum_{k<=N/4} (4k)*floor(N/(4k))
// = S(N) - 4*S(N/4).
// Therefore
// T(r) = 1 + 8 * (S(N) - 4*S(N/4)).
//
// We need T(1e8) mod 1e9+7 with N=1e16. We compute S(N) mod M in O(sqrt N) using a
// split hyperbola method: handle i<=sqrt(N) directly, and for i>sqrt(N) group by q=floor(N/i).
using u64 = std::uint64_t;
using u128 = unsigned __int128;
static constexpr u64 MOD = 1000000007ULL;
static constexpr u64 INV2 = 500000004ULL; // (MOD+1)/2
static inline u64 mod_add(u64 a, u64 b) {
a += b;
if (a >= MOD) a -= MOD;
return a;
}
static inline u64 mod_sub(u64 a, u64 b) {
return (a >= b) ? (a - b) : (a + MOD - b);
}
static inline u64 mod_mul(u64 a, u64 b) {
return (u64)((u128)a * (u128)b % (u128)MOD);
}
static u64 isqrt_u64(u64 x) {
u64 r = (u64)std::sqrt((long double)x);
while ((u128)(r + 1) * (u128)(r + 1) <= (u128)x) ++r;
while ((u128)r * (u128)r > (u128)x) --r;
return r;
}
static inline u64 sum_arith_mod(u64 l, u64 r) {
// sum_{i=l}^r i mod MOD
const u64 cnt = (r - l + 1) % MOD;
const u64 lr = (l % MOD + r % MOD) % MOD;
return mod_mul(mod_mul(lr, cnt), INV2);
}
static u64 S_mod(u64 N) {
if (N == 0) return 0;
const u64 M = isqrt_u64(N);
u64 ans = 0;
// i <= M
for (u64 i = 1; i <= M; ++i) {
const u64 q = N / i;
ans = mod_add(ans, mod_mul(i % MOD, q % MOD));
}
// i > M, grouped by q = floor(N/i)
const u64 qmax = N / (M + 1);
for (u64 q = 1; q <= qmax; ++q) {
u64 l = N / (q + 1) + 1;
const u64 r = N / q;
if (r <= M) continue;
if (l <= M) l = M + 1;
if (l > r) continue;
ans = mod_add(ans, mod_mul(q % MOD, sum_arith_mod(l, r)));
}
return ans;
}
static u64 T_mod(u64 r) {
const u64 N = r * r;
const u64 sN = S_mod(N);
const u64 sN4 = S_mod(N / 4);
const u64 term = mod_sub(sN, mod_mul(4 % MOD, sN4));
const u64 t = mod_add(1, mod_mul(8 % MOD, term));
return t;
}
int main() {
// Validation points from the statement.
if (T_mod(2) != 89ULL) {
std::cerr << "Validation failed: T(2)\n";
return 1;
}
if (T_mod(5) != 3121ULL) {
std::cerr << "Validation failed: T(5)\n";
return 1;
}
if (T_mod(100) != 493490641ULL) {
std::cerr << "Validation failed: T(100)\n";
return 1;
}
{
const u64 expected = 49348022079085897ULL % MOD;
if (T_mod(10000) != expected) {
std::cerr << "Validation failed: T(1e4)\n";
return 1;
}
}
std::cout << T_mod(100000000ULL) << "\n";
return 0;
}
Python
import math
from multiprocessing import Pool, cpu_count
MOD = 1000000007
INV2 = 500000004
def worker_part1(args):
start, end, N = args
ans = 0
for i in range(start, end):
ans += i * (N // i)
return ans % MOD
def worker_part2(args):
start, end, M, N = args
ans = 0
for q in range(start, end):
l = N // (q + 1) + 1
r = N // q
if r <= M: continue
if l <= M: l = M + 1
if l > r: continue
cnt = (r - l + 1) % MOD
lr = (l + r) % MOD
s = (lr * cnt * INV2) % MOD
ans = (ans + (q % MOD) * s) % MOD
return ans % MOD
def S_mod(N):
if N == 0: return 0
M = math.isqrt(N)
threads = max(1, cpu_count())
tasks1 = []
chunk1 = (M + threads - 1) // threads
for t in range(threads):
s = 1 + t * chunk1
e = min(M + 1, 1 + (t + 1) * chunk1)
if s < e:
tasks1.append((s, e, N))
qmax = N // (M + 1)
tasks2 = []
chunk2 = (qmax + threads - 1) // threads
for t in range(threads):
s = 1 + t * chunk2
e = min(qmax + 1, 1 + (t + 1) * chunk2)
if s < e:
tasks2.append((s, e, M, N))
ans = 0
if len(tasks1) + len(tasks2) > 2:
with Pool(threads) as pool:
ans1 = pool.map(worker_part1, tasks1)
ans2 = pool.map(worker_part2, tasks2)
ans = (sum(ans1) + sum(ans2)) % MOD
else:
for t in tasks1:
ans = (ans + worker_part1(t)) % MOD
for t in tasks2:
ans = (ans + worker_part2(t)) % MOD
return ans % MOD
def T_mod(r):
N = r * r
sN = S_mod(N)
sN4 = S_mod(N // 4)
term = (sN - (4 * sN4) % MOD) % MOD
if term < 0: term += MOD
t = (1 + 8 * term) % MOD
return t
def solve():
return str(T_mod(100000000))
if __name__ == '__main__':
print(solve())
Java
import java.util.ArrayList;
import java.util.List;
import java.util.concurrent.ForkJoinPool;
import java.util.concurrent.RecursiveTask;
public class Euler596 {
static final long MOD = 1000000007L;
static final long INV2 = 500000004L;
static long modAdd(long a, long b) {
a += b;
if (a >= MOD)
a -= MOD;
return a;
}
static long modSub(long a, long b) {
return (a >= b) ? (a - b) : (a + MOD - b);
}
static long isqrt(long x) {
long r = (long) Math.sqrt((double) x);
while (r > 0 && r * r > x)
r--;
while ((r + 1) * (r + 1) <= x && (r + 1) * (r + 1) > 0)
r++;
return r;
}
static class SWorkerPart1 extends RecursiveTask<Long> {
long start, end, N;
SWorkerPart1(long s, long e, long n) {
start = s;
end = e;
N = n;
}
protected Long compute() {
long ans = 0;
for (long i = start; i < end; i++) {
ans = (ans + (i % MOD) * ((N / i) % MOD)) % MOD;
}
return ans;
}
}
static class SWorkerPart2 extends RecursiveTask<Long> {
long start, end, M, N;
SWorkerPart2(long s, long e, long m, long n) {
start = s;
end = e;
M = m;
N = n;
}
protected Long compute() {
long ans = 0;
for (long q = start; q < end; q++) {
long l = N / (q + 1) + 1;
long r = N / q;
if (r <= M)
continue;
if (l <= M)
l = M + 1;
if (l > r)
continue;
long cnt = (r - l + 1) % MOD;
long lr = (l % MOD + r % MOD) % MOD;
long sVal = (((lr * cnt) % MOD) * INV2) % MOD;
ans = (ans + (q % MOD) * sVal) % MOD;
}
return ans;
}
}
static long SMod(long N) {
if (N == 0)
return 0;
long M = isqrt(N);
int threads = Math.max(1, Runtime.getRuntime().availableProcessors());
ForkJoinPool pool = new ForkJoinPool(threads);
long ans = 0;
List<SWorkerPart1> tasks1 = new ArrayList<>();
long chunk1 = (M + threads - 1) / threads;
for (int t = 0; t < threads; t++) {
long s = 1 + t * chunk1;
long e = Math.min(M + 1, 1 + (t + 1) * chunk1);
if (s < e)
tasks1.add(new SWorkerPart1(s, e, N));
}
long qmax = N / (M + 1);
List<SWorkerPart2> tasks2 = new ArrayList<>();
long chunk2 = (qmax + threads - 1) / threads;
for (int t = 0; t < threads; t++) {
long s = 1 + t * chunk2;
long e = Math.min(qmax + 1, 1 + (t + 1) * chunk2);
if (s < e)
tasks2.add(new SWorkerPart2(s, e, M, N));
}
if (tasks1.size() + tasks2.size() > 2) {
for (int i = 1; i < tasks1.size(); i++)
tasks1.get(i).fork();
for (int i = 1; i < tasks2.size(); i++)
tasks2.get(i).fork();
if (!tasks1.isEmpty())
ans = modAdd(ans, tasks1.get(0).compute());
if (!tasks2.isEmpty())
ans = modAdd(ans, tasks2.get(0).compute());
for (int i = 1; i < tasks1.size(); i++)
ans = modAdd(ans, tasks1.get(i).join());
for (int i = 1; i < tasks2.size(); i++)
ans = modAdd(ans, tasks2.get(i).join());
} else {
for (SWorkerPart1 task : tasks1)
ans = modAdd(ans, task.compute());
for (SWorkerPart2 task : tasks2)
ans = modAdd(ans, task.compute());
}
return ans;
}
static long TMod(long r) {
long N = r * r;
long sN = SMod(N);
long sN4 = SMod(N / 4);
long term = modSub(sN, (4 * sN4) % MOD);
long t = modAdd(1, (8 * term) % MOD);
return t;
}
public static String solve() {
return Long.toString(TMod(100000000L));
}
public static void main(String[] args) {
System.out.println(solve());
}
}