Problem 457: A Polynomial Modulo the Square of a Prime

View on Project Euler

Project Euler Problem 457 Solution

EulerSolve provides an optimized solution for Project Euler Problem 457, A Polynomial Modulo the Square of a Prime, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary For each odd prime \(p\), let \(R(p)\) be the least positive integer \(n\) satisfying $$n^2-3n-1\equiv 0\pmod{p^2}.$$ The goal is to compute $$S(L)=\sum_{p\le L} R(p)$$ for \(L=10^7\). Useful checkpoints are \(S(10)=5\), \(S(100)=1752\), and \(S(1000)=6728355\). Mathematical Approach Write $$f(n)=n^2-3n-1.$$ The implementation does not search over all \(n\). Instead it first determines when roots can exist modulo \(p\), computes those roots explicitly, and then lifts them once to modulo \(p^2\). Step 1: Rewrite the Congruence Multiplying by \(4\) and completing the square gives $$4f(n)=4n^2-12n-4=(2n-3)^2-13.$$ Because \(p\) is odd, \(2\) is invertible modulo \(p\). Therefore $$f(n)\equiv 0\pmod{p}\iff (2n-3)^2\equiv 13\pmod{p}.$$ So the congruence has roots modulo \(p\) exactly when \(13\) is a quadratic residue modulo \(p\). Step 2: Which Primes Can Contribute? For odd primes \(p\neq 13\), quadratic reciprocity gives $$\left(\frac{13}{p}\right)=\left(\frac{p}{13}\right),$$ because \(13\equiv 1\pmod{4}\). The nonzero quadratic residues modulo \(13\) are $$1,\ 3,\ 4,\ 9,\ 10,\ 12.$$ Hence roots can exist only for primes in those six residue classes modulo \(13\). This explains the residue-class filter used by the implementation before any expensive modular square-root work is attempted. The prime \(p=13\) must be handled separately....

Detailed mathematical approach

Problem Summary

For each odd prime \(p\), let \(R(p)\) be the least positive integer \(n\) satisfying

$$n^2-3n-1\equiv 0\pmod{p^2}.$$

The goal is to compute

$$S(L)=\sum_{p\le L} R(p)$$

for \(L=10^7\). Useful checkpoints are \(S(10)=5\), \(S(100)=1752\), and \(S(1000)=6728355\).

Mathematical Approach

Write

$$f(n)=n^2-3n-1.$$

The implementation does not search over all \(n\). Instead it first determines when roots can exist modulo \(p\), computes those roots explicitly, and then lifts them once to modulo \(p^2\).

Step 1: Rewrite the Congruence

Multiplying by \(4\) and completing the square gives

$$4f(n)=4n^2-12n-4=(2n-3)^2-13.$$

Because \(p\) is odd, \(2\) is invertible modulo \(p\). Therefore

$$f(n)\equiv 0\pmod{p}\iff (2n-3)^2\equiv 13\pmod{p}.$$

So the congruence has roots modulo \(p\) exactly when \(13\) is a quadratic residue modulo \(p\).

Step 2: Which Primes Can Contribute?

For odd primes \(p\neq 13\), quadratic reciprocity gives

$$\left(\frac{13}{p}\right)=\left(\frac{p}{13}\right),$$

because \(13\equiv 1\pmod{4}\). The nonzero quadratic residues modulo \(13\) are

$$1,\ 3,\ 4,\ 9,\ 10,\ 12.$$

Hence roots can exist only for primes in those six residue classes modulo \(13\). This explains the residue-class filter used by the implementation before any expensive modular square-root work is attempted.

The prime \(p=13\) must be handled separately. Modulo \(13\), the congruence becomes

$$ (2n-3)^2\equiv 0\pmod{13}, $$

so there is the double root \(n\equiv 8\pmod{13}\). But

$$f(8)=8^2-3\cdot 8-1=39,$$

