Problem 989: Fibonacci Sum
View on Project EulerProject Euler Problem 989 Solution
EulerSolve provides an optimized solution for Project Euler Problem 989, Fibonacci Sum, with C++, Python, Java, and a step-by-step mathematical explanation.
Problem Summary The quantity to compute is \[ \sum_{n=1}^{L} F_n G(n) \pmod{10^9+9}, \qquad L=10^{14}, \] where \(F_n\) is the Fibonacci sequence and \(G(n)\) counts the residue classes \(x \pmod n\) satisfying \[ x^2 \equiv x+1 \pmod n. \] The published checkpoint is \[ \sum_{n=1}^{10^3} F_n G(n)\equiv 190950976 \pmod{10^9+9}. \] A direct sweep up to \(10^{14}\) is hopeless, so the solution rewrites \(G(n)\) in arithmetic terms, turns the Fibonacci weight into exponential weights via Binet's formula, and then evaluates the resulting sums with Möbius inversion and a sliding-window count for a binary quadratic form. Mathematical Approach Prime-power behavior of \(G(n)\) The congruence is the same as \[ x^2-x-1 \equiv 0 \pmod n, \] whose discriminant is \(5\). That immediately explains the local structure. For \(2^e\) there are no solutions, so \(G(2^e)=0\). For \(5\) there is exactly one solution, namely \(x\equiv 3 \pmod 5\), but it does not lift to \(25\), so \[ G(5)=1, \qquad G(5^e)=0 \ \text{ for } e\ge 2. \] For an odd prime \(p\neq 5\), the polynomial has roots modulo \(p\) exactly when \(5\) is a quadratic residue modulo \(p\). By quadratic reciprocity this happens precisely for \[ p\equiv 1,4 \pmod 5. \] In that case there are two distinct roots modulo \(p\), the derivative \(2x-1\) is nonzero at each root, and Hensel lifting gives two roots modulo every \(p^e\)....
Detailed mathematical approach
Problem Summary
The quantity to compute is
\[ \sum_{n=1}^{L} F_n G(n) \pmod{10^9+9}, \qquad L=10^{14}, \]
where \(F_n\) is the Fibonacci sequence and \(G(n)\) counts the residue classes \(x \pmod n\) satisfying
\[ x^2 \equiv x+1 \pmod n. \]
The published checkpoint is
\[ \sum_{n=1}^{10^3} F_n G(n)\equiv 190950976 \pmod{10^9+9}. \]
A direct sweep up to \(10^{14}\) is hopeless, so the solution rewrites \(G(n)\) in arithmetic terms, turns the Fibonacci weight into exponential weights via Binet's formula, and then evaluates the resulting sums with Möbius inversion and a sliding-window count for a binary quadratic form.
Mathematical Approach
Prime-power behavior of \(G(n)\)
The congruence is the same as
\[ x^2-x-1 \equiv 0 \pmod n, \]
whose discriminant is \(5\). That immediately explains the local structure.
For \(2^e\) there are no solutions, so \(G(2^e)=0\). For \(5\) there is exactly one solution, namely \(x\equiv 3 \pmod 5\), but it does not lift to \(25\), so
\[ G(5)=1, \qquad G(5^e)=0 \ \text{ for } e\ge 2. \]
For an odd prime \(p\neq 5\), the polynomial has roots modulo \(p\) exactly when \(5\) is a quadratic residue modulo \(p\). By quadratic reciprocity this happens precisely for
\[ p\equiv 1,4 \pmod 5. \]
In that case there are two distinct roots modulo \(p\), the derivative \(2x-1\) is nonzero at each root, and Hensel lifting gives two roots modulo every \(p^e\). If \(p\equiv 2,3 \pmod 5\), there are no roots at any power of \(p\). Hence
\[ G(p^e)= \begin{cases} 2, & p\equiv 1,4 \pmod 5,\\ 0, & p\equiv 2,3 \pmod 5, \end{cases} \qquad (p\neq 5,\ p \text{ odd}). \]
By the Chinese remainder theorem, \(G\) is multiplicative. Therefore, if
\[ n=5^\varepsilon \prod_{i=1}^{k} p_i^{a_i}\prod_{j=1}^{m} q_j^{b_j}, \]
where every \(p_i\equiv 1,4 \pmod 5\) and every \(q_j\) is either \(2\) or congruent to \(2\) or \(3\) modulo \(5\), then \(G(n)=0\) unless \(m=0\) and \(\varepsilon\in\{0,1\}\). In the surviving case,
\[ G(n)=2^k. \]
From roots to the norm form \(Q(a,b)=a^2-ab-b^2\)
The next step moves to the quadratic ring generated by a root of \(t^2-t-1\). Let \(\varphi\) and \(\psi=1-\varphi\) be the two roots of \(t^2-t-1=0\), so \(\varphi^2=\varphi+1\). In the ring \(\mathbb Z[\varphi]\), the norm of \(a-b\varphi\) is
\[ N(a-b\varphi)=(a-b\varphi)(a-b\psi)=a^2-ab-b^2. \]
This binary quadratic form is exactly the arithmetic object used by the solver:
\[ Q(a,b)=a^2-ab-b^2. \]
The prime factors that contribute to \(G(n)\) are precisely the primes that split in \(\mathbb Z[\varphi]\). Choosing one root of \(x^2-x-1\) at each split prime power is equivalent to choosing one prime factor above each rational prime. Multiplying those local choices produces an algebraic integer of norm \(n\).
Two issues remain: units and non-primitive representations. Multiplying by a unit does not change the norm, so one imposes a reduction region to select a single representative. The implementations use the classical reduced region
\[ a\ge 2b>0. \]
Also, if \(\gcd(a,b)>1\), then the representation is not primitive and corresponds to extra square factors. After enforcing both conditions, one obtains the key identity
\[ G(n)=\#\{(a,b):\ a\ge 2b>0,\ \gcd(a,b)=1,\ Q(a,b)=n\}. \]
Fibonacci weights become exponential weights
The modulus \(10^9+9\) is prime, \(5\) has a square root modulo this prime, and the two modular roots \(\varphi,\psi\) of \(t^2-t-1\) satisfy
\[ \varphi\psi=-1. \]
So Binet's formula is valid in the finite field:
\[ F_n=\frac{\varphi^n-\psi^n}{\sqrt 5}. \]
That converts the original sum into two weighted counts of quadratic-form values. Define
\[ P_w(L)= \sum_{\substack{a\ge 2b>0\\ \gcd(a,b)=1\\ Q(a,b)\le L}} w^{Q(a,b)}. \]
Then the desired answer is
\[ \sum_{n\le L} F_n G(n) = \frac{P_\varphi(L)-P_\psi(L)}{\sqrt 5} \pmod{10^9+9}. \]
Removing the coprimality condition with Möbius inversion
The form \(Q\) is homogeneous of degree two:
\[ Q(ga,gb)=g^2Q(a,b). \]
Let
\[ A_w(L)= \sum_{\substack{a\ge 2b>0\\ Q(a,b)\le L}} w^{Q(a,b)} \]
be the same sum without \(\gcd(a,b)=1\). Möbius inversion then gives
\[ P_w(L)= \sum_{g\le \sqrt L} \mu(g)\, A_{w^{g^2}}\!\left(\left\lfloor \frac{L}{g^2}\right\rfloor\right). \]
So the primitive count is recovered by inclusion-exclusion over the common divisor \(g\).
Diagonalizing the form and splitting by parity
The decisive algebraic identity is
\[ 4Q(a,b)=(2a-b)^2-5b^2. \]
Set
\[ u=2a-b,\qquad v=b. \]
Because \(a\ge 2b>0\), we have \(u\ge 3v>0\). Also \(u\equiv v \pmod 2\), so only two parity branches occur.
If \(u=2m\) and \(v=2t\), then
\[ Q(a,b)=m^2-5t^2, \qquad m\ge 3t, \qquad t\ge 1. \]
If \(u=2m+1\) and \(v=2t+1\), then
\[ Q(a,b)=m(m+1)-5t(t+1)-1, \qquad m\ge 3t+1, \qquad t\ge 0. \]
Therefore \(A_w(L)\) is the sum of two one-dimensional window sums:
\[ \sum_{t\ge 1} \sum_{m=3t}^{\lfloor \sqrt{L+5t^2}\rfloor} w^{m^2-5t^2}, \]
\[ \sum_{t\ge 0} \sum_{m=3t+1}^{\left\lfloor(\sqrt{4L+20t^2+20t+5}-1)/2\right\rfloor} w^{m(m+1)-5t(t+1)-1}. \]
Worked example: \(n=11\)
Since \(11\equiv 1 \pmod 5\), there are two solutions of \(x^2\equiv x+1 \pmod{11}\), namely \(x\equiv 4\) and \(x\equiv 8\). So \(G(11)=2\).
The reduced primitive representations of \(11\) by \(Q(a,b)\) are
\[ Q(4,1)=16-4-1=11, \qquad Q(5,2)=25-10-4=11. \]
The first pair lands in the odd branch: \((u,v)=(2\cdot 4-1,1)=(7,1)\), so \(u=2m+1\), \(v=2t+1\) with \((m,t)=(3,0)\), and indeed
\[ m(m+1)-5t(t+1)-1=3\cdot 4-0-1=11. \]
The second pair lands in the even branch: \((u,v)=(8,2)\), so \((m,t)=(4,1)\), and
\[ m^2-5t^2=16-5=11. \]
This small case shows why the pair count agrees with \(G(n)\) and why both parity branches are needed in the final summation.
How the Code Works
Arithmetic precomputation and verification
The C++, Python, and Java implementations begin by working modulo \(10^9+9\). They verify that the chosen modular constants really satisfy \(\sqrt 5^2=5\), \(\varphi^2=\varphi+1\), \(\psi^2=\psi+1\), and \(\varphi\psi=-1\). They also compare three views of the problem on small inputs: direct brute force for the congruence, the prime-factor formula for \(G(n)\), and the reduced-pair count for the quadratic form. Finally, they check the published sample sum for \(L=10^3\).
Evaluating one non-primitive weighted sum
For a fixed weight \(w\), the implementation computes \(A_w(L)\) by scanning the two parity branches separately. It does not recompute \(w^{m^2}\) or \(w^{m(m+1)}\) from scratch. Instead it updates these powers multiplicatively, using the identities
\[ w^{(m+1)^2}=w^{m^2}w^{2m+1}, \qquad w^{(m+1)(m+2)}=w^{m(m+1)}w^{2m+2}. \]
As \(t\) grows, the allowed interval of \(m\) moves to the right. The implementation keeps a running window sum: new terms are appended when the upper bound increases, and expired terms are removed when the lower bound advances. Each admissible power enters once and leaves once.
Möbius sweep and final combination
After that, the implementation iterates over \(g\le \sqrt L\) with \(\mu(g)\neq 0\), evaluates the previous routine at the scaled limit \(\lfloor L/g^2\rfloor\), and combines the results with the Möbius sign. It precomputes the sequences \(\varphi^{g^2}\) and \(\varphi^{-g^2}\), which are enough for both branches because
\[ \psi=-\varphi^{-1}. \]
So \(\psi^{g^2}\) differs from \(\varphi^{-g^2}\) only by the parity-dependent sign of \(g^2\). The final step subtracts the \(\psi\)-sum from the \(\varphi\)-sum and multiplies by \(1/\sqrt 5\). The C++ and Java implementations can split the Möbius range across threads, while the Python implementation supports the same mathematics either serially or across processes.
Complexity Analysis
The precomputed Möbius array and the tables of \(\varphi^{g^2}\) and \(\varphi^{-g^2}\) have size \(O(\sqrt L)\), so memory usage is \(O(\sqrt L)\).
For one fixed \(g\), the scaled limit is \(L/g^2\), and the two sliding-window branches together cost \(O(\sqrt{L/g^2})=O(\sqrt L/g)\). Summing over all \(g\le \sqrt L\) yields
\[ O\!\left(\sum_{g\le \sqrt L}\frac{\sqrt L}{g}\right) =O(\sqrt L\log L). \]
That is the reason the method can handle \(L=10^{14}\): it never iterates over all \(n\le L\), and it never enumerates all primitive pairs individually.
Footnotes and References
- Project Euler problem page: Project Euler 989
- Fibonacci numbers and Binet's formula: Wikipedia - Fibonacci number
- Hensel lifting: Wikipedia - Hensel's lemma
- Quadratic reciprocity and residue classes mod \(5\): Wikipedia - Quadratic reciprocity
- Binary quadratic forms: Wikipedia - Binary quadratic form
- Möbius inversion: Wikipedia - Möbius inversion formula
Problem 989 source code
C++
#include <algorithm>
#include <cassert>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <cstdlib>
#include <iomanip>
#include <iostream>
#include <limits>
#include <pthread.h>
#include <string>
#include <thread>
#include <vector>
namespace {
using i128 = __int128_t;
using u32 = std::uint32_t;
using u64 = std::uint64_t;
using u128 = unsigned __int128;
constexpr u64 MOD = 1'000'000'009ULL;
constexpr u64 TARGET_LIMIT = 100'000'000'000'000ULL;
constexpr u64 SAMPLE_LIMIT = 1'000ULL;
constexpr u64 SAMPLE_SUM = 190'950'976ULL;
constexpr u64 SQRT5_MOD = 383'008'016ULL;
constexpr u64 PHI_MOD = 691'504'013ULL;
constexpr u64 PSI_MOD = 308'495'997ULL;
struct Options {
bool allow_multithreading = true;
unsigned requested_threads = 0;
};
void usage() {
std::cerr
<< "Usage:\n"
<< " ./Euler989 validate [check_max] [--single-thread] [--threads=N]\n"
<< " ./Euler989 sum <limit> [--single-thread] [--threads=N]\n"
<< " ./Euler989 answer [--single-thread] [--threads=N]\n";
}
u64 mod_mul(const u64 a, const u64 b) {
return static_cast<u64>((static_cast<u128>(a) * static_cast<u128>(b)) % static_cast<u128>(MOD));
}
u64 mod_add(const u64 a, const u64 b) {
const u64 s = a + b;
return s >= MOD ? s - MOD : s;
}
u64 mod_sub(const u64 a, const u64 b) {
return a >= b ? a - b : a + MOD - b;
}
u64 mod_neg(const u64 a) {
return a == 0 ? 0 : MOD - a;
}
u64 mod_pow(u64 base, u64 exp) {
u64 result = 1;
while (exp > 0) {
if ((exp & 1ULL) != 0ULL) {
result = mod_mul(result, base);
}
base = mod_mul(base, base);
exp >>= 1U;
}
return result;
}
u64 isqrt(const u64 n) {
u64 x = static_cast<u64>(std::sqrt(static_cast<long double>(n)));
while ((x + 1) <= n / (x + 1)) {
++x;
}
while (x > n / x) {
--x;
}
return x;
}
std::vector<u32> sieve_primes(const int limit) {
if (limit < 2) {
return {};
}
std::vector<bool> is_prime(static_cast<std::size_t>(limit + 1), true);
std::vector<u32> primes;
is_prime[0] = false;
is_prime[1] = false;
for (int i = 2; i <= limit; ++i) {
if (!is_prime[static_cast<std::size_t>(i)]) {
continue;
}
primes.push_back(static_cast<u32>(i));
if (i > limit / i) {
continue;
}
for (int j = i * i; j <= limit; j += i) {
is_prime[static_cast<std::size_t>(j)] = false;
}
}
return primes;
}
std::vector<std::int8_t> mobius_sieve(const int limit) {
std::vector<std::int8_t> mu(static_cast<std::size_t>(limit + 1), 0);
std::vector<int> primes;
std::vector<int> least(static_cast<std::size_t>(limit + 1), 0);
mu[1] = 1;
for (int i = 2; i <= limit; ++i) {
if (least[static_cast<std::size_t>(i)] == 0) {
least[static_cast<std::size_t>(i)] = i;
primes.push_back(i);
mu[static_cast<std::size_t>(i)] = -1;
}
for (const int p : primes) {
const int ip = i * p;
if (ip > limit || p > least[static_cast<std::size_t>(i)]) {
break;
}
least[static_cast<std::size_t>(ip)] = p;
if (p == least[static_cast<std::size_t>(i)]) {
mu[static_cast<std::size_t>(ip)] = 0;
break;
}
mu[static_cast<std::size_t>(ip)] = -mu[static_cast<std::size_t>(i)];
}
}
return mu;
}
u32 g_bruteforce(const u64 n) {
u32 count = 0;
for (u64 x = 0; x < n; ++x) {
if ((x * x + n - x - 1) % n == 0) {
++count;
}
}
return count;
}
u32 g_from_factorization(u64 n, const std::vector<u32>& primes) {
if (n == 1) {
return 1;
}
u32 result = 1;
for (const u32 p32 : primes) {
const u64 p = static_cast<u64>(p32);
if (p > n / p) {
break;
}
if (n % p != 0) {
continue;
}
u32 exponent = 0;
while (n % p == 0) {
n /= p;
++exponent;
}
if (p == 2) {
return 0;
}
if (p == 5) {
if (exponent >= 2) {
return 0;
}
continue;
}
switch (p % 5) {
case 1:
case 4:
result = static_cast<u32>(result * 2U);
break;
case 2:
case 3:
return 0;
default:
std::abort();
}
}
if (n > 1) {
if (n == 2) {
return 0;
}
if (n == 5) {
return result;
}
switch (n % 5) {
case 1:
case 4:
result = static_cast<u32>(result * 2U);
break;
case 2:
case 3:
return 0;
default:
std::abort();
}
}
return result;
}
u64 q_form(const u64 a, const u64 b) {
return a * a - a * b - b * b;
}
u64 gcd(u64 a, u64 b) {
while (b != 0) {
const u64 t = a % b;
a = b;
b = t;
}
return a;
}
u32 reduced_pair_count(const u64 n) {
u32 count = 0;
const u64 max_b = isqrt(n);
for (u64 b = 1; b <= max_b; ++b) {
for (u64 a = 2 * b;; ++a) {
const u64 q = q_form(a, b);
if (q > n) {
break;
}
if (q == n && gcd(a, b) == 1) {
++count;
}
}
}
return count;
}
u64 direct_pair_sum(const u64 limit, const u64 base) {
u64 acc = 0;
const u64 max_b = isqrt(limit);
for (u64 b = 1; b <= max_b; ++b) {
for (u64 a = 2 * b;; ++a) {
const u64 q = q_form(a, b);
if (q > limit) {
break;
}
if (gcd(a, b) == 1) {
acc = mod_add(acc, mod_pow(base, q));
}
}
}
return acc;
}
struct PowerSeq {
u64 v = 0;
u64 term = 1;
u64 delta = 1;
u64 delta_step = 1;
static PowerSeq square(const u64 base) {
const u64 base_sq = mod_mul(base, base);
return {0, 1, base, base_sq};
}
static PowerSeq triangular(const u64 base) {
const u64 base_sq = mod_mul(base, base);
return {0, 1, base_sq, base_sq};
}
void step() {
term = mod_mul(term, delta);
delta = mod_mul(delta, delta_step);
++v;
}
void extend_through(const u64 target, u64& window) {
while (v <= target) {
window = mod_add(window, term);
step();
}
}
void trim_before(const u64 target, u64& window) {
while (v < target) {
window = mod_sub(window, term);
step();
}
}
};
u64 nonprimitive_sum(const u64 limit, const u64 w, const u64 winv) {
if (limit == 0) {
return 0;
}
const u64 sqrt_limit = isqrt(limit);
const u64 winv5 = mod_pow(winv, 5);
const u64 winv10 = mod_mul(winv5, winv5);
u64 ans = 0;
const u64 even_t_max = sqrt_limit / 2;
if (even_t_max > 0) {
PowerSeq add_seq = PowerSeq::square(w);
PowerSeq trim_seq = PowerSeq::square(w);
u64 window = 0;
u64 factor = 1;
u64 ratio = winv5;
for (u64 t = 1; t <= even_t_max; ++t) {
factor = mod_mul(factor, ratio);
ratio = mod_mul(ratio, winv10);
const u64 vmax = isqrt(limit + 5 * t * t);
add_seq.extend_through(vmax, window);
trim_seq.trim_before(3 * t, window);
ans = mod_add(ans, mod_mul(factor, window));
}
}
const u64 odd_t_max = (sqrt_limit - 1) / 2;
PowerSeq add_seq = PowerSeq::triangular(w);
PowerSeq trim_seq = PowerSeq::triangular(w);
u64 window = 0;
u64 factor = winv;
u64 ratio = winv10;
for (u64 t = 0; t <= odd_t_max; ++t) {
const u64 disc = 4 * limit + 20 * t * t + 20 * t + 5;
const u64 vmax = (isqrt(disc) - 1) / 2;
add_seq.extend_through(vmax, window);
trim_seq.trim_before(3 * t + 1, window);
ans = mod_add(ans, mod_mul(factor, window));
factor = mod_mul(factor, ratio);
ratio = mod_mul(ratio, winv10);
}
return ans;
}
std::vector<u32> square_powers(const u64 base, const int max_g) {
std::vector<u32> out(static_cast<std::size_t>(max_g + 1), 0);
out[0] = 1;
if (max_g == 0) {
return out;
}
const u64 base_sq = mod_mul(base, base);
u64 value = 1;
u64 ratio = base;
for (int g = 1; g <= max_g; ++g) {
value = mod_mul(value, ratio);
out[static_cast<std::size_t>(g)] = static_cast<u32>(value);
ratio = mod_mul(ratio, base_sq);
}
return out;
}
u64 normalize_signed_mod(i128 value) {
value %= static_cast<i128>(MOD);
if (value < 0) {
value += static_cast<i128>(MOD);
}
return static_cast<u64>(value);
}
unsigned choose_thread_count(const bool allow_multithreading,
const unsigned requested_threads,
const int workload) {
if (!allow_multithreading || workload <= 1) {
return 1U;
}
unsigned threads = requested_threads;
if (threads == 0U) {
threads = std::thread::hardware_concurrency();
if (threads == 0U) {
threads = 1U;
}
}
if (threads > static_cast<unsigned>(workload)) {
threads = static_cast<unsigned>(workload);
}
return std::max(1U, threads);
}
struct PrimitiveWorkerTask {
u64 limit = 0;
int start_g = 1;
int end_g = 1;
const std::vector<std::int8_t>* mu = nullptr;
const std::vector<u32>* phi_sq = nullptr;
const std::vector<u32>* phi_inv_sq = nullptr;
i128 phi_total = 0;
i128 psi_total = 0;
};
void* primitive_worker_entry(void* arg) {
auto& task = *static_cast<PrimitiveWorkerTask*>(arg);
for (int g = task.start_g; g < task.end_g; ++g) {
const int mu_g = static_cast<int>((*task.mu)[static_cast<std::size_t>(g)]);
if (mu_g == 0) {
continue;
}
const u64 gg = static_cast<u64>(g) * static_cast<u64>(g);
const u64 scaled_limit = task.limit / gg;
const u64 phi_w = static_cast<u64>((*task.phi_sq)[static_cast<std::size_t>(g)]);
const u64 phi_winv = static_cast<u64>((*task.phi_inv_sq)[static_cast<std::size_t>(g)]);
const u64 psi_w = (g & 1) == 0 ? phi_winv : mod_neg(phi_winv);
const u64 psi_winv = (g & 1) == 0 ? phi_w : mod_neg(phi_w);
const i128 phi_value = static_cast<i128>(nonprimitive_sum(scaled_limit, phi_w, phi_winv));
const i128 psi_value = static_cast<i128>(nonprimitive_sum(scaled_limit, psi_w, psi_winv));
if (mu_g > 0) {
task.phi_total += phi_value;
task.psi_total += psi_value;
} else {
task.phi_total -= phi_value;
task.psi_total -= psi_value;
}
}
return nullptr;
}
std::pair<u64, u64> primitive_sums(const u64 limit,
const std::vector<std::int8_t>& mu,
const std::vector<u32>& phi_sq,
const std::vector<u32>& phi_inv_sq,
const Options& options) {
const int max_g = static_cast<int>(isqrt(limit));
const unsigned thread_count = choose_thread_count(options.allow_multithreading,
options.requested_threads,
max_g);
std::vector<pthread_t> threads(static_cast<std::size_t>(thread_count));
std::vector<PrimitiveWorkerTask> tasks(static_cast<std::size_t>(thread_count));
for (unsigned t = 0; t < thread_count; ++t) {
PrimitiveWorkerTask& task = tasks[static_cast<std::size_t>(t)];
task.limit = limit;
task.start_g = 1 + static_cast<int>((static_cast<u64>(max_g) * t) / thread_count);
task.end_g = 1 + static_cast<int>((static_cast<u64>(max_g) * (t + 1U)) / thread_count);
task.mu = μ
task.phi_sq = &phi_sq;
task.phi_inv_sq = &phi_inv_sq;
task.phi_total = 0;
task.psi_total = 0;
}
for (unsigned t = 0; t < thread_count; ++t) {
const int rc = pthread_create(&threads[static_cast<std::size_t>(t)],
nullptr,
primitive_worker_entry,
&tasks[static_cast<std::size_t>(t)]);
assert(rc == 0);
}
for (unsigned t = 0; t < thread_count; ++t) {
const int rc = pthread_join(threads[static_cast<std::size_t>(t)], nullptr);
assert(rc == 0);
}
i128 phi_total = 0;
i128 psi_total = 0;
for (const PrimitiveWorkerTask& task : tasks) {
phi_total += task.phi_total;
psi_total += task.psi_total;
}
return {normalize_signed_mod(phi_total), normalize_signed_mod(psi_total)};
}
u64 solve(const u64 limit, const Options& options) {
const int max_g = static_cast<int>(isqrt(limit));
const std::vector<std::int8_t> mu = mobius_sieve(max_g);
const u64 phi_inv = mod_pow(PHI_MOD, MOD - 2);
const std::vector<u32> phi_sq = square_powers(PHI_MOD, max_g);
const std::vector<u32> phi_inv_sq = square_powers(phi_inv, max_g);
const auto [phi_sum, psi_sum] = primitive_sums(limit, mu, phi_sq, phi_inv_sq, options);
const u64 inv_sqrt5 = mod_pow(SQRT5_MOD, MOD - 2);
return mod_mul(mod_sub(phi_sum, psi_sum), inv_sqrt5);
}
u64 checksum_via_factorization(const u64 limit, const std::vector<u32>& primes) {
u64 acc = 0;
u64 f_prev = 0;
u64 f_cur = 1;
for (u64 n = 1; n <= limit; ++n) {
const u64 g = static_cast<u64>(g_from_factorization(n, primes));
acc = mod_add(acc, mod_mul(f_cur, g));
const u64 f_next = mod_add(f_prev, f_cur);
f_prev = f_cur;
f_cur = f_next;
}
return acc;
}
void validate(const u64 check_max, const Options& options) {
assert(mod_mul(SQRT5_MOD, SQRT5_MOD) == 5);
assert(mod_sub(mod_mul(PHI_MOD, PHI_MOD), PHI_MOD) == 1);
assert(mod_sub(mod_mul(PSI_MOD, PSI_MOD), PSI_MOD) == 1);
assert(mod_mul(PHI_MOD, PSI_MOD) == MOD - 1);
std::cout << "Checkpoint 1 passed: modular sqrt(5), phi, and psi are consistent.\n";
const int prime_limit = static_cast<int>(isqrt(check_max)) + 10;
const std::vector<u32> primes = sieve_primes(prime_limit);
for (u64 n = 1; n <= check_max; ++n) {
const u32 brute = g_bruteforce(n);
const u32 factorized = g_from_factorization(n, primes);
assert(brute == factorized);
}
std::cout << "Checkpoint 2 passed: G(n) factorization matches brute force.\n";
for (u64 n = 1; n <= check_max; ++n) {
const u32 reduced_pairs = reduced_pair_count(n);
const u32 factorized = g_from_factorization(n, primes);
assert(reduced_pairs == factorized);
}
std::cout << "Checkpoint 3 passed: reduced quadratic-form pairs match G(n).\n";
const u64 direct_limit = std::min<u64>(check_max, 250);
const u64 phi_pair_sum = direct_pair_sum(direct_limit, PHI_MOD);
const u64 psi_pair_sum = direct_pair_sum(direct_limit, PSI_MOD);
const int max_g = static_cast<int>(isqrt(direct_limit));
const std::vector<std::int8_t> mu = mobius_sieve(max_g);
const u64 phi_inv = mod_pow(PHI_MOD, MOD - 2);
const std::vector<u32> phi_sq = square_powers(PHI_MOD, max_g);
const std::vector<u32> phi_inv_sq = square_powers(phi_inv, max_g);
const auto [phi_fast, psi_fast] = primitive_sums(direct_limit, mu, phi_sq, phi_inv_sq, options);
assert(phi_fast == phi_pair_sum);
assert(psi_fast == psi_pair_sum);
std::cout << "Checkpoint 4 passed: Möbius/parity solver matches direct pair enumeration.\n";
const int sample_prime_limit = static_cast<int>(isqrt(SAMPLE_LIMIT)) + 10;
const std::vector<u32> sample_primes = sieve_primes(sample_prime_limit);
const u64 sample_fast = solve(SAMPLE_LIMIT, options);
const u64 sample_factorized = checksum_via_factorization(SAMPLE_LIMIT, sample_primes);
assert(sample_fast == SAMPLE_SUM);
assert(sample_fast == sample_factorized);
std::cout << "Checkpoint 5 passed: sample checksum equals " << sample_fast << ".\n";
}
bool parse_u64(const std::string& text, u64& value) {
if (text.empty()) {
return false;
}
u64 parsed = 0;
for (const char c : text) {
if (c < '0' || c > '9') {
return false;
}
parsed = parsed * 10 + static_cast<u64>(c - '0');
}
value = parsed;
return true;
}
bool parse_unsigned_after_prefix(const std::string& arg, const char* prefix, unsigned& value) {
const std::string p(prefix);
if (arg.rfind(p, 0) != 0) {
return false;
}
const std::string tail = arg.substr(p.size());
if (tail.empty()) {
return false;
}
u64 parsed = 0;
if (!parse_u64(tail, parsed) || parsed > static_cast<u64>(std::numeric_limits<unsigned>::max())) {
return false;
}
value = static_cast<unsigned>(parsed);
return true;
}
bool parse_command_options(int argc,
char** argv,
int start_index,
Options& options,
std::vector<std::string>& positional) {
for (int i = start_index; i < argc; ++i) {
const std::string arg(argv[i]);
if (arg == "--single-thread") {
options.allow_multithreading = false;
continue;
}
unsigned threads = 0;
if (parse_unsigned_after_prefix(arg, "--threads=", threads)) {
options.requested_threads = threads;
continue;
}
positional.push_back(arg);
}
return true;
}
double elapsed_seconds(const std::chrono::steady_clock::time_point started) {
const auto elapsed = std::chrono::steady_clock::now() - started;
return std::chrono::duration<double>(elapsed).count();
}
} // namespace
int main(int argc, char** argv) {
if (argc < 2) {
usage();
return 0;
}
const std::string command(argv[1]);
Options options;
std::vector<std::string> positional;
if (!parse_command_options(argc, argv, 2, options, positional)) {
usage();
return 1;
}
if (command == "validate") {
if (positional.size() > 1) {
usage();
return 1;
}
u64 check_max = 200;
if (!positional.empty() && !parse_u64(positional[0], check_max)) {
usage();
return 1;
}
const auto started = std::chrono::steady_clock::now();
validate(check_max, options);
std::cout << std::fixed << std::setprecision(3)
<< "Validation completed in " << elapsed_seconds(started) << "s.\n";
return 0;
}
if (command == "sum") {
if (positional.size() != 1) {
usage();
return 1;
}
u64 limit = 0;
if (!parse_u64(positional[0], limit)) {
usage();
return 1;
}
const auto started = std::chrono::steady_clock::now();
const u64 answer = solve(limit, options);
std::cout << answer << '\n';
std::cerr << std::fixed << std::setprecision(3)
<< "Computed in " << elapsed_seconds(started) << "s.\n";
return 0;
}
if (command == "answer") {
if (!positional.empty()) {
usage();
return 1;
}
const auto started = std::chrono::steady_clock::now();
const u64 answer = solve(TARGET_LIMIT, options);
std::cout << answer << '\n';
std::cerr << std::fixed << std::setprecision(3)
<< "Computed in " << elapsed_seconds(started) << "s.\n";
return 0;
}
usage();
return 0;
}
Python
from __future__ import annotations
import math
import multiprocessing as mp
import sys
from array import array
MOD = 1_000_000_009
TARGET_LIMIT = 100_000_000_000_000
SAMPLE_LIMIT = 1_000
SAMPLE_SUM = 190_950_976
SQRT5_MOD = 383_008_016
PHI_MOD = 691_504_013
PSI_MOD = 308_495_997
class Options:
__slots__ = ("allow_multiprocessing", "requested_processes")
def __init__(self):
self.allow_multiprocessing = True
self.requested_processes = 0
_WORK_LIMIT = 0
_WORK_MU = None
_WORK_PHI_SQ = None
_WORK_PHI_INV_SQ = None
def mod_mul(a, b):
return (a * b) % MOD
def mod_add(a, b):
s = a + b
return s - MOD if s >= MOD else s
def mod_sub(a, b):
return a - b if a >= b else a + MOD - b
def mod_neg(a):
return 0 if a == 0 else MOD - a
def mod_pow(base, exp):
return pow(base, exp, MOD)
def normalize_signed_mod(value):
return value % MOD
def sieve_primes(limit):
if limit < 2:
return []
is_prime = bytearray(b"\x01") * (limit + 1)
is_prime[0] = 0
is_prime[1] = 0
primes = []
for i in range(2, limit + 1):
if not is_prime[i]:
continue
primes.append(i)
if i > limit // i:
continue
step = i
start = i * i
is_prime[start : limit + 1 : step] = b"\x00" * (((limit - start) // step) + 1)
return primes
def mobius_sieve(limit):
mu = array("b", [0]) * (limit + 1)
least = array("I", [0]) * (limit + 1)
primes = []
mu[1] = 1
for i in range(2, limit + 1):
if least[i] == 0:
least[i] = i
primes.append(i)
mu[i] = -1
li = least[i]
mui = mu[i]
for p in primes:
ip = i * p
if ip > limit or p > li:
break
least[ip] = p
if p == li:
mu[ip] = 0
break
mu[ip] = -mui
return mu
def g_bruteforce(n):
count = 0
for x in range(n):
if (x * x + n - x - 1) % n == 0:
count += 1
return count
def g_from_factorization(n, primes):
if n == 1:
return 1
result = 1
for p in primes:
if p > n // p:
break
if n % p != 0:
continue
exponent = 0
while n % p == 0:
n //= p
exponent += 1
if p == 2:
return 0
if p == 5:
if exponent >= 2:
return 0
continue
r = p % 5
if r == 1 or r == 4:
result *= 2
elif r == 2 or r == 3:
return 0
else:
raise RuntimeError("unreachable")
if n > 1:
if n == 2:
return 0
if n == 5:
return result
r = n % 5
if r == 1 or r == 4:
result *= 2
elif r == 2 or r == 3:
return 0
else:
raise RuntimeError("unreachable")
return result
def q_form(a, b):
return a * a - a * b - b * b
def gcd(a, b):
while b != 0:
a, b = b, a % b
return a
def reduced_pair_count(n):
count = 0
max_b = math.isqrt(n)
for b in range(1, max_b + 1):
a = 2 * b
while True:
q = q_form(a, b)
if q > n:
break
if q == n and gcd(a, b) == 1:
count += 1
a += 1
return count
def direct_pair_sum(limit, base):
acc = 0
max_b = math.isqrt(limit)
for b in range(1, max_b + 1):
a = 2 * b
while True:
q = q_form(a, b)
if q > limit:
break
if gcd(a, b) == 1:
acc = mod_add(acc, mod_pow(base, q))
a += 1
return acc
class PowerSeq:
__slots__ = ("v", "term", "delta", "delta_step")
def __init__(self, v, term, delta, delta_step):
self.v = v
self.term = term
self.delta = delta
self.delta_step = delta_step
@staticmethod
def square(base):
base_sq = mod_mul(base, base)
return PowerSeq(0, 1, base, base_sq)
@staticmethod
def triangular(base):
base_sq = mod_mul(base, base)
return PowerSeq(0, 1, base_sq, base_sq)
def step(self):
self.term = mod_mul(self.term, self.delta)
self.delta = mod_mul(self.delta, self.delta_step)
self.v += 1
def extend_through(self, target, window):
while self.v <= target:
window += self.term
if window >= MOD:
window -= MOD
self.step()
return window
def trim_before(self, target, window):
while self.v < target:
window -= self.term
if window < 0:
window += MOD
self.step()
return window
def nonprimitive_sum(limit, w, winv):
if limit == 0:
return 0
sqrt_limit = math.isqrt(limit)
winv5 = mod_pow(winv, 5)
winv10 = mod_mul(winv5, winv5)
ans = 0
even_t_max = sqrt_limit // 2
if even_t_max > 0:
add_seq = PowerSeq.square(w)
trim_seq = PowerSeq.square(w)
window = 0
factor = 1
ratio = winv5
for t in range(1, even_t_max + 1):
factor = mod_mul(factor, ratio)
ratio = mod_mul(ratio, winv10)
vmax = math.isqrt(limit + 5 * t * t)
window = add_seq.extend_through(vmax, window)
window = trim_seq.trim_before(3 * t, window)
ans = mod_add(ans, mod_mul(factor, window))
odd_t_max = (sqrt_limit - 1) // 2
add_seq = PowerSeq.triangular(w)
trim_seq = PowerSeq.triangular(w)
window = 0
factor = winv
ratio = winv10
for t in range(0, odd_t_max + 1):
disc = 4 * limit + 20 * t * t + 20 * t + 5
vmax = (math.isqrt(disc) - 1) // 2
window = add_seq.extend_through(vmax, window)
window = trim_seq.trim_before(3 * t + 1, window)
ans = mod_add(ans, mod_mul(factor, window))
factor = mod_mul(factor, ratio)
ratio = mod_mul(ratio, winv10)
return ans
def square_powers(base, max_g):
out = array("I", [0]) * (max_g + 1)
out[0] = 1
if max_g == 0:
return out
base_sq = mod_mul(base, base)
value = 1
ratio = base
for g in range(1, max_g + 1):
value = mod_mul(value, ratio)
out[g] = value
ratio = mod_mul(ratio, base_sq)
return out
def choose_process_count(allow_multiprocessing, requested_processes, workload):
if not allow_multiprocessing or workload <= 1 or workload < 1_000_000:
return 1
processes = requested_processes
if processes == 0:
processes = min(mp.cpu_count() or 1, 16)
if processes > workload:
processes = workload
return max(1, processes)
def _set_worker_state(limit, mu, phi_sq, phi_inv_sq):
global _WORK_LIMIT, _WORK_MU, _WORK_PHI_SQ, _WORK_PHI_INV_SQ
_WORK_LIMIT = limit
_WORK_MU = mu
_WORK_PHI_SQ = phi_sq
_WORK_PHI_INV_SQ = phi_inv_sq
def _primitive_worker(bounds):
start_g, end_g = bounds
limit = _WORK_LIMIT
mu = _WORK_MU
phi_sq = _WORK_PHI_SQ
phi_inv_sq = _WORK_PHI_INV_SQ
phi_total = 0
psi_total = 0
for g in range(start_g, end_g):
mu_g = mu[g]
if mu_g == 0:
continue
gg = g * g
scaled_limit = limit // gg
phi_w = phi_sq[g]
phi_winv = phi_inv_sq[g]
if (g & 1) == 0:
psi_w = phi_winv
psi_winv = phi_w
else:
psi_w = mod_neg(phi_winv)
psi_winv = mod_neg(phi_w)
phi_value = nonprimitive_sum(scaled_limit, phi_w, phi_winv)
psi_value = nonprimitive_sum(scaled_limit, psi_w, psi_winv)
if mu_g > 0:
phi_total += phi_value
psi_total += psi_value
else:
phi_total -= phi_value
psi_total -= psi_value
return phi_total, psi_total
def build_work_chunks(max_g, process_count):
chunk_count = max(process_count * 16, 1)
chunk_size = max(1, (max_g + chunk_count - 1) // chunk_count)
bounds = []
start_g = 1
while start_g <= max_g:
end_g = min(max_g + 1, start_g + chunk_size)
bounds.append((start_g, end_g))
start_g = end_g
return bounds
def primitive_sums(limit, mu, phi_sq, phi_inv_sq, options):
max_g = math.isqrt(limit)
process_count = choose_process_count(
options.allow_multiprocessing,
options.requested_processes,
max_g,
)
_set_worker_state(limit, mu, phi_sq, phi_inv_sq)
if process_count == 1:
phi_total, psi_total = _primitive_worker((1, max_g + 1))
return normalize_signed_mod(phi_total), normalize_signed_mod(psi_total)
try:
ctx = mp.get_context("fork")
except ValueError:
phi_total, psi_total = _primitive_worker((1, max_g + 1))
return normalize_signed_mod(phi_total), normalize_signed_mod(psi_total)
bounds = build_work_chunks(max_g, process_count)
with ctx.Pool(process_count) as pool:
totals = list(pool.imap_unordered(_primitive_worker, bounds, chunksize=1))
phi_total = sum(item[0] for item in totals)
psi_total = sum(item[1] for item in totals)
return normalize_signed_mod(phi_total), normalize_signed_mod(psi_total)
def solve(limit, options=None):
if options is None:
options = Options()
max_g = math.isqrt(limit)
mu = mobius_sieve(max_g)
phi_inv = mod_pow(PHI_MOD, MOD - 2)
phi_sq = square_powers(PHI_MOD, max_g)
phi_inv_sq = square_powers(phi_inv, max_g)
phi_sum, psi_sum = primitive_sums(limit, mu, phi_sq, phi_inv_sq, options)
inv_sqrt5 = mod_pow(SQRT5_MOD, MOD - 2)
return mod_mul(mod_sub(phi_sum, psi_sum), inv_sqrt5)
def checksum_via_factorization(limit, primes):
acc = 0
f_prev = 0
f_cur = 1
for n in range(1, limit + 1):
g = g_from_factorization(n, primes)
acc = mod_add(acc, mod_mul(f_cur, g))
f_prev, f_cur = f_cur, mod_add(f_prev, f_cur)
return acc
def run_checkpoints():
options = Options()
options.allow_multiprocessing = False
assert mod_mul(SQRT5_MOD, SQRT5_MOD) == 5
assert mod_sub(mod_mul(PHI_MOD, PHI_MOD), PHI_MOD) == 1
assert mod_sub(mod_mul(PSI_MOD, PSI_MOD), PSI_MOD) == 1
assert mod_mul(PHI_MOD, PSI_MOD) == MOD - 1
check_max = 250
prime_limit = math.isqrt(check_max) + 10
primes = sieve_primes(prime_limit)
for n in range(1, check_max + 1):
brute = g_bruteforce(n)
factorized = g_from_factorization(n, primes)
assert brute == factorized
for n in range(1, check_max + 1):
reduced_pairs = reduced_pair_count(n)
factorized = g_from_factorization(n, primes)
assert reduced_pairs == factorized
direct_limit = min(check_max, 250)
phi_pair_sum = direct_pair_sum(direct_limit, PHI_MOD)
psi_pair_sum = direct_pair_sum(direct_limit, PSI_MOD)
max_g = math.isqrt(direct_limit)
mu = mobius_sieve(max_g)
phi_inv = mod_pow(PHI_MOD, MOD - 2)
phi_sq = square_powers(PHI_MOD, max_g)
phi_inv_sq = square_powers(phi_inv, max_g)
phi_fast, psi_fast = primitive_sums(direct_limit, mu, phi_sq, phi_inv_sq, options)
assert phi_fast == phi_pair_sum
assert psi_fast == psi_pair_sum
sample_prime_limit = math.isqrt(SAMPLE_LIMIT) + 10
sample_primes = sieve_primes(sample_prime_limit)
sample_fast = solve(SAMPLE_LIMIT, options)
sample_factorized = checksum_via_factorization(SAMPLE_LIMIT, sample_primes)
assert sample_fast == SAMPLE_SUM
assert sample_fast == sample_factorized
def usage():
print(
"Usage:\n"
" python Euler989.py [--skip-checkpoints] [--single-thread] [--threads=N]\n"
" python Euler989.py validate [check_max]\n"
" python Euler989.py sum <limit> [--single-thread] [--threads=N]\n"
" python Euler989.py answer [--single-thread] [--threads=N]",
file=sys.stderr,
)
def parse_unsigned_after_prefix(arg, prefix):
if not arg.startswith(prefix):
return None
tail = arg[len(prefix) :]
if not tail or not tail.isdigit():
return None
return int(tail)
def parse_command_options(args):
options = Options()
positional = []
for arg in args:
if arg in ("--single-process", "--single-thread"):
options.allow_multiprocessing = False
continue
processes = parse_unsigned_after_prefix(arg, "--processes=")
if processes is not None:
options.requested_processes = processes
continue
threads = parse_unsigned_after_prefix(arg, "--threads=")
if threads is not None:
options.requested_processes = threads
continue
positional.append(arg)
return options, positional
def main(argv):
args = list(argv[1:])
should_run_checkpoints = True
if "--skip-checkpoints" in args:
should_run_checkpoints = False
args.remove("--skip-checkpoints")
if should_run_checkpoints:
run_checkpoints()
options, positional = parse_command_options(args)
if not positional:
print(solve(TARGET_LIMIT, options))
return 0
command = positional[0]
if command == "validate":
if len(positional) > 2:
usage()
return 1
check_max = 250 if len(positional) == 1 else int(positional[1])
prime_limit = math.isqrt(check_max) + 10
primes = sieve_primes(prime_limit)
for n in range(1, check_max + 1):
assert g_bruteforce(n) == g_from_factorization(n, primes)
assert reduced_pair_count(n) == g_from_factorization(n, primes)
print("ok")
return 0
if command == "sum" and len(positional) == 2:
print(solve(int(positional[1]), options))
return 0
if command == "answer" and len(positional) == 1:
print(solve(TARGET_LIMIT, options))
return 0
usage()
return 1
if __name__ == "__main__":
raise SystemExit(main(sys.argv))
Java
import java.util.ArrayList;
import java.util.Arrays;
public class Euler989 {
private static final long MOD = 1_000_000_009L;
private static final long TARGET_LIMIT = 100_000_000_000_000L;
private static final long SAMPLE_LIMIT = 1_000L;
private static final long SAMPLE_SUM = 190_950_976L;
private static final long SQRT5_MOD = 383_008_016L;
private static final long PHI_MOD = 691_504_013L;
private static final long PSI_MOD = 308_495_997L;
private static final class Options {
boolean allowMultithreading = true;
int requestedThreads = 0;
}
private static final class IntList {
private int[] data = new int[16];
private int size = 0;
void add(int value) {
if (size == data.length) {
data = Arrays.copyOf(data, data.length * 2);
}
data[size++] = value;
}
int get(int index) {
return data[index];
}
int size() {
return size;
}
int[] toArray() {
return Arrays.copyOf(data, size);
}
}
private static final class PowerSeq {
long v;
long term;
long delta;
long deltaStep;
PowerSeq(long v, long term, long delta, long deltaStep) {
this.v = v;
this.term = term;
this.delta = delta;
this.deltaStep = deltaStep;
}
static PowerSeq square(long base) {
long baseSq = modMul(base, base);
return new PowerSeq(0L, 1L, base, baseSq);
}
static PowerSeq triangular(long base) {
long baseSq = modMul(base, base);
return new PowerSeq(0L, 1L, baseSq, baseSq);
}
void step() {
term = modMul(term, delta);
delta = modMul(delta, deltaStep);
++v;
}
long extendThrough(long target, long window) {
while (v <= target) {
window = modAdd(window, term);
step();
}
return window;
}
long trimBefore(long target, long window) {
while (v < target) {
window = modSub(window, term);
step();
}
return window;
}
}
private static final class PrimitiveWorker implements Runnable {
long limit;
int startG;
int endG;
byte[] mu;
int[] phiSq;
int[] phiInvSq;
long phiTotal;
long psiTotal;
@Override
public void run() {
long localPhi = 0L;
long localPsi = 0L;
for (int g = startG; g < endG; ++g) {
int muG = mu[g];
if (muG == 0) {
continue;
}
long gg = (long) g * (long) g;
long scaledLimit = limit / gg;
long phiW = phiSq[g] & 0xFFFFFFFFL;
long phiWInv = phiInvSq[g] & 0xFFFFFFFFL;
long psiW = (g & 1) == 0 ? phiWInv : modNeg(phiWInv);
long psiWInv = (g & 1) == 0 ? phiW : modNeg(phiW);
long phiValue = nonprimitiveSum(scaledLimit, phiW, phiWInv);
long psiValue = nonprimitiveSum(scaledLimit, psiW, psiWInv);
if (muG > 0) {
localPhi += phiValue;
localPsi += psiValue;
} else {
localPhi -= phiValue;
localPsi -= psiValue;
}
}
phiTotal = localPhi;
psiTotal = localPsi;
}
}
private static void usage() {
System.err.println(
"Usage:\n"
+ " java Euler989 [--skip-checkpoints]\n"
+ " java Euler989 validate [check_max] [--single-thread] [--threads=N]\n"
+ " java Euler989 sum <limit> [--single-thread] [--threads=N]\n"
+ " java Euler989 answer [--single-thread] [--threads=N]");
}
private static long modMul(long a, long b) {
return (a * b) % MOD;
}
private static long modAdd(long a, long b) {
long s = a + b;
return s >= MOD ? s - MOD : s;
}
private static long modSub(long a, long b) {
return a >= b ? a - b : a + MOD - b;
}
private static long modNeg(long a) {
return a == 0L ? 0L : MOD - a;
}
private static long modPow(long base, long exp) {
long result = 1L;
long cur = base % MOD;
while (exp > 0L) {
if ((exp & 1L) != 0L) {
result = modMul(result, cur);
}
cur = modMul(cur, cur);
exp >>= 1;
}
return result;
}
private static long isqrt(long n) {
long x = (long) Math.sqrt((double) n);
while ((x + 1L) <= n / (x + 1L)) {
++x;
}
while (x > n / x) {
--x;
}
return x;
}
private static int[] sievePrimes(int limit) {
if (limit < 2) {
return new int[0];
}
boolean[] isPrime = new boolean[limit + 1];
Arrays.fill(isPrime, true);
isPrime[0] = false;
isPrime[1] = false;
IntList primes = new IntList();
for (int i = 2; i <= limit; ++i) {
if (!isPrime[i]) {
continue;
}
primes.add(i);
if (i > limit / i) {
continue;
}
for (int j = i * i; j <= limit; j += i) {
isPrime[j] = false;
}
}
return primes.toArray();
}
private static byte[] mobiusSieve(int limit) {
byte[] mu = new byte[limit + 1];
int[] least = new int[limit + 1];
IntList primes = new IntList();
mu[1] = 1;
for (int i = 2; i <= limit; ++i) {
if (least[i] == 0) {
least[i] = i;
primes.add(i);
mu[i] = -1;
}
int li = least[i];
byte mui = mu[i];
for (int idx = 0; idx < primes.size(); ++idx) {
int p = primes.get(idx);
long ip = (long) i * (long) p;
if (ip > limit || p > li) {
break;
}
least[(int) ip] = p;
if (p == li) {
mu[(int) ip] = 0;
break;
}
mu[(int) ip] = (byte) (-mui);
}
}
return mu;
}
private static int gBruteforce(long n) {
int count = 0;
for (long x = 0; x < n; ++x) {
if ((x * x + n - x - 1L) % n == 0L) {
++count;
}
}
return count;
}
private static int gFromFactorization(long n, int[] primes) {
if (n == 1L) {
return 1;
}
int result = 1;
for (int p32 : primes) {
long p = p32;
if (p > n / p) {
break;
}
if (n % p != 0L) {
continue;
}
int exponent = 0;
while (n % p == 0L) {
n /= p;
++exponent;
}
if (p == 2L) {
return 0;
}
if (p == 5L) {
if (exponent >= 2) {
return 0;
}
continue;
}
long r = p % 5L;
if (r == 1L || r == 4L) {
result *= 2;
} else if (r == 2L || r == 3L) {
return 0;
} else {
throw new IllegalStateException("unreachable");
}
}
if (n > 1L) {
if (n == 2L) {
return 0;
}
if (n == 5L) {
return result;
}
long r = n % 5L;
if (r == 1L || r == 4L) {
result *= 2;
} else if (r == 2L || r == 3L) {
return 0;
} else {
throw new IllegalStateException("unreachable");
}
}
return result;
}
private static long qForm(long a, long b) {
return a * a - a * b - b * b;
}
private static long gcd(long a, long b) {
while (b != 0L) {
long t = a % b;
a = b;
b = t;
}
return a;
}
private static int reducedPairCount(long n) {
int count = 0;
long maxB = isqrt(n);
for (long b = 1L; b <= maxB; ++b) {
for (long a = 2L * b;; ++a) {
long q = qForm(a, b);
if (q > n) {
break;
}
if (q == n && gcd(a, b) == 1L) {
++count;
}
}
}
return count;
}
private static long directPairSum(long limit, long base) {
long acc = 0L;
long maxB = isqrt(limit);
for (long b = 1L; b <= maxB; ++b) {
for (long a = 2L * b;; ++a) {
long q = qForm(a, b);
if (q > limit) {
break;
}
if (gcd(a, b) == 1L) {
acc = modAdd(acc, modPow(base, q));
}
}
}
return acc;
}
private static long nonprimitiveSum(long limit, long w, long winv) {
if (limit == 0L) {
return 0L;
}
long sqrtLimit = isqrt(limit);
long winv5 = modPow(winv, 5L);
long winv10 = modMul(winv5, winv5);
long ans = 0L;
long evenTMax = sqrtLimit / 2L;
if (evenTMax > 0L) {
PowerSeq addSeq = PowerSeq.square(w);
PowerSeq trimSeq = PowerSeq.square(w);
long window = 0L;
long factor = 1L;
long ratio = winv5;
for (long t = 1L; t <= evenTMax; ++t) {
factor = modMul(factor, ratio);
ratio = modMul(ratio, winv10);
long vmax = isqrt(limit + 5L * t * t);
window = addSeq.extendThrough(vmax, window);
window = trimSeq.trimBefore(3L * t, window);
ans = modAdd(ans, modMul(factor, window));
}
}
long oddTMax = (sqrtLimit - 1L) / 2L;
PowerSeq addSeq = PowerSeq.triangular(w);
PowerSeq trimSeq = PowerSeq.triangular(w);
long window = 0L;
long factor = winv;
long ratio = winv10;
for (long t = 0L; t <= oddTMax; ++t) {
long disc = 4L * limit + 20L * t * t + 20L * t + 5L;
long vmax = (isqrt(disc) - 1L) / 2L;
window = addSeq.extendThrough(vmax, window);
window = trimSeq.trimBefore(3L * t + 1L, window);
ans = modAdd(ans, modMul(factor, window));
factor = modMul(factor, ratio);
ratio = modMul(ratio, winv10);
}
return ans;
}
private static int[] squarePowers(long base, int maxG) {
int[] out = new int[maxG + 1];
out[0] = 1;
if (maxG == 0) {
return out;
}
long baseSq = modMul(base, base);
long value = 1L;
long ratio = base;
for (int g = 1; g <= maxG; ++g) {
value = modMul(value, ratio);
out[g] = (int) value;
ratio = modMul(ratio, baseSq);
}
return out;
}
private static long normalizeSignedMod(long value) {
long res = value % MOD;
if (res < 0L) {
res += MOD;
}
return res;
}
private static int chooseThreadCount(boolean allowMultithreading, int requestedThreads, int workload) {
if (!allowMultithreading || workload <= 1) {
return 1;
}
int threads = requestedThreads;
if (threads == 0) {
threads = Runtime.getRuntime().availableProcessors();
if (threads == 0) {
threads = 1;
}
}
if (threads > workload) {
threads = workload;
}
return Math.max(1, threads);
}
private static long[] primitiveSums(long limit, byte[] mu, int[] phiSq, int[] phiInvSq, Options options)
throws InterruptedException {
int maxG = (int) isqrt(limit);
int threadCount = chooseThreadCount(options.allowMultithreading, options.requestedThreads, maxG);
PrimitiveWorker[] tasks = new PrimitiveWorker[threadCount];
Thread[] threads = new Thread[threadCount];
for (int t = 0; t < threadCount; ++t) {
PrimitiveWorker task = new PrimitiveWorker();
task.limit = limit;
task.startG = 1 + (int) (((long) maxG * (long) t) / (long) threadCount);
task.endG = 1 + (int) (((long) maxG * (long) (t + 1)) / (long) threadCount);
task.mu = mu;
task.phiSq = phiSq;
task.phiInvSq = phiInvSq;
tasks[t] = task;
threads[t] = new Thread(task);
threads[t].start();
}
long phiTotal = 0L;
long psiTotal = 0L;
for (int t = 0; t < threadCount; ++t) {
threads[t].join();
phiTotal += tasks[t].phiTotal;
psiTotal += tasks[t].psiTotal;
}
return new long[] { normalizeSignedMod(phiTotal), normalizeSignedMod(psiTotal) };
}
private static long solve(long limit, Options options) throws InterruptedException {
int maxG = (int) isqrt(limit);
byte[] mu = mobiusSieve(maxG);
long phiInv = modPow(PHI_MOD, MOD - 2L);
int[] phiSq = squarePowers(PHI_MOD, maxG);
int[] phiInvSq = squarePowers(phiInv, maxG);
long[] sums = primitiveSums(limit, mu, phiSq, phiInvSq, options);
long invSqrt5 = modPow(SQRT5_MOD, MOD - 2L);
return modMul(modSub(sums[0], sums[1]), invSqrt5);
}
private static long checksumViaFactorization(long limit, int[] primes) {
long acc = 0L;
long fPrev = 0L;
long fCur = 1L;
for (long n = 1L; n <= limit; ++n) {
long g = gFromFactorization(n, primes);
acc = modAdd(acc, modMul(fCur, g));
long fNext = modAdd(fPrev, fCur);
fPrev = fCur;
fCur = fNext;
}
return acc;
}
private static void require(boolean condition, String message) {
if (!condition) {
throw new IllegalStateException(message);
}
}
private static void runCheckpoints(Options options) throws InterruptedException {
require(modMul(SQRT5_MOD, SQRT5_MOD) == 5L, "sqrt(5) mismatch");
require(modSub(modMul(PHI_MOD, PHI_MOD), PHI_MOD) == 1L, "phi mismatch");
require(modSub(modMul(PSI_MOD, PSI_MOD), PSI_MOD) == 1L, "psi mismatch");
require(modMul(PHI_MOD, PSI_MOD) == MOD - 1L, "phi*psi mismatch");
int checkMax = 250;
int[] primes = sievePrimes((int) isqrt(checkMax) + 10);
for (long n = 1L; n <= checkMax; ++n) {
require(gBruteforce(n) == gFromFactorization(n, primes), "factorization mismatch");
require(reducedPairCount(n) == gFromFactorization(n, primes), "pair-count mismatch");
}
long directLimit = Math.min(checkMax, 250);
long phiPairSum = directPairSum(directLimit, PHI_MOD);
long psiPairSum = directPairSum(directLimit, PSI_MOD);
int maxG = (int) isqrt(directLimit);
byte[] mu = mobiusSieve(maxG);
long phiInv = modPow(PHI_MOD, MOD - 2L);
int[] phiSq = squarePowers(PHI_MOD, maxG);
int[] phiInvSq = squarePowers(phiInv, maxG);
long[] fast = primitiveSums(directLimit, mu, phiSq, phiInvSq, options);
require(fast[0] == phiPairSum, "phi pair sum mismatch");
require(fast[1] == psiPairSum, "psi pair sum mismatch");
int[] samplePrimes = sievePrimes((int) isqrt(SAMPLE_LIMIT) + 10);
long sampleFast = solve(SAMPLE_LIMIT, options);
long sampleFactorized = checksumViaFactorization(SAMPLE_LIMIT, samplePrimes);
require(sampleFast == SAMPLE_SUM, "sample sum mismatch");
require(sampleFast == sampleFactorized, "sample factorization mismatch");
}
private static boolean parseU64(String text, long[] out) {
if (text.isEmpty()) {
return false;
}
long value = 0L;
for (int i = 0; i < text.length(); ++i) {
char c = text.charAt(i);
if (c < '0' || c > '9') {
return false;
}
value = value * 10L + (long) (c - '0');
}
out[0] = value;
return true;
}
private static boolean parseCommandOptions(String[] args, int startIndex, Options options, ArrayList<String> positional) {
for (int i = startIndex; i < args.length; ++i) {
String arg = args[i];
if ("--single-thread".equals(arg)) {
options.allowMultithreading = false;
continue;
}
if ("--single-process".equals(arg)) {
options.allowMultithreading = false;
continue;
}
if (arg.startsWith("--threads=")) {
String tail = arg.substring("--threads=".length());
if (tail.isEmpty()) {
return false;
}
try {
options.requestedThreads = Integer.parseInt(tail);
} catch (NumberFormatException ex) {
return false;
}
continue;
}
if (arg.startsWith("--processes=")) {
String tail = arg.substring("--processes=".length());
if (tail.isEmpty()) {
return false;
}
try {
options.requestedThreads = Integer.parseInt(tail);
} catch (NumberFormatException ex) {
return false;
}
continue;
}
positional.add(arg);
}
return true;
}
public static void main(String[] args) throws Exception {
Options options = new Options();
ArrayList<String> positional = new ArrayList<>();
boolean skipCheckpoints = false;
ArrayList<String> cleaned = new ArrayList<>();
for (String arg : args) {
if ("--skip-checkpoints".equals(arg)) {
skipCheckpoints = true;
} else {
cleaned.add(arg);
}
}
args = cleaned.toArray(new String[0]);
if (!skipCheckpoints) {
runCheckpoints(options);
}
if (args.length == 0) {
System.out.println(solve(TARGET_LIMIT, options));
return;
}
String command = args[0];
options = new Options();
positional.clear();
if (!parseCommandOptions(args, 1, options, positional)) {
usage();
return;
}
if ("validate".equals(command)) {
long checkMax = 250L;
if (positional.size() > 1) {
usage();
return;
}
if (positional.size() == 1) {
long[] parsed = new long[1];
if (!parseU64(positional.get(0), parsed)) {
usage();
return;
}
checkMax = parsed[0];
}
int[] primes = sievePrimes((int) isqrt(checkMax) + 10);
for (long n = 1L; n <= checkMax; ++n) {
require(gBruteforce(n) == gFromFactorization(n, primes), "factorization mismatch");
require(reducedPairCount(n) == gFromFactorization(n, primes), "pair-count mismatch");
}
System.out.println("ok");
return;
}
if ("sum".equals(command) && positional.size() == 1) {
long[] parsed = new long[1];
if (!parseU64(positional.get(0), parsed)) {
usage();
return;
}
System.out.println(solve(parsed[0], options));
return;
}
if ("answer".equals(command) && positional.isEmpty()) {
System.out.println(solve(TARGET_LIMIT, options));
return;
}
usage();
}
}