Problem 676: Matching Digit Sums

View on Project Euler

Project Euler Problem 676 Solution

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

Problem Summary Fix integers \(k\) and \(\ell\) with \(3 \le k \le 6\) and \(1 \le \ell \le k-2\). Write a positive integer \(x\) in binary as $$x=\sum_{p\ge 0} b_p 2^p,\qquad b_p\in\{0,1\}.$$ For each period \(r\), define the weighted binary digit sum $$B_r(x)=\sum_{p\ge 0} b_p\,2^{p \bmod r}.$$ For a fixed pair \((k,\ell)\), the quantity of interest is $$S_{k,\ell}(n)=\sum_{1\le x\le n,\ B_k(x)=B_\ell(x)} x \pmod{10^{16}}.$$ The full Project Euler task asks for the sum of \(S_{k,\ell}(10^{16})\) over all admissible pairs \((k,\ell)\). Direct enumeration is impossible, so the solution turns the digit condition into a bounded-state dynamic program over the bits of \(n\). Mathematical Approach The central trick is to replace the equality of two weighted digit sums by a single signed balance equation. Once that balance is tracked while scanning the binary digits of \(n\), the algorithm can count all admissible numbers and add their values without listing them individually....

Detailed mathematical approach

Problem Summary

Fix integers \(k\) and \(\ell\) with \(3 \le k \le 6\) and \(1 \le \ell \le k-2\). Write a positive integer \(x\) in binary as

$$x=\sum_{p\ge 0} b_p 2^p,\qquad b_p\in\{0,1\}.$$

For each period \(r\), define the weighted binary digit sum

$$B_r(x)=\sum_{p\ge 0} b_p\,2^{p \bmod r}.$$

For a fixed pair \((k,\ell)\), the quantity of interest is

$$S_{k,\ell}(n)=\sum_{1\le x\le n,\ B_k(x)=B_\ell(x)} x \pmod{10^{16}}.$$

The full Project Euler task asks for the sum of \(S_{k,\ell}(10^{16})\) over all admissible pairs \((k,\ell)\). Direct enumeration is impossible, so the solution turns the digit condition into a bounded-state dynamic program over the bits of \(n\).

Mathematical Approach

The central trick is to replace the equality of two weighted digit sums by a single signed balance equation. Once that balance is tracked while scanning the binary digits of \(n\), the algorithm can count all admissible numbers and add their values without listing them individually.

Step 1: Rewrite the matching condition as a signed sum

For each bit position \(p\), define

$$\delta_p=2^{p \bmod k}-2^{p \bmod \ell}.$$

Then the difference between the two weighted sums is

$$B_k(x)-B_\ell(x)=\sum_{p\ge 0} b_p\,\delta_p.$$

Therefore the matching condition is equivalent to

$$\sum_{p\ge 0} b_p\,\delta_p=0.$$

Each bit set to \(1\) contributes a fixed signed amount to a running balance, while each bit set to \(0\) contributes nothing. The coefficient pattern is periodic because it depends only on \(p \bmod k\) and \(p \bmod \ell\), but the program only needs the finitely many positions that appear below the highest bit of \(n\).

Step 2: Bound the range of reachable balances

Let

$$L=\lfloor\log_2 n\rfloor+1$$

be the number of relevant bit positions when \(n>0\). Over these positions, define the safe bound

$$A=\sum_{p=0}^{L-1} |\delta_p|.$$

No partial or final balance can ever leave the interval

$$-A\le d\le A.$$

This observation turns the problem into a finite-state DP. Instead of storing balances in a dictionary, the implementation shifts the interval by an offset and uses ordinary arrays of width

$$2A+1.$$

That is why the method stays compact even though the original condition is phrased over all integers \(x\le n\).

Step 3: Use prefix-constrained binary digit DP

Process the bits of \(n\) from the most significant bit down to the least significant bit. After deciding the higher bits, a DP state records

$$\text{state}=(d,e),$$