which is divisible by \(13\) but not by \(13^2=169\). Therefore the root modulo \(13\) does not lift to a root modulo \(13^2\), and \(p=13\) contributes nothing.

Step 3: Explicit Roots Modulo \(p\)

Assume now that \(p\neq 13\) and that \(13\) is a quadratic residue modulo \(p\). If

$$s^2\equiv 13\pmod{p},$$

then from \(2n-3\equiv \pm s\pmod{p}\) we obtain the two roots

$$r_1\equiv \frac{3+s}{2}\pmod{p},\qquad r_2\equiv \frac{3-s}{2}\pmod{p}.$$

A modular square root \(s\) is found with the Tonelli-Shanks algorithm. Since \(2^{-1}\equiv (p+1)/2\pmod{p}\), the division by \(2\) is just another modular multiplication.

Step 4: Lift Each Root to Modulo \(p^2\)

Let \(r\) be one of the roots modulo \(p\). Any lift to modulo \(p^2\) has the form

$$n=r+tp$$

for some \(t\in\{0,1,\dots,p-1\}\). Expand \(f(r+tp)\):

$$f(r+tp)=f(r)+tp\,f'(r)+t^2p^2,$$

where

$$f'(x)=2x-3.$$

Reducing modulo \(p^2\) removes the last term, so the condition \(f(r+tp)\equiv 0\pmod{p^2}\) becomes

$$f(r)+tp\,f'(r)\equiv 0\pmod{p^2}.$$

Because \(r\) is already a root modulo \(p\), the value \(f(r)\) is divisible by \(p\). Dividing by \(p\) yields the linear congruence

$$t\,f'(r)\equiv -\frac{f(r)}{p}\pmod{p},$$

and therefore

