Problem 635: Subset Sums

View on Project Euler

Project Euler Problem 635 Solution

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

Problem Summary For a prime \(p\) and an integer \(q\), let \(A_q(p)\) denote the number of \(p\)-element subsets of \(\{1,2,\dots,qp\}\) whose element-sum is divisible by \(p\). Problem 635 asks for $$\sum_{p\le 10^8}\bigl(A_2(p)+A_3(p)\bigr)\pmod{10^9+9}.$$ A direct subset search is hopeless: even for a single large prime, the number of candidate subsets is astronomically large. The solution works because, for prime moduli, the counting problem collapses to compact closed forms involving binomial coefficients. Mathematical Approach The main task is to count \(p\)-element subsets whose sum is congruent to \(0 \pmod p\). A roots-of-unity filter extracts exactly those sums, and the residue pattern of \(\{1,\dots,qp\}\) makes the nontrivial filter terms identical. Step 1: Encode \(p\)-element subsets with a generating function Fix \(q\in\{2,3\}\) and a prime \(p\). Consider the bivariate generating function $$F_q(x,y)=\prod_{n=1}^{qp}(1+y x^n).$$ Choosing the term \(y x^n\) means that \(n\) is included in the subset; choosing \(1\) means it is omitted. Therefore the coefficient of \(y^k x^m\) counts \(k\)-element subsets with total sum \(m\). In particular, \(A_q(p)\) is obtained by taking \(k=p\) and summing only those coefficients whose exponent of \(x\) is a multiple of \(p\)....

Detailed mathematical approach

Problem Summary

For a prime \(p\) and an integer \(q\), let \(A_q(p)\) denote the number of \(p\)-element subsets of \(\{1,2,\dots,qp\}\) whose element-sum is divisible by \(p\). Problem 635 asks for

$$\sum_{p\le 10^8}\bigl(A_2(p)+A_3(p)\bigr)\pmod{10^9+9}.$$

A direct subset search is hopeless: even for a single large prime, the number of candidate subsets is astronomically large. The solution works because, for prime moduli, the counting problem collapses to compact closed forms involving binomial coefficients.

Mathematical Approach

The main task is to count \(p\)-element subsets whose sum is congruent to \(0 \pmod p\). A roots-of-unity filter extracts exactly those sums, and the residue pattern of \(\{1,\dots,qp\}\) makes the nontrivial filter terms identical.

Step 1: Encode \(p\)-element subsets with a generating function

Fix \(q\in\{2,3\}\) and a prime \(p\). Consider the bivariate generating function

$$F_q(x,y)=\prod_{n=1}^{qp}(1+y x^n).$$

Choosing the term \(y x^n\) means that \(n\) is included in the subset; choosing \(1\) means it is omitted. Therefore the coefficient of \(y^k x^m\) counts \(k\)-element subsets with total sum \(m\).

In particular, \(A_q(p)\) is obtained by taking \(k=p\) and summing only those coefficients whose exponent of \(x\) is a multiple of \(p\).

Step 2: Use a roots-of-unity filter to keep only sums divisible by \(p\)

Let \(\omega\) be a primitive \(p\)-th root of unity. The standard identity

$$\frac{1}{p}\sum_{j=0}^{p-1}\omega^{jm}= \begin{cases} 1,& p\mid m,\\ 0,& p\nmid m \end{cases}$$

filters out exactly the exponents \(m\) that are divisible by \(p\). Hence

$$A_q(p)=\frac{1}{p}\sum_{j=0}^{p-1}[y^p]\prod_{n=1}^{qp}(1+y\omega^{jn}).$$

Now group the numbers \(1,2,\dots,qp\) by their residues modulo \(p\). Each residue class appears exactly \(q\) times, so when \(j\neq 0\), multiplication by \(j\) merely permutes the residue classes and we get

$$\prod_{n=1}^{qp}(1+y\omega^{jn})=\left(\prod_{r=0}^{p-1}(1+y\omega^r)\right)^q.$$

Step 3: Collapse the nonzero filter terms

The cyclotomic product identity

$$\prod_{r=0}^{p-1}(1-t\omega^r)=1-t^p$$

gives, after substituting \(t=-y\),

$$\prod_{r=0}^{p-1}(1+y\omega^r)=1-(-y)^p.$$

Therefore every nonzero filter term contributes

$$[y^p](1-(-y)^p)^q.$$

If \(p\) is odd, then \(1-(-y)^p=1+y^p\), so the coefficient of \(y^p\) is simply \(q\). The \(j=0\) term is different: it is

