Problem 521: Smallest Prime Factor

View on Project Euler

Project Euler Problem 521 Solution

EulerSolve provides an optimized solution for Project Euler Problem 521, Smallest Prime Factor, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary For \(n\ge 2\), let \(s(n)\) be the smallest prime factor of \(n\). We must evaluate $$S(N)=\sum_{n=2}^{N} s(n), \qquad N=10^{12},$$ and return the result modulo \(10^9\). A direct sieve up to \(N\) would require linear memory and essentially linear time, so the implementation instead tracks how composite numbers disappear when primes are processed in increasing order. Mathematical Approach The key idea is that every composite number is removed exactly once, namely when its smallest prime factor is reached, while primes are never removed. The algorithm turns that observation into coupled recurrences for filtered counts and filtered sums. Step 1: Split the answer into prime and composite parts If \(n\) is prime, then \(s(n)=n\). If \(n\) is composite, then \(s(n)\le \sqrt n\le \sqrt N\). Therefore $$S(N)=\sum_{\substack{q\le N\\ q\text{ prime}}} q+\sum_{p\le \sqrt N} p\,M_p(N),$$ where \(M_p(N)\) denotes the number of composite integers \(n\le N\) whose smallest prime factor is exactly \(p\). So the entire task becomes: count those composites efficiently, and also recover the sum of all primes up to \(N\)....

Detailed mathematical approach

Problem Summary

For \(n\ge 2\), let \(s(n)\) be the smallest prime factor of \(n\). We must evaluate

$$S(N)=\sum_{n=2}^{N} s(n), \qquad N=10^{12},$$

and return the result modulo \(10^9\). A direct sieve up to \(N\) would require linear memory and essentially linear time, so the implementation instead tracks how composite numbers disappear when primes are processed in increasing order.

Mathematical Approach

The key idea is that every composite number is removed exactly once, namely when its smallest prime factor is reached, while primes are never removed. The algorithm turns that observation into coupled recurrences for filtered counts and filtered sums.

Step 1: Split the answer into prime and composite parts

If \(n\) is prime, then \(s(n)=n\). If \(n\) is composite, then \(s(n)\le \sqrt n\le \sqrt N\). Therefore

$$S(N)=\sum_{\substack{q\le N\\ q\text{ prime}}} q+\sum_{p\le \sqrt N} p\,M_p(N),$$

where \(M_p(N)\) denotes the number of composite integers \(n\le N\) whose smallest prime factor is exactly \(p\). So the entire task becomes: count those composites efficiently, and also recover the sum of all primes up to \(N\).

Step 2: Define the filtered set that survives before prime \(p\)

For a prime threshold \(p\), define the surviving set

$$\mathcal{R}_p(x)=\left\{m\in\{2,\dots,x\}: s(m)\ge p\right\}\cup\left\{q\le x:q\text{ prime},\ q<p\right\}.$$

This set contains two kinds of numbers: primes smaller than \(p\), which are kept forever, and numbers whose prime factors are all at least \(p\), which have not yet been removed. Now define

$$C_p(x)=\#\mathcal{R}_p(x), \qquad T_p(x)=\sum_{m\in\mathcal{R}_p(x)} m.$$

At the start no filtering has happened, so for \(p=2\) we simply have

$$C_2(x)=x-1,\qquad T_2(x)=\sum_{m=2}^{x} m=\frac{x(x+1)}{2}-1.$$

A useful consequence is that \(C_p(p)-C_p(p-1)=1\) exactly when \(p\) is prime, because every composite number has already disappeared before its own turn.

Step 3: Count composites whose smallest prime factor is \(p\)

Take a prime \(p\). Any composite \(n\le x\) with \(s(n)=p\) can be written uniquely as

$$n=p\,m,$$

where \(m\ge p\) and every prime factor of \(m\) is at least \(p\). Equivalently, \(m\) lies in \(\mathcal{R}_p(\lfloor x/p\rfloor)\), but the elements below \(p\) must be excluded because they are precisely the smaller primes. Hence

$$M_p(x)=C_p\!\left(\left\lfloor\frac{x}{p}\right\rfloor\right)-C_p(p-1).$$

The weighted contribution of all such composites is therefore

