Problem 900: DistribuNim II

View on Project Euler

Project Euler Problem 900 Solution

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

Problem Summary For each positive integer \(n\), let \(t(n)\) be the smallest offset \(k\) such that the DistribuNim II position consisting of \(n\) piles of size \(n\) and one additional pile of size \(n+k\) is a losing position. The problem asks for $$S(N)=\sum_{n=1}^{2^N} t(n)\pmod{900497239}.$$ A direct search over game states is hopeless for large \(N\). The efficient solution therefore avoids game-tree exploration and instead uses a closed formula for each \(t(n)\), followed by an exact summation over bit-length blocks. Mathematical Approach The key fact used by the optimized solution is that the losing offset is controlled by the low bits of \(n^2\). Step 1: Closed Form for \(t(n)\) Let $$m=\lfloor \log_2 n \rfloor+1,\qquad 2^{m-1}\le n<2^m.$$ Define $$\varepsilon(n)=\begin{cases} 0, & n\text{ even},\\ 2^m-1, & n\text{ odd}. \end{cases}$$ Then the implementations use the identity $$t(n)=\bigl(\varepsilon(n)-n^2\bigr)\bmod 2^m.$$ Equivalently, $$t(n)=\begin{cases} (-n^2)\bmod 2^m, & n\text{ even},\\ (-n^2-1)\bmod 2^m, & n\text{ odd}. \end{cases}$$ So once the bit-length \(m\) is known, we only need the lowest \(m\) bits of \(n^2\). Step 2: Sum by Fixed Bit-Length For each \(m\ge 1\), define the block sum $$B_m=\sum_{n=2^{m-1}}^{2^m-1} t(n).$$ If \(n=2^r\), then \(n\) is even and \(n^2=2^{2r}\) is divisible by \(2^{r+1}\), so \(t(2^r)=0\)....

Detailed mathematical approach

Problem Summary

For each positive integer \(n\), let \(t(n)\) be the smallest offset \(k\) such that the DistribuNim II position consisting of \(n\) piles of size \(n\) and one additional pile of size \(n+k\) is a losing position. The problem asks for

$$S(N)=\sum_{n=1}^{2^N} t(n)\pmod{900497239}.$$

A direct search over game states is hopeless for large \(N\). The efficient solution therefore avoids game-tree exploration and instead uses a closed formula for each \(t(n)\), followed by an exact summation over bit-length blocks.

Mathematical Approach

The key fact used by the optimized solution is that the losing offset is controlled by the low bits of \(n^2\).

Step 1: Closed Form for \(t(n)\)

Let

$$m=\lfloor \log_2 n \rfloor+1,\qquad 2^{m-1}\le n<2^m.$$

Define

$$\varepsilon(n)=\begin{cases} 0, & n\text{ even},\\ 2^m-1, & n\text{ odd}. \end{cases}$$

Then the implementations use the identity

$$t(n)=\bigl(\varepsilon(n)-n^2\bigr)\bmod 2^m.$$

Equivalently,

$$t(n)=\begin{cases} (-n^2)\bmod 2^m, & n\text{ even},\\ (-n^2-1)\bmod 2^m, & n\text{ odd}. \end{cases}$$

So once the bit-length \(m\) is known, we only need the lowest \(m\) bits of \(n^2\).

Step 2: Sum by Fixed Bit-Length

For each \(m\ge 1\), define the block sum

$$B_m=\sum_{n=2^{m-1}}^{2^m-1} t(n).$$

If \(n=2^r\), then \(n\) is even and \(n^2=2^{2r}\) is divisible by \(2^{r+1}\), so \(t(2^r)=0\). Therefore the endpoint \(n=2^N\) contributes nothing, and

$$S(N)=\sum_{m=1}^{N} B_m.$$

Now write

$$n=2^{m-1}+j,\qquad 0\le j<2^{m-1}.$$

Then

$$n^2=(2^{m-1}+j)^2\equiv j^2 \pmod{2^m},$$

because both \(2^{2m-2}\) and \(2^m j\) are multiples of \(2^m\). For \(m\ge 2\), the number \(2^{m-1}\) is even, so \(n\) and \(j\) have the same parity. Hence