$$[y^p](1+y)^{qp}=\binom{qp}{p}.$$

Step 4: Derive the closed forms for odd primes

For every odd prime \(p\), the \(p-1\) nonzero filter terms all contribute \(q\). Combining them with the \(j=0\) term yields

$$A_q(p)=\frac{\binom{qp}{p}+q(p-1)}{p}.$$

Specializing to the two quantities in the problem gives

$$A_2(p)=\frac{\binom{2p}{p}+2(p-1)}{p},$$

$$A_3(p)=\frac{\binom{3p}{p}+3(p-1)}{p}$$

for every odd prime \(p\).

Step 5: Handle the special case \(p=2\)

When \(p=2\), the same identity becomes

$$\prod_{r=0}^{1}(1+y\omega^r)=1-y^2,$$

so the nonzero filter term contributes \([y^2](1-y^2)^q=-q\) instead of \(+q\). Hence

$$A_q(2)=\frac{\binom{2q}{2}-q}{2}.$$

In particular,

$$A_2(2)=2,\qquad A_3(2)=6,$$

which explains why the implementation treats \(p=2\) separately.

Worked Example: \(p=5\)

For \(p=5\), both formulas apply directly:

$$A_2(5)=\frac{\binom{10}{5}+2\cdot 4}{5}=\frac{252+8}{5}=52,$$

$$A_3(5)=\frac{\binom{15}{5}+3\cdot 4}{5}=\frac{3003+12}{5}=603.$$

So the prime \(5\) contributes \(52+603=655\) to the final sum. As another useful checkpoint,

$$\sum_{p\le 10}A_2(p)=A_2(2)+A_2(3)+A_2(5)+A_2(7)=2+8+52+492=554,$$

which matches the internal verification used by the implementation.

How the Code Works

The C++, Python, and Java implementations first sieve all primes up to \(10^8\). They never store complete factorial and inverse-factorial tables up to \(3\cdot 10^8\); that would be unnecessarily large. Instead, they keep only a few checkpoint values for each prime.

A forward scan computes factorials modulo \(10^9+9\) and records the values needed later at \(p-1\), \(2p\), and \(3p\). Then the implementation computes the inverse of \((3\cdot 10^8)!\) once using fast modular exponentiation and performs a backward scan to reconstruct inverse factorials, recording only the checkpoints at \(p\) and \(2p\).

Those saved values are enough to recover

$$\frac{1}{p}=\frac{(p-1)!}{p!},\qquad \binom{2p}{p}=\frac{(2p)!}{p!\,p!},\qquad \binom{3p}{p}=\frac{(3p)!}{p!\,(2p)!}$$

in constant time for each prime. The modulus is prime and larger than \(3\cdot 10^8\), so every factorial value in that range is invertible modulo \(10^9+9\). After reconstructing the binomial terms, the program applies the \(p=2\) special case or the odd-prime formula, adds the two contributions, and reduces modulo \(10^9+9\) throughout.

Complexity Analysis

Let \(L=10^8\). The prime sieve costs \(O(L\log\log L)\) time. The forward factorial scan and the backward inverse-factorial scan both run to \(3L\), so they add \(O(L)\) time. Once the checkpoint tables are ready, each prime is processed in \(O(1)\) time. The overall running time is therefore \(O(L\log\log L)\).

Memory usage consists of an odd-only sieve plus a constant number of checkpoint arrays indexed by primes. Asymptotically this is \(O(L+\pi(L))\), and the important practical point is that the implementation avoids storing all factorials and inverse factorials up to \(3L\).

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=635
  2. Root-of-unity filter: Wikipedia - Root of unity filter
  3. Binomial coefficient: Wikipedia - Binomial coefficient
  4. Fermat's little theorem: Wikipedia - Fermat's little theorem
  5. Sieve of Eratosthenes: Wikipedia - Sieve of Eratosthenes

Problem 635 source code

C++

#include <algorithm>
#include <cassert>
#include <cstdint>
#include <iostream>
#include <vector>

using i64 = long long;

static constexpr int MOD = 1'000'000'009;

static inline int mod_add(int a, int b) {
    int s = a + b;
    if (s >= MOD) s -= MOD;
    return s;
}

static inline int mod_mul(i64 a, i64 b) { return (int)(a * b % MOD); }

static int mod_pow(int a, int e) {
    i64 r = 1, x = a;
    while (e > 0) {
        if (e & 1) r = (r * x) % MOD;
        x = (x * x) % MOD;
        e >>= 1;
    }
    return (int)r;
}