$$p\,M_p(N)=p\left(C_p\!\left(\left\lfloor\frac{N}{p}\right\rfloor\right)-C_p(p-1)\right).$$

There is no extra \(+p\) term here because the prime \(p\) itself remains in the filtered sum and is added later together with the other primes.

Step 4: Update the filtered counts and filtered sums

After the composites with smallest prime factor \(p\) have been identified, they must be removed from later stages. The count recurrence is

$$C_{p^+}(x)=C_p(x)-\left(C_p\!\left(\left\lfloor\frac{x}{p}\right\rfloor\right)-C_p(p-1)\right),$$

where \(p^+\) denotes the state immediately after processing \(p\). The corresponding sum recurrence removes the actual composite values \(p\,m\):

$$T_{p^+}(x)=T_p(x)-p\left(T_p\!\left(\left\lfloor\frac{x}{p}\right\rfloor\right)-T_p(p-1)\right).$$

Once every prime \(p\le \sqrt N\) has been processed, each composite number has been removed exactly once, so the surviving filtered sum is just the sum of all primes up to \(N\):

$$T_{\mathrm{final}}(N)=\sum_{\substack{q\le N\\ q\text{ prime}}} q.$$

Step 5: Compress all queried arguments to \(O(\sqrt N)\) values

Let

$$v=\left\lfloor\sqrt N\right\rfloor.$$

Every argument needed by the recurrence is either at most \(v\), or has the form \(\left\lfloor N/i\right\rfloor\) for some \(1\le i\le v\). This is the standard floor-quotient compression: the number of distinct values is only \(O(\sqrt N)\), not \(O(N)\). So the implementation stores two synchronized views,

$$C_p(x),\ T_p(x)\quad \text{for }1\le x\le v,$$

and

$$C_p\!\left(\left\lfloor\frac{N}{i}\right\rfloor\right),\ T_p\!\left(\left\lfloor\frac{N}{i}\right\rfloor\right)\quad \text{for }1\le i\le v.$$

That is why the memory usage stays near \(O(\sqrt N)\).

Worked Example: \(N=10\)

The smallest prime factors are

$$s(2),\dots,s(10)=2,3,2,5,2,7,2,3,2,$$

so the correct total is \(28\).

Initially, \(T_2(10)=2+3+\cdots+10=54\).

For \(p=2\),

$$M_2(10)=C_2(5)-C_2(1)=4-0=4,$$

corresponding to \(4,6,8,10\). Their weighted contribution is \(2\cdot 4=8\), and their actual sum \(4+6+8+10=28\) is removed from the filtered sum, leaving \(26\).

For \(p=3\), the current surviving set up to \(10\) is \(\{2,3,5,7,9\}\). Then

$$M_3(10)=C_3(3)-C_3(2)=2-1=1,$$

which corresponds to \(9\). Its weighted contribution is \(3\), and removing \(9\) leaves \(17\), exactly the sum of the primes \(2+3+5+7\).

Therefore

$$S(10)=17+8+3=28.$$

How the Code Works

The C++, Python, and Java implementations all begin with \(v=\lfloor\sqrt N\rfloor\) and build four tables: filtered counts and filtered sums for small arguments \(x\le v\), plus the same quantities for large arguments of the form \(\lfloor N/i\rfloor\). The initial values are the unfiltered counts \(x-1\) and the unfiltered sums \(\frac{x(x+1)}{2}-1\).

They then scan \(p=2,3,\dots,v\). A position is treated as prime exactly when the filtered count changes between \(p-1\) and \(p\). For each prime, the implementation first adds

$$p\left(C_p\!\left(\left\lfloor\frac{N}{p}\right\rfloor\right)-C_p(p-1)\right)$$

to the answer accumulator, accounting for every composite whose smallest prime factor is \(p\). It then applies the recurrence for \(C\) and \(T\) to every stored argument, selecting the small-table or large-table source according to whether the quotient lies below \(\sqrt N\).

The C++ and Python implementations keep exact integer sums during the whole process and reduce modulo \(10^9\) only at the end. The Java implementation keeps the sum tables modulo \(10^9\) during the updates while retaining exact counts; this is valid because only linear combinations of the sums are used, and the final answer itself is required modulo \(10^9\). After the prime loop finishes, the remaining filtered sum is the sum of all primes up to \(N\), so adding it to the accumulator yields \(S(N)\).