$$t(2^{m-1}+j)=\begin{cases} (-j^2)\bmod 2^m, & j\text{ even},\\ (-j^2-1)\bmod 2^m, & j\text{ odd}. \end{cases}$$

Step 3: Evaluate One Whole Block

Split the block into even and odd offsets:

$$j=2u\quad\text{or}\quad j=2u+1,\qquad 0\le u<2^{m-2}.$$

Then the even contribution is built from residues of \(4u^2\), while the odd contribution is built from residues of \(4u(u+1)+1\). Summing those two parity classes over a complete block gives exact formulas

$$B_{2k}=4^{2k-1}+2\cdot 8^{k-1}-2^{2k},\qquad k\ge 1,$$

$$B_{2k+1}=4^{2k}+8^k-2^{2k+1},\qquad k\ge 0.$$

The first few values are

$$B_1=0,\qquad B_2=2,\qquad B_3=16,\qquad B_4=64,\qquad B_5=288,$$

which already shows the geometric structure that the code exploits.

Step 4: Sum the Block Formulas

Every block contributes a \(4^{m-1}\) term, so

$$\sum_{m=1}^{N} 4^{m-1}=\frac{4^N-1}{3}.$$

The \(8\)-power corrections come in adjacent pairs. Block \(2k+1\) contributes \(8^k\), and block \(2k+2\) contributes \(2\cdot 8^k\), so one complete pair contributes \(3\cdot 8^k\). If

$$K=\left\lfloor \frac{N}{2}\right\rfloor,$$

then the total \(8\)-power contribution is

$$\frac{3(8^K-1)}{7}+\mathbf{1}_{N\text{ odd}}\,8^K.$$

The subtractive pieces add up to

$$\sum_{m=1}^{N}2^m=2^{N+1}-2.$$

Therefore

$$\boxed{S(N)=\frac{4^N-1}{3}+\frac{3(8^{\lfloor N/2\rfloor}-1)}{7}+\mathbf{1}_{N\text{ odd}}\,8^{\lfloor N/2\rfloor}-(2^{N+1}-2)\pmod{900497239}.}$$

Step 5: Worked Example

Take \(n=6\). Here \(m=3\) because \(4\le 6<8\). Since \(6\) is even,

$$t(6)=(-6^2)\bmod 8=(-36)\bmod 8=4.$$

Take \(n=5\). Again \(m=3\), but now \(5\) is odd, so

$$t(5)=(-5^2-1)\bmod 8=(-26)\bmod 8=6.$$

For \(N=4\), the block totals are

$$B_1=0,\qquad B_2=2,\qquad B_3=16,\qquad B_4=64,$$

hence

$$S(4)=0+2+16+64=82.$$

This matches the direct sum of \(t(n)\) for \(1\le n\le 16\).

How the Code Works

The C++, Python, and Java implementations all follow the same fast path. First they compute the modular inverses needed for the fractions \(1/3\) and \(1/7\). Next they evaluate \(4^N\), \(8^{\lfloor N/2\rfloor}\), and \(2^{N+1}\) with fast modular exponentiation. Those three powers are then assembled into

$$\text{sumA}=\frac{4^N-1}{3},\qquad \text{sumB}=\frac{3(8^{\lfloor N/2\rfloor}-1)}{7}+\mathbf{1}_{N\text{ odd}}\,8^{\lfloor N/2\rfloor},\qquad \text{sumC}=2^{N+1}-2,$$

and the final answer is \(\text{sumA}+\text{sumB}-\text{sumC}\) modulo \(900497239\). The C++ implementation also keeps small-scale validation logic: it compares the closed form with direct summation for small \(N\) and checks the first few losing offsets against brute force, but the production result is obtained entirely from the closed formula.

Complexity Analysis

The optimized method uses only a constant number of modular exponentiations, so the running time is \(O(\log N)\) and the memory usage is \(O(1)\). By contrast, the defining sum already has \(2^N\) terms before considering any game-state search, so the closed form is what makes the problem tractable.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=900
  2. Impartial combinatorial games: Wikipedia - Impartial game
  3. Modular arithmetic: Wikipedia - Modular arithmetic
  4. Quadratic residues: Wikipedia - Quadratic residue
  5. Geometric series: Wikipedia - Geometric series
  6. Modular exponentiation: Wikipedia - Modular exponentiation