where \(d\) is the current balance and \(e\in\{0,1\}\) is a prefix-equality flag. The value \(e=1\) means the chosen prefix is still exactly equal to the prefix of \(n\); the value \(e=0\) means the constructed number is already smaller, so the remaining bits may be chosen freely.

If the current bit of \(n\) is \(n_p\in\{0,1\}\) and we choose the new bit \(u\in\{0,1\}\), then in the equal-prefix layer we must respect \(u\le n_p\). The balance update is

$$d'=d+u\,\delta_p.$$

The new flag becomes

$$e'=\begin{cases} 1,& e=1\ \text{and}\ u=n_p,\\ 0,& \text{otherwise}. \end{cases}$$

This is the standard digit-DP tightness idea, specialized to binary digits and to the balance defined above.

Step 4: Carry both counts and value sums

For every reachable state, the DP stores two quantities modulo \(10^{16}\):

$$C(d,e)=\text{number of represented prefixes},\qquad V(d,e)=\text{sum of their numeric values}.$$

When the next bit \(u\) is appended at position \(p\), every represented number gains \(u\,2^p\). Therefore the transition rules are

$$C'(d',e')\equiv C'(d',e')+C(d,e)\pmod{10^{16}},$$

$$V'(d',e')\equiv V'(d',e')+V(d,e)+u\,2^p\,C(d,e)\pmod{10^{16}}.$$

The second formula is what makes the solution efficient: a single transition adds the contribution of an entire family of numbers at once, rather than adding admissible integers one by one.

Step 5: Read the answer from balance \(0\)

After all \(L\) bits are processed, the condition \(B_k(x)=B_\ell(x)\) is satisfied exactly when the final balance is zero. Hence

$$S_{k,\ell}(n)=V_{\mathrm{final}}(0,0)+V_{\mathrm{final}}(0,1)\pmod{10^{16}}.$$

The state for \(x=0\) is harmless: it also ends with balance \(0\), but it contributes value \(0\), so the sum over positive integers is unchanged.

Worked Example: \(n=10\), \(k=3\), \(\ell=1\)

Because \(10=1010_2\), the relevant positions are \(p=0,1,2,3\). Since \(p \bmod 1=0\) always, the coefficients are

$$\delta_0=2^0-2^0=0,\qquad \delta_1=2^1-2^0=1,\qquad \delta_2=2^2-2^0=3,\qquad \delta_3=2^0-2^0=0.$$

So bits at positions \(0\) and \(3\) do not affect the balance, while bits at positions \(1\) and \(2\) add \(1\) and \(3\). Among the integers \(1\le x\le 10\), the only values with total balance \(0\) are

$$1=0001_2,\qquad 8=1000_2,\qquad 9=1001_2.$$

Their sum is

$$1+8+9=18,$$

which matches the first checkpoint used by the C++, Python, and Java implementations. Those implementations also verify the larger checkpoints \(292\) for \(n=100\) and \(19{,}173{,}952\) for \(n=10^6\) when \((k,\ell)=(3,1)\).

How the Code Works

The implementations begin with the trivial case \(n=0\), for which the required sum is \(0\). Otherwise they compute the bit length of \(n\), generate the coefficient list \(\delta_0,\dots,\delta_{L-1}\), and add their absolute values to determine the balance range.

Next they maintain four rolling arrays: counts and value sums for the smaller-prefix layer, and counts and value sums for the equal-prefix layer. For each bit position, fresh arrays are created for the next layer, the allowable bit choices are applied, and all arithmetic is reduced modulo \(10^{16}\).

The C++ implementation uses 128-bit intermediate multiplication to keep modular products safe, the Python implementation relies on arbitrary-precision integers, and the Java implementation uses repeated doubling so multiplication cannot overflow a signed 64-bit value. After evaluating one pair \((k,\ell)\), the implementation adds that contribution to the running total over all required pairs and prints the final result as a 16-digit decimal string.

Complexity Analysis

For fixed \(n\), \(k\), and \(\ell\), let

$$L=\lfloor\log_2 n\rfloor+1,\qquad A=\sum_{p=0}^{L-1} |\delta_p|.$$