$$t\equiv -\frac{f(r)/p}{f'(r)}\pmod{p}.$$

For \(p\neq 13\), we have \(f'(r)=2r-3\equiv \pm s\not\equiv 0\pmod{p}\), so the inverse exists. Thus each root modulo \(p\) lifts to exactly one root modulo \(p^2\).

Step 5: Determine \(R(p)\)

The two roots \(r_1\) and \(r_2\) modulo \(p\) produce two lifted roots \(n_1\) and \(n_2\) modulo \(p^2\). The definition of \(R(p)\) asks for the least positive solution, so

$$R(p)=\min(n_1,n_2).$$

Summing this quantity over all contributing primes gives the desired value \(S(L)\).

Worked Example: \(p=3\)

Here \(13\equiv 1\pmod{3}\), so we may take \(s\equiv 1\). The two roots modulo \(3\) are

$$r_1\equiv \frac{3+1}{2}\equiv 2\pmod{3},\qquad r_2\equiv \frac{3-1}{2}\equiv 1\pmod{3}.$$

For \(r=2\),

$$f(2)=4-6-1=-3,\qquad f'(2)=1.$$

So

$$t\equiv -\frac{-3/3}{1}\equiv 1\pmod{3},$$

and the lifted root is

$$n=2+1\cdot 3=5.$$

For \(r=1\),

$$f(1)=1-3-1=-3,\qquad f'(1)=-1\equiv 2\pmod{3}.$$

Since \(2^{-1}\equiv 2\pmod{3}\),

$$t\equiv -\frac{-3/3}{2}\equiv 2\pmod{3},$$

which gives

$$n=1+2\cdot 3=7.$$

Therefore \(R(3)=5\), and because \(3\) is the only contributing prime up to \(10\), we recover the checkpoint \(S(10)=5\).

How the Code Works

The C++, Python, and Java implementations all use the same structure. They begin with a prime sieve up to \(L\). After excluding \(2\) and \(13\), they keep only primes whose residue modulo \(13\) is one of \(1,3,4,9,10,12\). For each remaining prime, the implementation computes a square root of \(13\) modulo \(p\) via Tonelli-Shanks, constructs the two roots modulo \(p\), applies the Hensel correction above to obtain the two roots modulo \(p^2\), and adds the smaller lifted value to the running sum. Fast modular exponentiation is reused for Euler-criterion tests, modular inverses, and the Tonelli-Shanks substeps.

Complexity Analysis

Generating all primes up to \(L\) with a sieve of Eratosthenes costs \(O(L\log\log L)\) time and \(O(L)\) memory. The residue-class filter removes roughly half of the odd primes immediately. Each surviving prime then needs only a constant amount of modular arithmetic plus one Tonelli-Shanks square-root computation, whose cost is polylogarithmic in \(p\). In practice the sieve dominates the memory usage, and the entire method is easily fast enough for \(L=10^7\).

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=457
  2. Quadratic reciprocity: Wikipedia - Quadratic reciprocity
  3. Quadratic residue and Legendre symbol: Wikipedia - Quadratic residue
  4. Hensel's lemma: Wikipedia - Hensel's lemma
  5. Tonelli-Shanks algorithm: Wikipedia - Tonelli-Shanks algorithm

Problem 457 source code

C++

#include <array>
#include <cstdint>
#include <iostream>
#include <string>
#include <vector>
#include <algorithm>
#include <functional>

namespace {

using u32 = std::uint32_t;
using u64 = std::uint64_t;
using i64 = std::int64_t;
using i128 = __int128_t;
using u128 = __uint128_t;

struct Options {
    u32 l = 10'000'000U;
    bool run_checkpoints = true;
};

bool parse_u32_after_prefix(const std::string& arg, const std::string& prefix, u32& out) {
    if (arg.rfind(prefix, 0U) != 0U) {
        return false;
    }
    const std::string tail = arg.substr(prefix.size());
    if (tail.empty()) {
        return false;
    }
    try {
        out = static_cast<u32>(std::stoul(tail));
    } catch (...) {
        return false;
    }
    return true;
}

bool parse_arguments(int argc, char** argv, Options& options) {
    for (int i = 1; i < argc; ++i) {
        const std::string arg(argv[i]);
        if (arg == "--skip-checkpoints") {
            options.run_checkpoints = false;
            continue;
        }
        if (parse_u32_after_prefix(arg, "--l=", options.l)) {
            continue;
        }
        std::cerr << "Unknown argument: " << arg << '\n';
        return false;
    }
    return options.l >= 2U;
}

u32 mod_pow(u64 base, u64 exp, const u32 mod) {
    u64 result = 1ULL;
    base %= mod;
    while (exp > 0ULL) {
        if ((exp & 1ULL) != 0ULL) {
            result = (result * base) % mod;
        }
        base = (base * base) % mod;
        exp >>= 1ULL;
    }
    return static_cast<u32>(result);
}

u32 tonelli_shanks(const u32 n, const u32 p) {
    if (p == 2U) {
        return n & 1U;
    }
    if (n == 0U) {
        return 0U;
    }
    if (mod_pow(n, (p - 1U) / 2U, p) != 1U) {
        return 0U;
    }
    if ((p & 3U) == 3U) {
        return mod_pow(n, (p + 1U) / 4U, p);
    }

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

    u32 z = 2U;
    while (mod_pow(z, (p - 1U) / 2U, p) != p - 1U) {
        ++z;
    }

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

    while (t != 1U) {
        u64 tt = t;
        u64 i = 0U;
        while (tt != 1U && i < m) {
            tt = (tt * tt) % p;
            ++i;
        }

        const u64 shift = m - i - 1U;
        const u64 b = mod_pow(static_cast<u32>(c), 1ULL << shift, p);
        r = (r * b) % p;
        const u64 b2 = (b * b) % p;
        t = (t * b2) % p;
        c = b2;
        m = i;
    }

    return static_cast<u32>(r);
}

u64 lift_root(const u32 p, const u32 r) {
    const u32 deriv = static_cast<u32>((2ULL * r + p - 3ULL) % p);
    const u32 inv_deriv = mod_pow(deriv, p - 2U, p);

    const i128 fr = static_cast<i128>(r) * static_cast<i128>(r) - 3 * static_cast<i128>(r) - 1;
    const i64 q = static_cast<i64>(fr / static_cast<i128>(p));

    i64 neg_q = -(q % static_cast<i64>(p));
    neg_q %= static_cast<i64>(p);
    if (neg_q < 0) {
        neg_q += p;
    }

    const u32 t = static_cast<u32>((static_cast<u64>(neg_q) * inv_deriv) % p);
    return static_cast<u64>(r) + static_cast<u64>(t) * static_cast<u64>(p);
}

u128 solve(const u32 limit) {
    std::vector<bool> is_prime(static_cast<std::size_t>(limit) + 1U, true);
    if (limit >= 0U) {
        is_prime[0] = false;
    }
    if (limit >= 1U) {
        is_prime[1] = false;
    }

    for (u32 i = 2U; static_cast<u64>(i) * i <= limit; ++i) {
        if (!is_prime[i]) {
            continue;
        }
        for (u32 j = i * i; j <= limit; j += i) {
            is_prime[j] = false;
        }
    }

    std::array<bool, 13> residue{};
    residue.fill(false);
    residue[1] = true;
    residue[3] = true;
    residue[4] = true;
    residue[9] = true;
    residue[10] = true;
    residue[12] = true;

    u128 sum = 0;
    for (u32 p = 2U; p <= limit; ++p) {
        if (!is_prime[p]) {
            continue;
        }
        if (p == 2U || p == 13U) {
            continue;
        }
        if (!residue[p % 13U]) {
            continue;
        }

        const u32 s = tonelli_shanks(13U % p, p);
        const u32 inv2 = (p + 1U) / 2U;

        const u32 r1 = static_cast<u32>((static_cast<u64>(3U + s) * inv2) % p);
        const u32 r2 = static_cast<u32>((static_cast<u64>(3U + p - s) * inv2) % p);

        const u64 n1 = lift_root(p, r1);
        const u64 n2 = lift_root(p, r2);
        sum += (n1 < n2 ? n1 : n2);
    }

    return sum;
}

std::string to_string_u128(u128 value) {
    if (value == 0) {
        return "0";
    }
    std::string out;
    while (value > 0) {
        const unsigned digit = static_cast<unsigned>(value % 10);
        out.push_back(static_cast<char>('0' + digit));
        value /= 10;
    }
    std::reverse(out.begin(), out.end());
    return out;
}

bool run_checkpoints() {
    if (solve(10U) != 5U) {
        std::cerr << "Checkpoint failed: SR(10)\n";
        return false;
    }
    if (solve(100U) != 1'752U) {
        std::cerr << "Checkpoint failed: SR(100)\n";
        return false;
    }
    if (solve(1000U) != 6'728'355ULL) {
        std::cerr << "Checkpoint failed: SR(1000)\n";
        return false;
    }
    return true;
}

}  // namespace

int main(int argc, char** argv) {
    Options options;
    if (!parse_arguments(argc, argv, options)) {
        return 1;
    }

    if (options.run_checkpoints && !run_checkpoints()) {
        return 2;
    }

    std::cout << to_string_u128(solve(options.l)) << '\n';
    return 0;
}

Python

import sys
import math

def mod_pow(base, exp, mod):
    result = 1
    base %= mod
    while exp > 0:
        if exp & 1:
            result = (result * base) % mod
        base = (base * base) % mod
        exp >>= 1
    return result

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

def lift_root(p, r):
    deriv = (2 * r + p - 3) % p
    inv_deriv = mod_pow(deriv, p - 2, p)
    
    fr = r * r - 3 * r - 1
    q = fr // p
    
    neg_q = -q % p
    if neg_q < 0: neg_q += p
    
    t = (neg_q * inv_deriv) % p
    return r + t * p

def solve_limit(limit):
    is_prime = bytearray([1]) * (limit + 1)
    if limit >= 0: is_prime[0] = 0
    if limit >= 1: is_prime[1] = 0
    
    for i in range(2, math.isqrt(limit) + 1):
        if is_prime[i]:
            is_prime[i*i : limit+1 : i] = bytearray([0]) * len(range(i*i, limit+1, i))
            
    residue = [False] * 13
    for i in [1, 3, 4, 9, 10, 12]:
        residue[i] = True
        
    total_sum = 0
    for p in range(2, limit + 1):
        if not is_prime[p] or p == 2 or p == 13:
            continue
            
        if not residue[p % 13]:
            continue
            
        s = tonelli_shanks(13 % p, p)
        inv2 = (p + 1) // 2
        
        r1 = ((3 + s) * inv2) % p
        r2 = ((3 + p - s) * inv2) % p
        
        n1 = lift_root(p, r1)
        n2 = lift_root(p, r2)
        total_sum += n1 if n1 < n2 else n2
        
    return total_sum

def solve():
    return str(solve_limit(10000000))

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

Java

public class Euler457 {

    static long modPow(long base, long exp, long mod) {
        long result = 1;
        base %= mod;
        while (exp > 0) {
            if ((exp & 1) != 0) {
                result = (result * base) % mod;
            }
            base = (base * base) % mod;
            exp >>= 1;
        }
        return result;
    }

    static long tonelliShanks(long n, long p) {
        if (p == 2)
            return n & 1;
        if (n == 0)
            return 0;
        if (modPow(n, (p - 1) / 2, p) != 1)
            return 0;
        if ((p & 3) == 3)
            return modPow(n, (p + 1) / 4, p);

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

        long z = 2;
        while (modPow(z, (p - 1) / 2, p) != p - 1) {
            z++;
        }

        long m = s;
        long c = modPow(z, q, p);
        long t = modPow(n, q, p);
        long r = modPow(n, (q + 1) / 2, p);

        while (t != 1) {
            long tt = t;
            long i = 0;
            while (tt != 1 && i < m) {
                tt = (tt * tt) % p;
                i++;
            }

            long shift = m - i - 1;
            long b = modPow(c, 1L << shift, p);
            r = (r * b) % p;
            long b2 = (b * b) % p;
            t = (t * b2) % p;
            c = b2;
            m = i;
        }

        return r;
    }

    static long liftRoot(long p, long r) {
        long deriv = (2 * r + p - 3) % p;
        long invDeriv = modPow(deriv, p - 2, p);

        long fr = r * r - 3 * r - 1;
        long q = fr / p;

        long negQ = -(q % p);
        negQ %= p;
        if (negQ < 0) {
            negQ += p;
        }

        long t = (negQ * invDeriv) % p;
        return r + t * p;
    }

    public static String solve() {
        int limit = 10000000;
        boolean[] isPrime = new boolean[limit + 1];
        for (int i = 2; i <= limit; i++)
            isPrime[i] = true;

        for (int i = 2; (long) i * i <= limit; i++) {
            if (isPrime[i]) {
                for (int j = i * i; j <= limit; j += i) {
                    isPrime[j] = false;
                }
            }
        }

        boolean[] residue = new boolean[13];
        residue[1] = true;
        residue[3] = true;
        residue[4] = true;
        residue[9] = true;
        residue[10] = true;
        residue[12] = true;

        long sum = 0;
        for (int p = 2; p <= limit; p++) {
            if (!isPrime[p] || p == 2 || p == 13)
                continue;
            if (!residue[p % 13])
                continue;

            long s = tonelliShanks(13 % p, p);
            long inv2 = (p + 1) / 2;

            long r1 = ((3 + s) * inv2) % p;
            long r2 = ((3 + p - s) * inv2) % p;

            long n1 = liftRoot(p, r1);
            long n2 = liftRoot(p, r2);

            sum += (n1 < n2) ? n1 : n2;
        }

        return Long.toString(sum);
    }

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