Problem 900 source code

C++

#include <algorithm>
#include <cassert>
#include <cstdint>
#include <future>
#include <iostream>
#include <numeric>
#include <unordered_map>
#include <utility>
#include <vector>

using namespace std;

static constexpr long long MOD = 900497239LL;

static long long mod_mul(long long a, long long b) {
    return static_cast<long long>((__int128)a * b % MOD);
}

static long long mod_pow(long long base, long long exp) {
    long long result = 1 % MOD;
    base %= MOD;
    while (exp > 0) {
        if (exp & 1LL) {
            result = mod_mul(result, base);
        }
        base = mod_mul(base, base);
        exp >>= 1LL;
    }
    return result;
}

static long long mod_inv(long long a) {
    long long b = MOD;
    long long x0 = 1, y0 = 0;
    long long x1 = 0, y1 = 1;
    while (b != 0) {
        long long q = a / b;
        long long t = a % b;
        a = b;
        b = t;
        t = x0 - q * x1;
        x0 = x1;
        x1 = t;
        t = y0 - q * y1;
        y0 = y1;
        y1 = t;
    }
    if (a != 1) {
        cerr << "Modular inverse does not exist.\n";
        exit(1);
    }
    long long res = x0 % MOD;
    if (res < 0) res += MOD;
    return res;
}

static long long compute_S_closed(long long N) {
    long long inv3 = mod_inv(3);
    long long inv7 = mod_inv(7);

    long long pow4N = mod_pow(4, N);
    long long sumA = mod_mul((pow4N - 1 + MOD) % MOD, inv3);

    long long sumC = (mod_pow(2, N + 1) - 2) % MOD;
    if (sumC < 0) sumC += MOD;

    long long K = N / 2;
    long long pow8K = mod_pow(8, K);
    long long geo = mod_mul((pow8K - 1 + MOD) % MOD, inv7);
    long long sumB = mod_mul(3, geo);
    if (N % 2 == 1) {
        sumB += pow8K;
        if (sumB >= MOD) sumB -= MOD;
    }

    long long ans = (sumA + sumB - sumC) % MOD;
    if (ans < 0) ans += MOD;
    return ans;
}

static inline int bitlen_u64(uint64_t x) {
    return 64 - __builtin_clzll(x);
}

static uint64_t t_formula_u64(uint64_t n) {
    int m = bitlen_u64(n);
    uint64_t mask = (m == 64) ? ~0ULL : ((1ULL << m) - 1ULL);
    uint64_t r = static_cast<uint64_t>((__int128)n * n) & mask;
    uint64_t target = (n & 1ULL) ? mask : 0ULL;
    return (target - r) & mask;
}

static long long compute_S_direct_small(int N) {
    uint64_t lim = 1ULL << N;
    __int128 acc = 0;
    for (uint64_t n = 1; n <= lim; ++n) {
        acc += t_formula_u64(n);
    }
    return static_cast<long long>(acc % MOD);
}

struct VecHash {
    size_t operator()(const vector<int> &v) const noexcept {
        size_t h = 1469598103934665603ULL;
        for (int x : v) {
            h ^= static_cast<size_t>(x) + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2);
        }
        return h;
    }
};

struct BruteSolver {
    unordered_map<vector<int>, bool, VecHash> memo;

    bool is_losing(vector<int> piles) {
        sort(piles.begin(), piles.end());
        auto it = memo.find(piles);
        if (it != memo.end()) return it->second;

        int m = static_cast<int>(piles.size());
        int s = piles[0];

        long long maxSum = 0;
        for (int a : piles) {
            maxSum += (a - 1);
        }
        if (maxSum < s) {
            memo[piles] = true;
            return true;
        }

        vector<int> take(m, 0);
        function<bool(int, int)> dfs = [&](int i, int remaining) -> bool {
            if (i == m) {
                if (remaining != 0) return false;
                vector<int> next = piles;
                for (int j = 0; j < m; ++j) {
                    next[j] -= take[j];
                }
                if (is_losing(next)) return true;
                return false;
            }
            int ai = piles[i];
            int up = min(ai - 1, remaining);
            for (int x = 0; x <= up; ++x) {
                take[i] = x;
                if (dfs(i + 1, remaining - x)) return true;
            }
            take[i] = 0;
            return false;
        };

        bool hasWinningMove = dfs(0, s);
        bool res = !hasWinningMove;
        memo[piles] = res;
        return res;
    }