The DP uses two prefix layers and \(2A+1\) balance states. Therefore the running time is

$$O\bigl(L(2A+1)\bigr)=O(LA),$$

and the memory usage with rolling arrays is

$$O(2A+1)=O(A).$$

In the actual problem there are only ten admissible pairs \((k,\ell)\), and the bit length of \(10^{16}\) is small, so this bounded-state DP is easily fast enough in all three languages.

Footnotes and References

  1. Problem page: Project Euler 676
  2. Binary number: Wikipedia - Binary number
  3. Dynamic programming: Wikipedia - Dynamic programming
  4. Modular arithmetic: Wikipedia - Modular arithmetic
  5. Digit DP overview: GeeksforGeeks - Digit DP Introduction

Problem 676 source code

C++

#include <cassert>
#include <cstdint>
#include <cstdlib>
#include <iomanip>
#include <iostream>
#include <vector>

namespace {

using u64 = std::uint64_t;
using u128 = __uint128_t;

constexpr u64 MOD16 = 10'000'000'000'000'000ULL;

u64 add_mod(u64 a, u64 b) {
    const u64 s = a + b;
    if (s >= MOD16 || s < a) {
        return s - MOD16;
    }
    return s;
}

u64 mul_mod(u64 a, u64 b) {
    return static_cast<u64>((static_cast<u128>(a) * static_cast<u128>(b)) % MOD16);
}

u64 M(u64 n, int k, int l) {
    if (n == 0ULL) {
        return 0ULL;
    }

    const int max_bit = 64 - __builtin_clzll(n);

    std::vector<int> coeff(static_cast<std::size_t>(max_bit), 0);
    int max_abs = 0;
    for (int p = 0; p < max_bit; ++p) {
        coeff[static_cast<std::size_t>(p)] =
            (1 << (p % k)) - (1 << (p % l));
        max_abs += std::abs(coeff[static_cast<std::size_t>(p)]);
    }

    const int offset = max_abs;
    const int width = 2 * max_abs + 1;

    std::vector<u64> cnt_lo(static_cast<std::size_t>(width), 0ULL);
    std::vector<u64> cnt_hi(static_cast<std::size_t>(width), 0ULL);
    std::vector<u64> sum_lo(static_cast<std::size_t>(width), 0ULL);
    std::vector<u64> sum_hi(static_cast<std::size_t>(width), 0ULL);

    cnt_hi[static_cast<std::size_t>(offset)] = 1ULL;

    for (int pos = max_bit - 1; pos >= 0; --pos) {
        std::vector<u64> ncnt_lo(static_cast<std::size_t>(width), 0ULL);
        std::vector<u64> ncnt_hi(static_cast<std::size_t>(width), 0ULL);
        std::vector<u64> nsum_lo(static_cast<std::size_t>(width), 0ULL);
        std::vector<u64> nsum_hi(static_cast<std::size_t>(width), 0ULL);

        const int lim = static_cast<int>((n >> pos) & 1ULL);
        const int c = coeff[static_cast<std::size_t>(pos)];
        const u64 bit_value = (1ULL << pos) % MOD16;

        for (int idx = 0; idx < width; ++idx) {
            const u64 c0 = cnt_lo[static_cast<std::size_t>(idx)];
            if (c0 != 0ULL) {
                const u64 s0 = sum_lo[static_cast<std::size_t>(idx)];

                ncnt_lo[static_cast<std::size_t>(idx)] += c0;
                nsum_lo[static_cast<std::size_t>(idx)] =
                    add_mod(nsum_lo[static_cast<std::size_t>(idx)], s0);

                const int j = idx + c;
                ncnt_lo[static_cast<std::size_t>(j)] += c0;
                const u64 add1 = add_mod(s0, mul_mod(c0, bit_value));
                nsum_lo[static_cast<std::size_t>(j)] =
                    add_mod(nsum_lo[static_cast<std::size_t>(j)], add1);
            }

            const u64 c1 = cnt_hi[static_cast<std::size_t>(idx)];
            if (c1 == 0ULL) {
                continue;
            }
            const u64 s1 = sum_hi[static_cast<std::size_t>(idx)];

            if (lim == 0) {
                ncnt_hi[static_cast<std::size_t>(idx)] += c1;
                nsum_hi[static_cast<std::size_t>(idx)] =
                    add_mod(nsum_hi[static_cast<std::size_t>(idx)], s1);
            } else {
                ncnt_lo[static_cast<std::size_t>(idx)] += c1;
                nsum_lo[static_cast<std::size_t>(idx)] =
                    add_mod(nsum_lo[static_cast<std::size_t>(idx)], s1);

                const int j = idx + c;
                ncnt_hi[static_cast<std::size_t>(j)] += c1;
                const u64 add1 = add_mod(s1, mul_mod(c1, bit_value));
                nsum_hi[static_cast<std::size_t>(j)] =
                    add_mod(nsum_hi[static_cast<std::size_t>(j)], add1);
            }
        }

        cnt_lo.swap(ncnt_lo);
        cnt_hi.swap(ncnt_hi);
        sum_lo.swap(nsum_lo);
        sum_hi.swap(nsum_hi);
    }

    return add_mod(sum_lo[static_cast<std::size_t>(offset)],
                   sum_hi[static_cast<std::size_t>(offset)]);
}

}  // namespace