static std::vector<int> primes_upto(int n) {
    std::vector<bool> comp(n / 2 + 1, false);
    std::vector<int> primes;
    primes.reserve(6'000'000);
    primes.push_back(2);
    for (int i = 3; (i64)i * i <= n; i += 2) {
        if (comp[i / 2]) continue;
        for (int j = i * i; j <= n; j += i * 2) comp[j / 2] = true;
    }
    for (int i = 3; i <= n; i += 2) {
        if (!comp[i / 2]) primes.push_back(i);
    }
    return primes;
}

int main() {
    static constexpr int L = 100'000'000;
    static constexpr int NMAX = 3 * L;

    const auto primes = primes_upto(L);
    const int m = (int)primes.size();

    std::vector<int> f_pminus1(m), f_2p(m), f_3p(m), if_p(m), if_2p(m);

    int fact = 1;
    int ip1 = 0, i2 = 0, i3 = 0;
    for (int i = 1; i <= NMAX; ++i) {
        fact = mod_mul(fact, i);
        while (ip1 < m && primes[ip1] - 1 == i) f_pminus1[ip1++] = fact;
        while (i2 < m && 2 * primes[i2] == i) f_2p[i2++] = fact;
        while (i3 < m && 3 * primes[i3] == i) f_3p[i3++] = fact;
    }
    assert(ip1 == m && i2 == m && i3 == m);

    int invfact = mod_pow(fact, MOD - 2);
    int jp = m - 1, j2 = m - 1;
    for (int i = NMAX; i >= 1; --i) {
        while (jp >= 0 && primes[jp] == i) if_p[jp--] = invfact;
        while (j2 >= 0 && 2 * primes[j2] == i) if_2p[j2--] = invfact;
        invfact = mod_mul(invfact, i);
    }
    assert(jp == -1 && j2 == -1);

    int s2 = 0, s3 = 0;
    int s2_10 = 0, s2_100 = 0, s3_100 = 0;
    for (int idx = 0; idx < m; ++idx) {
        const int p = primes[idx];

        int a2 = 0, a3 = 0;
        if (p == 2) {
            a2 = 2;
            a3 = 6;
        } else {
            const int inv_p = mod_mul(f_pminus1[idx], if_p[idx]);  // 1/p
            const int bin2 = mod_mul(f_2p[idx], mod_mul(if_p[idx], if_p[idx]));
            const int bin3 = mod_mul(f_3p[idx], mod_mul(if_p[idx], if_2p[idx]));

            a2 = mod_mul(mod_add(bin2, mod_mul(2, p - 1)), inv_p);
            a3 = mod_mul(mod_add(bin3, mod_mul(3, p - 1)), inv_p);
        }

        s2 = mod_add(s2, a2);
        s3 = mod_add(s3, a3);

        if (p <= 10) s2_10 = mod_add(s2_10, a2);
        if (p <= 100) {
            s2_100 = mod_add(s2_100, a2);
            s3_100 = mod_add(s3_100, a3);
        }
    }

    assert(s2_10 == 554);
    assert(s2_100 == 100433628);
    assert(s3_100 == 855618282);

    std::cout << mod_add(s2, s3) << "\n";
    return 0;
}

Python

