Problem 304: Primonacci

View on Project Euler

Project Euler Problem 304 Solution

EulerSolve provides an optimized solution for Project Euler Problem 304, Primonacci, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary Let \(p_1,p_2,\dots,p_m\) be the first \(m\) primes strictly larger than a huge starting value \(A\). The task is to compute $$\sum_{i=1}^{m}F_{p_i}\pmod{M},$$ where \(F_n\) is the Fibonacci sequence and the official parameters are extremely large. Mathematical Approach 1) Why prime generation must be segmented The primes we need live near \(A\approx 10^{14}\). A naive primality test for each number in that region would be far too slow, and a full sieve up to \(A\) is impossible in memory. The right tool is a segmented sieve. We scan blocks $$[L,H],$$ and mark composites inside that block using all base primes up to \(\sqrt{H}\). 2) Why primes up to \(\sqrt{H}\) are enough If \(x\in [L,H]\) is composite, then \(x=ab\) with \(1<a\le b\). Hence $$a\le \sqrt{x}\le \sqrt{H}.$$ So every composite in the segment has a prime divisor not exceeding \(\sqrt{H}\). Therefore, after marking multiples of all primes \(p\le \sqrt{H}\), every remaining unmarked number in the segment is prime. 3) Where marking starts inside a segment For a fixed base prime \(p\), the first multiple of \(p\) inside the current segment is $$\left\lceil\frac{L}{p}\right\rceil p.$$ But we must also avoid marking the prime \(p\) itself when \(p\in[L,H]\)....

Detailed mathematical approach

Problem Summary

Let \(p_1,p_2,\dots,p_m\) be the first \(m\) primes strictly larger than a huge starting value \(A\). The task is to compute

$$\sum_{i=1}^{m}F_{p_i}\pmod{M},$$

where \(F_n\) is the Fibonacci sequence and the official parameters are extremely large.

Mathematical Approach

1) Why prime generation must be segmented

The primes we need live near \(A\approx 10^{14}\). A naive primality test for each number in that region would be far too slow, and a full sieve up to \(A\) is impossible in memory. The right tool is a segmented sieve.

We scan blocks

$$[L,H],$$

and mark composites inside that block using all base primes up to \(\sqrt{H}\).

2) Why primes up to \(\sqrt{H}\) are enough

If \(x\in [L,H]\) is composite, then \(x=ab\) with \(1<a\le b\). Hence

$$a\le \sqrt{x}\le \sqrt{H}.$$

So every composite in the segment has a prime divisor not exceeding \(\sqrt{H}\). Therefore, after marking multiples of all primes \(p\le \sqrt{H}\), every remaining unmarked number in the segment is prime.

3) Where marking starts inside a segment

For a fixed base prime \(p\), the first multiple of \(p\) inside the current segment is

$$\left\lceil\frac{L}{p}\right\rceil p.$$

But we must also avoid marking the prime \(p\) itself when \(p\in[L,H]\). Therefore the correct starting point is

$$\max\!\left(p^2,\left\lceil\frac{L}{p}\right\rceil p\right).$$

This is exactly what the code computes before it walks through the segment in steps of \(p\).

4) Fibonacci values via fast doubling

Once the prime indices \(p_i\) are known, we still cannot generate Fibonacci numbers linearly up to index \(10^{14}\). The code instead computes \((F_n,F_{n+1})\) with fast doubling.

Using the addition formulas

$$F_{m+n}=F_mF_{n+1}+F_{m-1}F_n,$$

and substituting \(m=n=k\), one obtains the standard doubling identities

$$F_{2k}=F_k(2F_{k+1}-F_k),$$

$$F_{2k+1}=F_k^2+F_{k+1}^2.$$

Thus from \((F_k,F_{k+1})\) we can compute \((F_{2k},F_{2k+1})\), and then decide whether \(n\) is even or odd. Each recursive step halves the index, so the running time is \(O(\log n)\).

5) Why modular arithmetic can be applied early

All arithmetic is needed only modulo \(M\). Since Fibonacci recurrences use only addition, subtraction, and multiplication, we may reduce after every operation:

$$F_n \bmod M$$

can be computed without ever forming the enormous exact integer \(F_n\). This is why the solution remains fast even though the indices are gigantic.

6) Small worked checkpoint

The C++ code contains a cross-check on a much smaller range:

$$A=1000,\qquad m=30,\qquad M=10^9+7.$$

The first few primes after \(1000\) are

$$1009,\ 1013,\ 1019,\ 1021,\ 1031,\dots$$