int main() {
    assert(M(10ULL, 3, 1) == 18ULL);
    assert(M(100ULL, 3, 1) == 292ULL);
    assert(M(1'000'000ULL, 3, 1) == 19'173'952ULL);

    u64 ans = 0ULL;
    for (int k = 3; k <= 6; ++k) {
        for (int l = 1; l <= k - 2; ++l) {
            ans = add_mod(ans, M(10'000'000'000'000'000ULL, k, l));
        }
    }

    std::cout << std::setw(16) << std::setfill('0') << ans << "\n";
    return 0;
}

Python

MOD16 = 10**16

def add_mod(a, b):
    s = a + b
    if s >= MOD16:
        s -= MOD16
    return s

def mul_mod(a, b):
    return (a * b) % MOD16

def M(n, k, l):
    if n == 0:
        return 0

    max_bit = n.bit_length()

    coeff = [0] * max_bit
    max_abs = 0
    for p in range(max_bit):
        coeff[p] = (1 << (p % k)) - (1 << (p % l))
        max_abs += abs(coeff[p])

    offset = max_abs
    width = 2 * max_abs + 1

    cnt_lo = [0] * width
    cnt_hi = [0] * width
    sum_lo = [0] * width
    sum_hi = [0] * width

    cnt_hi[offset] = 1

    for pos in range(max_bit - 1, -1, -1):
        ncnt_lo = [0] * width
        ncnt_hi = [0] * width
        nsum_lo = [0] * width
        nsum_hi = [0] * width

        lim = (n >> pos) & 1
        c = coeff[pos]
        bit_value = (1 << pos) % MOD16

        for idx in range(width):
            c0 = cnt_lo[idx]
            if c0 != 0:
                s0 = sum_lo[idx]

                ncnt_lo[idx] = add_mod(ncnt_lo[idx], c0)
                nsum_lo[idx] = add_mod(nsum_lo[idx], s0)

                j = idx + c
                ncnt_lo[j] = add_mod(ncnt_lo[j], c0)
                add1 = add_mod(s0, mul_mod(c0, bit_value))
                nsum_lo[j] = add_mod(nsum_lo[j], add1)

            c1 = cnt_hi[idx]
            if c1 == 0:
                continue
            s1 = sum_hi[idx]

            if lim == 0:
                ncnt_hi[idx] = add_mod(ncnt_hi[idx], c1)
                nsum_hi[idx] = add_mod(nsum_hi[idx], s1)
            else:
                ncnt_lo[idx] = add_mod(ncnt_lo[idx], c1)
                nsum_lo[idx] = add_mod(nsum_lo[idx], s1)

                j = idx + c
                ncnt_hi[j] = add_mod(ncnt_hi[j], c1)
                add1 = add_mod(s1, mul_mod(c1, bit_value))
                nsum_hi[j] = add_mod(nsum_hi[j], add1)

        cnt_lo = ncnt_lo
        cnt_hi = ncnt_hi
        sum_lo = nsum_lo
        sum_hi = nsum_hi

    return add_mod(sum_lo[offset], sum_hi[offset])

def solve():
    ans = 0
    for k in range(3, 7):
        for l in range(1, k - 1):
            ans = add_mod(ans, M(10**16, k, l))

    return f"{ans:016d}"

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

Java

public class Euler676 {

    static final long MOD16 = 10000000000000000L;

    static long addMod(long a, long b) {
        long s = a + b;
        if (s >= MOD16 || s < a) {
            return s - MOD16;
        }
        return s;
    }

    static long mulMod(long a, long b) {
        long res = 0;
        a %= MOD16;
        b %= MOD16;
        while (b > 0) {
            if ((b & 1) != 0) {
                res += a;
                if (res >= MOD16)
                    res -= MOD16;
            }
            a += a;
            if (a >= MOD16)
                a -= MOD16;
            b >>= 1;
        }
        return res;
    }

    static int maxBit(long n) {
        if (n == 0)
            return 0;
        return 64 - Long.numberOfLeadingZeros(n);
    }

    static long M(long n, int k, int l) {
        if (n == 0) {
            return 0;
        }

        int max_bit = maxBit(n);

        int[] coeff = new int[max_bit];
        int max_abs = 0;
        for (int p = 0; p < max_bit; ++p) {
            coeff[p] = (1 << (p % k)) - (1 << (p % l));
            max_abs += Math.abs(coeff[p]);
        }

        int offset = max_abs;
        int width = 2 * max_abs + 1;

        long[] cntLo = new long[width];
        long[] cntHi = new long[width];
        long[] sumLo = new long[width];
        long[] sumHi = new long[width];

        cntHi[offset] = 1L;

        for (int pos = max_bit - 1; pos >= 0; --pos) {
            long[] ncntLo = new long[width];
            long[] ncntHi = new long[width];
            long[] nsumLo = new long[width];
            long[] nsumHi = new long[width];

            int lim = (int) ((n >> pos) & 1L);
            int c = coeff[pos];
            long bitValue = (1L << pos) % MOD16;

            for (int idx = 0; idx < width; ++idx) {
                long c0 = cntLo[idx];
                if (c0 != 0L) {
                    long s0 = sumLo[idx];

                    ncntLo[idx] = addMod(ncntLo[idx], c0);
                    nsumLo[idx] = addMod(nsumLo[idx], s0);

                    int j = idx + c;
                    ncntLo[j] = addMod(ncntLo[j], c0);
                    long add1 = addMod(s0, mulMod(c0, bitValue));
                    nsumLo[j] = addMod(nsumLo[j], add1);
                }

                long c1 = cntHi[idx];
                if (c1 == 0L) {
                    continue;
                }
                long s1 = sumHi[idx];

                if (lim == 0) {
                    ncntHi[idx] = addMod(ncntHi[idx], c1);
                    nsumHi[idx] = addMod(nsumHi[idx], s1);
                } else {
                    ncntLo[idx] = addMod(ncntLo[idx], c1);
                    nsumLo[idx] = addMod(nsumLo[idx], s1);

                    int j = idx + c;
                    ncntHi[j] = addMod(ncntHi[j], c1);
                    long add1 = addMod(s1, mulMod(c1, bitValue));
                    nsumHi[j] = addMod(nsumHi[j], add1);
                }
            }

            cntLo = ncntLo;
            cntHi = ncntHi;
            sumLo = nsumLo;
            sumHi = nsumHi;
        }

        return addMod(sumLo[offset], sumHi[offset]);
    }

    public static String solve() {
        long ans = 0L;
        for (int k = 3; k <= 6; ++k) {
            for (int l = 1; l <= k - 2; ++l) {
                ans = addMod(ans, M(10000000000000000L, k, l));
            }
        }
        return String.format("%016d", ans);
    }

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