Complexity Analysis

Let \(v=\lfloor\sqrt N\rfloor\). The tables use \(O(v)=O(\sqrt N)\) memory. For each prime \(p\le v\), the implementation updates one range of length about \(\min\!\left(v,\left\lfloor N/p^2\right\rfloor\right)\) in the large-domain tables and one range of length about \(\max(0,v-p^2+1)\) in the small-domain tables. Summed over all primes, the work is roughly

$$O\!\left(\sum_{p\le v}\left(\min\!\left(v,\frac{N}{p^2}\right)+\max(0,v-p^2)\right)\right),$$

which is about \(O\!\left(N^{3/4}/\log N\right)\) for this split-domain sieve. The important point is practical rather than asymptotic finesse: both time and memory are far below any linear-in-\(N\) method, which is why \(N=10^{12}\) is manageable.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=521
  2. Prime factor: Wikipedia — Prime factor
  3. Prime-counting function: Wikipedia — Prime-counting function
  4. Meissel-Lehmer algorithm: Wikipedia — Meissel-Lehmer algorithm
  5. Floor and ceiling functions: Wikipedia — Floor and ceiling functions

Problem 521 source code

C++

#include <cassert>
#include <cmath>
#include <cstdint>
#include <iomanip>
#include <iostream>
#include <string>
#include <vector>

namespace {

using u64 = std::uint64_t;
using u128 = unsigned __int128;

constexpr u64 kMod = 1'000'000'000ULL;

u64 isqrt_u64(u64 n) {
    u64 r = static_cast<u64>(std::sqrt(static_cast<long double>(n)));
    while ((r + 1) * (r + 1) <= n) {
        ++r;
    }
    while (r * r > n) {
        --r;
    }
    return r;
}

u128 triangular_minus_one(u64 n) {
    return static_cast<u128>(n) * static_cast<u128>(n + 1ULL) / 2U - 1U;
}

u128 min_factor_sum(u64 n) {
    const u64 v = isqrt_u64(n);

    std::vector<u64> s_cnt(v + 1ULL, 0ULL);
    std::vector<u128> s_sum(v + 1ULL, 0ULL);
    std::vector<u64> l_cnt(v + 1ULL, 0ULL);
    std::vector<u128> l_sum(v + 1ULL, 0ULL);
    std::vector<std::uint8_t> used(v + 1ULL, 0U);

    for (u64 i = 0; i <= v; ++i) {
        s_cnt[static_cast<std::size_t>(i)] = (i == 0ULL) ? 0ULL : (i - 1ULL);
        s_sum[static_cast<std::size_t>(i)] = triangular_minus_one(i);
        if (i == 0ULL) {
            l_cnt[0] = 0ULL;
            l_sum[0] = 0U;
        } else {
            const u64 q = n / i;
            l_cnt[static_cast<std::size_t>(i)] = q - 1ULL;
            l_sum[static_cast<std::size_t>(i)] = triangular_minus_one(q);
        }
    }

    u128 ret = 0U;
    for (u64 p = 2ULL; p <= v; ++p) {
        const std::size_t ps = static_cast<std::size_t>(p);
        if (s_cnt[ps] == s_cnt[ps - 1ULL]) {
            continue;
        }

        const u64 p_cnt = s_cnt[ps - 1ULL];
        const u128 p_sum = s_sum[ps - 1ULL];
        const u64 q = p * p;

        ret += static_cast<u128>(p) * static_cast<u128>(l_cnt[ps] - p_cnt);

        l_cnt[1] -= (l_cnt[ps] - p_cnt);
        l_sum[1] -= (l_sum[ps] - p_sum) * static_cast<u128>(p);

        const u64 interval = (p & 1ULL) + 1ULL;
        const u64 end = std::min(v, n / q);
        for (u64 i = p + interval; i <= end; i += interval) {
            const std::size_t is = static_cast<std::size_t>(i);
            if (used[is] != 0U) {
                continue;
            }
            const u64 d = i * p;
            if (d <= v) {
                const std::size_t ds = static_cast<std::size_t>(d);
                l_cnt[is] -= (l_cnt[ds] - p_cnt);
                l_sum[is] -= (l_sum[ds] - p_sum) * static_cast<u128>(p);
            } else {
                const u64 t = n / d;
                const std::size_t ts = static_cast<std::size_t>(t);
                l_cnt[is] -= (s_cnt[ts] - p_cnt);
                l_sum[is] -= (s_sum[ts] - p_sum) * static_cast<u128>(p);
            }
        }

        if (q <= v) {
            const u64 step = p * interval;
            for (u64 i = q; i < end; i += step) {
                used[static_cast<std::size_t>(i)] = 1U;
            }
        }

        for (u64 i = v; i >= q; --i) {
            const std::size_t is = static_cast<std::size_t>(i);
            const std::size_t ts = static_cast<std::size_t>(i / p);
            s_cnt[is] -= (s_cnt[ts] - p_cnt);
            s_sum[is] -= (s_sum[ts] - p_sum) * static_cast<u128>(p);
        }
    }

    return l_sum[1] + ret;
}

u64 brute_small(int n) {
    std::vector<int> spf(static_cast<std::size_t>(n + 1), 0);
    for (int i = 2; i <= n; ++i) {
        if (spf[static_cast<std::size_t>(i)] == 0) {
            spf[static_cast<std::size_t>(i)] = i;
            if (static_cast<u64>(i) * static_cast<u64>(i) <= static_cast<u64>(n)) {
                for (int m = i * i; m <= n; m += i) {
                    if (spf[static_cast<std::size_t>(m)] == 0) {
                        spf[static_cast<std::size_t>(m)] = i;
                    }
                }
            }
        }
    }

    u64 sum = 0ULL;
    for (int i = 2; i <= n; ++i) {
        sum += static_cast<u64>(spf[static_cast<std::size_t>(i)]);
    }
    return sum;
}

u64 mod_u128(u128 x, u64 mod) {
    return static_cast<u64>(x % static_cast<u128>(mod));
}

bool run_checkpoints() {
    if (min_factor_sum(100ULL) != 1'257ULL) {
        std::cerr << "Checkpoint failed: S(100)\n";
        return false;
    }
    if (min_factor_sum(100ULL) != static_cast<u128>(brute_small(100))) {
        std::cerr << "Checkpoint failed: solve/brute mismatch at 100\n";
        return false;
    }
    if (min_factor_sum(1'000ULL) != static_cast<u128>(brute_small(1'000))) {
        std::cerr << "Checkpoint failed: solve/brute mismatch at 1000\n";
        return false;
    }
    return true;
}

}  // namespace