def solve():
    MOD = 1000000009
    L = 100000000
    NMAX = 3 * L

    def mod_pow(a, e):
        r, x = 1, a
        while e > 0:
            if e & 1: r = r * x % MOD
            x = x * x % MOD; e >>= 1
        return r

    # Sieve primes up to L
    half = L // 2 + 1
    comp = bytearray(half)
    primes = [2]
    for i in range(3, int(L**0.5)+1, 2):
        if not comp[i//2]:
            for j in range(i*i, L+1, 2*i): comp[j//2] = 1
    for i in range(3, L+1, 2):
        if not comp[i//2]: primes.append(i)

    m = len(primes)
    # Compute factorials at key points
    f_pm1 = [0]*m; f_2p = [0]*m; f_3p = [0]*m; if_p = [0]*m; if_2p = [0]*m

    fact = 1; ip1 = i2 = i3 = 0
    for i in range(1, NMAX + 1):
        fact = fact * i % MOD
        while ip1 < m and primes[ip1]-1 == i: f_pm1[ip1] = fact; ip1 += 1
        while i2 < m and 2*primes[i2] == i: f_2p[i2] = fact; i2 += 1
        while i3 < m and 3*primes[i3] == i: f_3p[i3] = fact; i3 += 1

    invfact = mod_pow(fact, MOD-2)
    jp = m-1; j2 = m-1
    for i in range(NMAX, 0, -1):
        while jp >= 0 and primes[jp] == i: if_p[jp] = invfact; jp -= 1
        while j2 >= 0 and 2*primes[j2] == i: if_2p[j2] = invfact; j2 -= 1
        invfact = invfact * i % MOD

    s2 = s3 = 0
    for idx in range(m):
        p = primes[idx]
        if p == 2: a2, a3 = 2, 6
        else:
            inv_p = f_pm1[idx] * if_p[idx] % MOD
            bin2 = f_2p[idx] * if_p[idx] % MOD * if_p[idx] % MOD
            bin3 = f_3p[idx] * if_p[idx] % MOD * if_2p[idx] % MOD
            a2 = (bin2 + 2*(p-1)) % MOD * inv_p % MOD
            a3 = (bin3 + 3*(p-1)) % MOD * inv_p % MOD
        s2 = (s2 + a2) % MOD; s3 = (s3 + a3) % MOD

    return str((s2 + s3) % MOD)

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

Java

import java.util.ArrayList;

public class Euler635 {
    static final int MOD = 1000000009;

    static int mod_add(int a, int b) {
        int s = a + b;
        if (s >= MOD)
            s -= MOD;
        return s;
    }

    static int mod_mul(long a, long b) {
        return (int) (a * b % MOD);
    }

    static int mod_pow(int a, int e) {
        long r = 1, x = a;
        while (e > 0) {
            if ((e & 1) != 0)
                r = (r * x) % MOD;
            x = (x * x) % MOD;
            e >>= 1;
        }
        return (int) r;
    }

    static ArrayList<Integer> primes_upto(int n) {
        byte[] comp = new byte[n / 2 + 1];
        ArrayList<Integer> primes = new ArrayList<>();
        primes.add(2);
        for (int i = 3; (long) i * i <= n; i += 2) {
            if (comp[i / 2] == 0) {
                for (int j = i * i; j <= n; j += i * 2)
                    comp[j / 2] = 1;
            }
        }
        for (int i = 3; i <= n; i += 2) {
            if (comp[i / 2] == 0)
                primes.add(i);
        }
        return primes;
    }

    public static String solve() {
        int L = 100000000;
        int NMAX = 3 * L;

        ArrayList<Integer> primes = primes_upto(L);
        int m = primes.size();

        int[] f_pminus1 = new int[m];
        int[] f_2p = new int[m];
        int[] f_3p = new int[m];
        int[] if_p = new int[m];
        int[] if_2p = new int[m];

        int fact = 1;
        int ip1 = 0, i2 = 0, i3 = 0;
        for (int i = 1; i <= NMAX; ++i) {
            fact = mod_mul(fact, i);
            while (ip1 < m && primes.get(ip1) - 1 == i)
                f_pminus1[ip1++] = fact;
            while (i2 < m && 2 * primes.get(i2) == i)
                f_2p[i2++] = fact;
            while (i3 < m && 3 * primes.get(i3) == i)
                f_3p[i3++] = fact;
        }

        int invfact = mod_pow(fact, MOD - 2);
        int jp = m - 1, j2 = m - 1;
        for (int i = NMAX; i >= 1; --i) {
            while (jp >= 0 && primes.get(jp) == i)
                if_p[jp--] = invfact;
            while (j2 >= 0 && 2 * primes.get(j2) == i)
                if_2p[j2--] = invfact;
            invfact = mod_mul(invfact, i);
        }

        int s2 = 0, s3 = 0;
        for (int idx = 0; idx < m; ++idx) {
            int p = primes.get(idx);
            int a2 = 0, a3 = 0;
            if (p == 2) {
                a2 = 2;
                a3 = 6;
            } else {
                int inv_p = mod_mul(f_pminus1[idx], if_p[idx]);
                int bin2 = mod_mul(f_2p[idx], mod_mul(if_p[idx], if_p[idx]));
                int bin3 = mod_mul(f_3p[idx], mod_mul(if_p[idx], if_2p[idx]));

                a2 = mod_mul(mod_add(bin2, mod_mul(2, p - 1)), inv_p);
                a3 = mod_mul(mod_add(bin3, mod_mul(3, p - 1)), inv_p);
            }
            s2 = mod_add(s2, a2);
            s3 = mod_add(s3, a3);
        }

        return Integer.toString(mod_add(s2, s3));
    }

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