and the corresponding Fibonacci residues begin as

$$F_{1009}\equiv 529241575,\qquad F_{1013}\equiv 613716502,\qquad F_{1019}\equiv 506824938\pmod{10^9+7}.$$

For the first \(30\) such primes, the checkpoint sum is

$$\sum_{i=1}^{30}F_{p_i}\equiv 682050181\pmod{10^9+7}.$$

The program compares its segmented-sieve result with a brute-force prime search on that small input.

7) Safe preprocessing bounds

Because the main search scans segments whose upper endpoint is roughly \(A+\text{segment\_size}\), it is enough to precompute base primes up to a safe constant slightly larger than

$$\sqrt{A+\text{segment\_size}}.$$

For \(A\approx 10^{14}\), this is about \(10^7\), which explains the base-prime sieve limit used in the implementations.

How the Code Works

The solver has three clean stages:

1. generate base primes with an ordinary sieve up to a safe square-root bound;

2. run a segmented sieve to extract the first \(m\) primes after \(A\);

3. compute each \(F_{p_i}\bmod M\) by fast doubling and accumulate the sum modulo \(M\).

The helper fib_pair_mod(n, mod) returns \((F_n,F_{n+1})\), so each prime index is handled independently in logarithmic time.

Complexity Analysis

The segmented sieve is essentially linear in the total length of the scanned blocks, up to the usual harmonic marking cost from small primes. The Fibonacci part costs

$$O(m\log A),$$

because each of the \(m\) prime indices needs one logarithmic fast-doubling evaluation. Memory usage is modest: one segment bitmap plus the base-prime list.

Further Reading

  1. Problem page: https://projecteuler.net/problem=304
  2. Segmented sieve: https://en.wikipedia.org/wiki/Sieve_of_Eratosthenes#Segmented_sieve
  3. Fibonacci identities and fast doubling: https://cp-algorithms.com/algebra/fibonacci-numbers.html

Problem 304 source code

C++

#include <cmath>
#include <cstdint>
#include <iostream>
#include <string>
#include <utility>
#include <vector>

namespace {

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

struct Options {
    u64 start = 100000000000000ULL;
    int count = 100000;
    u64 modulo = 1234567891011ULL;
    bool run_checkpoints = true;
};

bool parse_u64_after_prefix(const std::string& arg, const std::string& prefix, u64& value) {
    if (arg.rfind(prefix, 0U) != 0U) {
        return false;
    }
    const std::string tail = arg.substr(prefix.size());
    if (tail.empty()) {
        return false;
    }
    u64 parsed = 0ULL;
    for (char c : tail) {
        if (c < '0' || c > '9') {
            return false;
        }
        parsed = parsed * 10ULL + static_cast<u64>(c - '0');
    }
    value = parsed;
    return true;
}

bool parse_int_after_prefix(const std::string& arg, const std::string& prefix, int& value) {
    if (arg.rfind(prefix, 0U) != 0U) {
        return false;
    }
    const std::string tail = arg.substr(prefix.size());
    if (tail.empty()) {
        return false;
    }
    int parsed = 0;
    for (char c : tail) {
        if (c < '0' || c > '9') {
            return false;
        }
        parsed = parsed * 10 + static_cast<int>(c - '0');
    }
    value = parsed;
    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_u64_after_prefix(arg, "--start=", options.start) ||
            parse_int_after_prefix(arg, "--count=", options.count) ||
            parse_u64_after_prefix(arg, "--modulo=", options.modulo)) {
            continue;
        }
        std::cerr << "Unknown argument: " << arg << '\n';
        return false;
    }
    return options.start >= 2ULL && options.count >= 1 && options.modulo >= 2ULL;
}

std::vector<int> sieve_base_primes(const int limit) {
    std::vector<std::uint8_t> is_prime(static_cast<std::size_t>(limit + 1), 1U);
    is_prime[0] = 0U;
    is_prime[1] = 0U;
    for (int p = 2; static_cast<u64>(p) * static_cast<u64>(p) <= static_cast<u64>(limit); ++p) {
        if (is_prime[static_cast<std::size_t>(p)] == 0U) {
            continue;
        }
        for (int q = p * p; q <= limit; q += p) {
            is_prime[static_cast<std::size_t>(q)] = 0U;
        }
    }
    std::vector<int> primes;
    for (int p = 2; p <= limit; ++p) {
        if (is_prime[static_cast<std::size_t>(p)] != 0U) {
            primes.push_back(p);
        }
    }
    return primes;
}