    int t_bruteforce(int n) {
        int m = 32 - __builtin_clz(n);
        int M = 1 << m;
        for (int k = 0; k < M; ++k) {
            vector<int> piles(n + 1, n);
            piles.back() = n + k;
            if (is_losing(piles)) return k;
        }
        return -1;
    }
};

static void run_validations() {
    {
        long long got = compute_S_closed(10);
        if (got != 361522) {
            cerr << "Validation failed: S(10) expected 361522, got " << got << "\n";
            exit(1);
        }
    }

    for (int N = 1; N <= 14; ++N) {
        long long a = compute_S_closed(N);
        long long b = compute_S_direct_small(N);
        if (a != b) {
            cerr << "Validation failed: N=" << N << " closed=" << a << " direct=" << b << "\n";
            exit(1);
        }
    }

    vector<future<void>> futs;
    for (int n = 1; n <= 8; ++n) {
        futs.push_back(async(launch::async, [n]() {
            BruteSolver bs;
            int brute = bs.t_bruteforce(n);
            uint64_t pred = t_formula_u64(static_cast<uint64_t>(n));
            if (static_cast<uint64_t>(brute) != pred) {
                cerr << "Validation failed: t(" << n << ") brute=" << brute
                     << " pred=" << pred << "\n";
                exit(1);
            }
        }));
    }
    for (auto &f : futs) {
        f.get();
    }
}

int main(int argc, char **argv) {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    bool validate = false;
    long long N = 10000;
    for (int i = 1; i < argc; ++i) {
        string s = argv[i];
        if (s == "--validate") {
            validate = true;
        } else {
            N = stoll(s);
        }
    }

    if (validate) {
        run_validations();
    }

    cout << compute_S_closed(N) << "\n";
    return 0;
}

Python

MOD = 900497239

def mod_mul(a, b):
    return (a * b) % MOD

def mod_pow(base, exp):
    return pow(base, exp, MOD)

def mod_inv(a):
    return pow(a, MOD - 2, MOD)

def compute_S_closed(N):
    inv3 = mod_inv(3)
    inv7 = mod_inv(7)

    pow4N = mod_pow(4, N)
    sumA = (pow4N - 1 + MOD) % MOD
    sumA = mod_mul(sumA, inv3)

    sumC = (mod_pow(2, N + 1) - 2) % MOD
    if sumC < 0:
        sumC += MOD

    K = N // 2
    pow8K = mod_pow(8, K)
    geo = mod_mul((pow8K - 1 + MOD) % MOD, inv7)
    sumB = mod_mul(3, geo)
    if N % 2 == 1:
        sumB = (sumB + pow8K) % MOD

    ans = (sumA + sumB - sumC) % MOD
    if ans < 0:
        ans += MOD
    return ans

def solve():
    return str(compute_S_closed(10000))

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

Java

public class Euler900 {
    static final long MOD = 900497239L;

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

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

    static long modInv(long a) {
        return modPow(a, MOD - 2);
    }

    static long computeSClosed(long N) {
        long inv3 = modInv(3);
        long inv7 = modInv(7);

        long pow4N = modPow(4, N);
        long sumA = modMul((pow4N - 1 + MOD) % MOD, inv3);

        long sumC = (modPow(2, N + 1) - 2) % MOD;
        if (sumC < 0)
            sumC += MOD;

        long K = N / 2;
        long pow8K = modPow(8, K);
        long geo = modMul((pow8K - 1 + MOD) % MOD, inv7);
        long sumB = modMul(3, geo);
        if (N % 2 == 1) {
            sumB += pow8K;
            if (sumB >= MOD)
                sumB -= MOD;
        }

        long ans = (sumA + sumB - sumC) % MOD;
        if (ans < 0)
            ans += MOD;
        return ans;
    }

    public static String solve() {
        return Long.toString(computeSClosed(10000));
    }

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