int main() {
    if (!run_checkpoints()) {
        return 1;
    }

    constexpr u64 n = 1'000'000'000'000ULL;
    const u128 full = min_factor_sum(n);
    const u64 answer = mod_u128(full, kMod);
    std::cout << std::setw(9) << std::setfill('0') << answer << '\n';
    return 0;
}

Python

import math

def solve():
    N = 10**12
    MOD = 10**9

    v = math.isqrt(N)

    s_cnt = list(range(v + 1))  # s_cnt[i] = i-1 for i>=1
    for i in range(v + 1):
        s_cnt[i] = i - 1 if i >= 1 else 0
    tri = lambda x: x * (x + 1) // 2 - 1 if x >= 1 else 0
    s_sum = [tri(i) for i in range(v + 1)]
    l_cnt = [0] * (v + 1)
    l_sum = [0] * (v + 1)
    for i in range(1, v + 1):
        q = N // i
        l_cnt[i] = q - 1
        l_sum[i] = q * (q + 1) // 2 - 1

    used = bytearray(v + 1)
    ret = 0

    for p in range(2, v + 1):
        if s_cnt[p] == s_cnt[p - 1]: continue
        p_cnt = s_cnt[p - 1]
        p_sum = s_sum[p - 1]
        q = p * p

        ret += p * (l_cnt[p] - p_cnt)
        l_cnt[1] -= (l_cnt[p] - p_cnt)
        l_sum[1] -= (l_sum[p] - p_sum) * p

        interval = (p & 1) + 1
        end = min(v, N // q)
        i = p + interval
        while i <= end:
            if not used[i]:
                d = i * p
                if d <= v:
                    l_cnt[i] -= (l_cnt[d] - p_cnt)
                    l_sum[i] -= (l_sum[d] - p_sum) * p
                else:
                    t = N // d
                    l_cnt[i] -= (s_cnt[t] - p_cnt)
                    l_sum[i] -= (s_sum[t] - p_sum) * p
            i += interval

        if q <= v:
            step = p * interval
            ii = q
            while ii < end:
                used[ii] = 1
                ii += step

        i = v
        while i >= q:
            t = i // p
            s_cnt[i] -= (s_cnt[t] - p_cnt)
            s_sum[i] -= (s_sum[t] - p_sum) * p
            i -= 1

    ans = (l_sum[1] + ret) % MOD
    return f"{ans:09d}"

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

Java

public class Euler521 {

    static final long kMod = 1000000000L;

    static long triangularMinusOneMod(long n) {
        // n*(n+1)/2 - 1 modulo 10^9
        long nMod = n % kMod;
        long nPlus1Mod = (n + 1) % kMod;

        // However, n*(n+1) must be divided by 2 before modulo.
        // We can just use exact math since n <= 10^12, n*(n+1) <= 10^24
        // We'll use BigInteger style approach but since it's just one op,
        // we can do:
        long resMod;
        if (n % 2 == 0) {
            long temp1 = (n / 2) % kMod;
            long temp2 = (n + 1) % kMod;
            resMod = (temp1 * temp2) % kMod;
        } else {
            long temp1 = n % kMod;
            long temp2 = ((n + 1) / 2) % kMod;
            resMod = (temp1 * temp2) % kMod;
        }
        resMod = (resMod - 1 + kMod) % kMod;
        return resMod;
    }

    public static void main(String[] args) {
        long N = 1000000000000L;
        int v = (int) Math.sqrt(N);

        long[] cnt_small = new long[v + 1];
        long[] cnt_large = new long[v + 1];

        long[] sum_small = new long[v + 1];
        long[] sum_large = new long[v + 1];

        for (int i = 1; i <= v; i++) {
            cnt_small[i] = i - 1;
            sum_small[i] = triangularMinusOneMod(i);

            long q = N / i;
            cnt_large[i] = q - 1;
            sum_large[i] = triangularMinusOneMod(q);
        }

        long ret = 0;

        for (int p = 2; p <= v; p++) {
            if (cnt_small[p] == cnt_small[p - 1])
                continue;

            long p_cnt = cnt_small[p - 1];
            long p_sum = sum_small[p - 1];
            long p2 = (long) p * p;

            long count_at_Np;
            if (N / p <= v) {
                count_at_Np = cnt_small[(int) (N / p)];
            } else {
                count_at_Np = cnt_large[p];
            }

            ret = (ret + (p % kMod) * ((count_at_Np - p_cnt) % kMod)) % kMod;

            int end_large = (int) Math.min(v, N / p2);

            for (int i = 1; i <= end_large; i++) {
                long d = (N / i) / p;
                long c, s;
                if (d <= v) {
                    c = cnt_small[(int) d];
                    s = sum_small[(int) d];
                } else {
                    c = cnt_large[(int) (N / d)];
                    s = sum_large[(int) (N / d)];
                }

                cnt_large[i] -= (c - p_cnt);
                long terms = (s - p_sum) % kMod;
                if (terms < 0)
                    terms += kMod;

                sum_large[i] = (sum_large[i] - terms * (p % kMod)) % kMod;
                if (sum_large[i] < 0)
                    sum_large[i] += kMod;
            }

            for (int i = v; i >= p2; i--) {
                int d = i / p;
                long c = cnt_small[d];
                long s = sum_small[d];

                cnt_small[i] -= (c - p_cnt);
                long terms = (s - p_sum) % kMod;
                if (terms < 0)
                    terms += kMod;

                sum_small[i] = (sum_small[i] - terms * (p % kMod)) % kMod;
                if (sum_small[i] < 0)
                    sum_small[i] += kMod;
            }
        }

        long ans = (sum_large[1] + ret) % kMod;
        if (ans < 0)
            ans += kMod;

        System.out.printf("%09d\n", ans);
    }
}