std::vector<u64> first_primes_after(u64 start, const int count) {
    const int base_limit = 20000000;  // sqrt(1e14 + margin) is around 1e7.
    const std::vector<int> base_primes = sieve_base_primes(base_limit);

    const u64 segment_size = 4000000ULL;
    std::vector<u64> out;
    out.reserve(static_cast<std::size_t>(count));

    u64 low = start + 1ULL;
    while (static_cast<int>(out.size()) < count) {
        const u64 high = low + segment_size - 1ULL;
        std::vector<std::uint8_t> is_prime(static_cast<std::size_t>(segment_size), 1U);

        for (int p : base_primes) {
            const u64 pp = static_cast<u64>(p);
            if (pp * pp > high) {
                break;
            }
            u64 first = (low + pp - 1ULL) / pp * pp;
            if (first < pp * pp) {
                first = pp * pp;
            }
            for (u64 x = first; x <= high; x += pp) {
                is_prime[static_cast<std::size_t>(x - low)] = 0U;
            }
        }

        for (u64 i = 0; i < segment_size && static_cast<int>(out.size()) < count; ++i) {
            const u64 value = low + i;
            if (value >= 2ULL && is_prime[static_cast<std::size_t>(i)] != 0U) {
                out.push_back(value);
            }
        }
        low = high + 1ULL;
    }
    return out;
}

std::pair<u64, u64> fib_pair_mod(const u64 n, const u64 mod) {
    if (n == 0ULL) {
        return {0ULL, 1ULL % mod};
    }
    const auto [a, b] = fib_pair_mod(n >> 1U, mod);
    const u64 two_b = static_cast<u64>((2ULL * b) % mod);
    const u64 c = static_cast<u64>((static_cast<u128>(a) * ((two_b + mod - a) % mod)) % mod);
    const u64 d = static_cast<u64>((static_cast<u128>(a) * a + static_cast<u128>(b) * b) % mod);
    if ((n & 1ULL) == 0ULL) {
        return {c, d};
    }
    return {d, static_cast<u64>((c + d) % mod)};
}

u64 solve(const u64 start, const int count, const u64 mod) {
    const std::vector<u64> primes = first_primes_after(start, count);
    u64 sum = 0ULL;
    for (u64 p : primes) {
        sum += fib_pair_mod(p, mod).first;
        if (sum >= mod) {
            sum %= mod;
        }
    }
    return sum % mod;
}

u64 solve_bruteforce(const u64 start, const int count, const u64 mod) {
    auto is_prime_trial = [](u64 x) {
        if (x < 2ULL) {
            return false;
        }
        if ((x & 1ULL) == 0ULL) {
            return x == 2ULL;
        }
        for (u64 d = 3ULL; d * d <= x; d += 2ULL) {
            if (x % d == 0ULL) {
                return false;
            }
        }
        return true;
    };

    std::vector<u64> primes;
    for (u64 x = start + 1ULL; static_cast<int>(primes.size()) < count; ++x) {
        if (is_prime_trial(x)) {
            primes.push_back(x);
        }
    }
    u64 sum = 0ULL;
    for (u64 p : primes) {
        sum = (sum + fib_pair_mod(p, mod).first) % mod;
    }
    return sum;
}

bool run_checkpoints() {
    if (fib_pair_mod(10ULL, 1000000007ULL).first != 55ULL) {
        std::cerr << "Checkpoint failed for Fibonacci doubling F(10)" << '\n';
        return false;
    }
    if (solve(1000ULL, 30, 1000000007ULL) != solve_bruteforce(1000ULL, 30, 1000000007ULL)) {
        std::cerr << "Checkpoint failed for brute cross-check on small prime range" << '\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 << solve(options.start, options.count, options.modulo) << '\n';
    return 0;
}

Python

def sieve_base_primes(limit):
    is_prime = bytearray(limit + 1)
    for i in range(2, limit + 1):
        is_prime[i] = 1
        
    p = 2
    while p * p <= limit:
        if is_prime[p]:
            for q in range(p * p, limit + 1, p):
                is_prime[q] = 0
        p += 1
        
    primes = [p for p in range(2, limit + 1) if is_prime[p]]
    return primes

def first_primes_after(start, count):
    base_limit = 10000000 + 1000000  # sqrt(10^14) is 10^7
    base_primes = sieve_base_primes(base_limit)
    
    segment_size = 4000000
    out = []
    
    low = start + 1
    while len(out) < count:
        high = low + segment_size - 1
        is_prime = bytearray(segment_size)
        for i in range(segment_size):
            is_prime[i] = 1
            
        for p in base_primes:
            if p * p > high:
                break
                
            first = ((low + p - 1) // p) * p
            if first < p * p:
                first = p * p
                
            start_idx = first - low
            for x in range(start_idx, segment_size, p):
                is_prime[x] = 0
                
        for i in range(segment_size):
            if len(out) >= count:
                break
            if is_prime[i]:
                value = low + i
                if value >= 2:
                    out.append(value)
                    
        low = high + 1
        
    return out

def fib_pair_mod(n, mod):
    if n == 0:
        return (0, 1 % mod)
        
    a, b = fib_pair_mod(n >> 1, mod)
    two_b = (2 * b) % mod
    c = (a * ((two_b + mod - a) % mod)) % mod
    d = (a * a + b * b) % mod
    
    if (n & 1) == 0:
        return (c, d)
    return (d, (c + d) % mod)

def solve(start=100000000000000, count=100000, mod=1234567891011):
    primes = first_primes_after(start, count)
    total_sum = 0
    for p in primes:
        total_sum += fib_pair_mod(p, mod)[0]
        if total_sum >= mod:
            total_sum %= mod
            
    return str(total_sum % mod)

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

Java

import java.util.ArrayList;
import java.util.List;

public class Euler304 {
    static List<Integer> sieveBasePrimes(int limit) {
        byte[] isPrime = new byte[limit + 1];
        for (int i = 2; i <= limit; ++i)
            isPrime[i] = 1;

        for (int p = 2; (long) p * p <= limit; ++p) {
            if (isPrime[p] == 1) {
                for (int q = p * p; q <= limit; q += p) {
                    isPrime[q] = 0;
                }
            }
        }

        List<Integer> primes = new ArrayList<>();
        for (int p = 2; p <= limit; ++p) {
            if (isPrime[p] == 1)
                primes.add(p);
        }
        return primes;
    }

    static List<Long> firstPrimesAfter(long start, int count) {
        int baseLimit = 11000000;
        List<Integer> basePrimes = sieveBasePrimes(baseLimit);

        int segmentSize = 4000000;
        List<Long> out = new ArrayList<>(count);

        long low = start + 1;
        byte[] isPrime = new byte[segmentSize];

        while (out.size() < count) {
            long high = low + segmentSize - 1;
            for (int i = 0; i < segmentSize; ++i)
                isPrime[i] = 1;

            for (int p : basePrimes) {
                long pp = (long) p;
                if (pp * pp > high)
                    break;

                long first = ((low + pp - 1) / pp) * pp;
                if (first < pp * pp) {
                    first = pp * pp;
                }

                int startIdx = (int) (first - low);
                for (int x = startIdx; x < segmentSize; x += p) {
                    isPrime[x] = 0;
                }
            }

            for (int i = 0; i < segmentSize && out.size() < count; ++i) {
                long value = low + i;
                if (value >= 2 && isPrime[i] == 1) {
                    out.add(value);
                }
            }

            low = high + 1;
        }

        return out;
    }

    static class Pair {
        long first, second;

        Pair(long first, long second) {
            this.first = first;
            this.second = second;
        }
    }

    static Pair fibPairMod(long n, long mod) {
        if (n == 0)
            return new Pair(0, 1 % mod);

        Pair half = fibPairMod(n >>> 1, mod);
        long a = half.first;
        long b = half.second;

        long twoB = (2L * b) % mod;

        long c = mulMod(a, (twoB + mod - a) % mod, mod);
        long d = (mulMod(a, a, mod) + mulMod(b, b, mod)) % mod;

        if ((n & 1) == 0) {
            return new Pair(c, d);
        }
        return new Pair(d, (c + d) % mod);
    }

    // multiply modulo handling overflow
    static long mulMod(long a, long b, long mod) {
        long res = 0;
        a %= mod;
        while (b > 0) {
            if ((b & 1) == 1) {
                res = (res + a) % mod;
            }
            a = (a * 2) % mod;
            b >>>= 1;
        }
        return res;
    }

    public static String solve() {
        long start = 100000000000000L;
        int count = 100000;
        long mod = 1234567891011L;

        List<Long> primes = firstPrimesAfter(start, count);
        long sum = 0;
        for (long p : primes) {
            sum += fibPairMod(p, mod).first;
            if (sum >= mod)
                sum %= mod;
        }

        return String.valueOf(sum % mod);
    }

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