Problem 271: Modular Cubes, Part 1

View on Project Euler

Project Euler Problem 271 Solution

EulerSolve provides an optimized solution for Project Euler Problem 271, Modular Cubes, Part 1, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary We want all residues \(x\) satisfying $$x^3\equiv1\pmod n,$$ and then we sum all nontrivial solutions, meaning all residues with \(1<x<n\). The default value used by the code is $$n=13082761331670030=2\cdot3\cdot5\cdot7\cdot11\cdot13\cdot17\cdot19\cdot23\cdot29\cdot31\cdot37\cdot41\cdot43,$$ which is squarefree. Mathematical Approach 1. Split the Problem Prime by Prime Because \(n\) is squarefree, we may write $$n=\prod_{i=1}^{k}p_i$$ with distinct primes \(p_i\). Then $$x^3\equiv1\pmod n$$ is equivalent to the simultaneous system $$x^3\equiv1\pmod{p_i}\qquad(i=1,\dots,k).$$ So the whole problem becomes: 1. find all cube roots of unity modulo each prime \(p_i\), 2. combine one local choice from each prime modulus using the Chinese Remainder Theorem. 2. How Many Cube Roots Exist Modulo a Prime? For a prime \(p\), the nonzero residues modulo \(p\) form a cyclic multiplicative group of size \(p-1\). In a cyclic group of order \(p-1\), the equation $$u^3=1$$ has exactly $$\gcd(3,p-1)$$ solutions. Therefore: 1. if \(p\equiv2\pmod3\), then \(\gcd(3,p-1)=1\), so the only root is \(1\), 2. if \(p\equiv1\pmod3\), then \(\gcd(3,p-1)=3\), so there are three roots. The code does not use a special generator argument; since every prime factor of the target \(n\) is small, it simply brute-forces all residues \(1\le x<p\) and keeps those with \(x^3\equiv1\pmod p\). 3....

Detailed mathematical approach

Problem Summary

We want all residues \(x\) satisfying

$$x^3\equiv1\pmod n,$$

and then we sum all nontrivial solutions, meaning all residues with \(1<x<n\). The default value used by the code is

$$n=13082761331670030=2\cdot3\cdot5\cdot7\cdot11\cdot13\cdot17\cdot19\cdot23\cdot29\cdot31\cdot37\cdot41\cdot43,$$

which is squarefree.

Mathematical Approach

1. Split the Problem Prime by Prime

Because \(n\) is squarefree, we may write

$$n=\prod_{i=1}^{k}p_i$$

with distinct primes \(p_i\). Then

$$x^3\equiv1\pmod n$$

is equivalent to the simultaneous system

$$x^3\equiv1\pmod{p_i}\qquad(i=1,\dots,k).$$

So the whole problem becomes:

1. find all cube roots of unity modulo each prime \(p_i\),

2. combine one local choice from each prime modulus using the Chinese Remainder Theorem.

2. How Many Cube Roots Exist Modulo a Prime?

For a prime \(p\), the nonzero residues modulo \(p\) form a cyclic multiplicative group of size \(p-1\). In a cyclic group of order \(p-1\), the equation

$$u^3=1$$

has exactly

$$\gcd(3,p-1)$$

solutions. Therefore:

1. if \(p\equiv2\pmod3\), then \(\gcd(3,p-1)=1\), so the only root is \(1\),

2. if \(p\equiv1\pmod3\), then \(\gcd(3,p-1)=3\), so there are three roots.

The code does not use a special generator argument; since every prime factor of the target \(n\) is small, it simply brute-forces all residues \(1\le x<p\) and keeps those with \(x^3\equiv1\pmod p\).

3. Small Local Examples

For \(p=7\), the roots are

$$R_7=\{1,2,4\},$$

because \(1^3\equiv2^3\equiv4^3\equiv1\pmod7\).

For \(p=13\), the roots are

$$R_{13}=\{1,3,9\}.$$

For primes such as \(5,11,17,23,29,41\), which are \(2\pmod3\), the local root set is just \(\{1\}\).

4. Chinese Remainder Reconstruction

Once we choose one root \(r_i\in R_{p_i}\) for every prime factor, there is a unique residue modulo \(n\) satisfying

$$x\equiv r_i\pmod{p_i}\qquad(i=1,\dots,k).$$

The implementation combines congruences two at a time. If we already know

$$x\equiv a_1\pmod{m_1},\qquad x\equiv a_2\pmod{m_2},$$

with \(\gcd(m_1,m_2)=1\), then the merged solution is

$$x=a_1+m_1\left((a_2-a_1)m_1^{-1}\bmod m_2\right).$$

This is exactly the formula implemented in crt_pair.

5. Why Enumeration Is Tiny

The default number \(n\) has 14 distinct prime factors. Among them, exactly six primes are \(1\pmod3\):

$$7,13,19,31,37,43.$$

Those contribute three local roots each. Every other prime contributes only one local root. Therefore the total number of global solutions is

$$3^6=729.$$

That is why a direct DFS over all CRT combinations is entirely practical.

6. Worked Checkpoint: \(n=7\)

The roots modulo \(7\) are \(\{1,2,4\}\). The trivial root \(1\) is excluded from the final sum, so the code returns

$$2+4=6.$$

This matches the checkpoint solve(7)=6.

7. Worked Checkpoint: \(n=91=7\cdot13\)

Here we combine

$$R_7=\{1,2,4\},\qquad R_{13}=\{1,3,9\}.$$

By CRT, each pair \((r_7,r_{13})\) gives exactly one solution modulo \(91\), so there are

$$3\cdot3=9$$

solutions in total. They are

$$1,\;9,\;16,\;22,\;29,\;53,\;74,\;79,\;81.$$

Excluding the trivial residue \(1\), the sum is

$$9+16+22+29+53+74+79+81=363,$$

which is exactly the checkpoint solve(91)=363.

8. Final Summation Rule

The DFS enumerates all global CRT solutions. At the leaf of the recursion, the code adds the residue only if

$$1<x<n.$$

That removes the always-present trivial solution \(x=1\) and keeps every nontrivial cube root of unity modulo \(n\).

How the Code Works

distinct_prime_factors extracts the distinct prime divisors of \(n\).

mod_pow tests local candidates by checking \(x^3\bmod p\).

crt_pair merges two congruences using an inverse computed by mod_inverse.

solve first builds roots_by_prime, then runs a DFS over all local root choices, updating the current CRT residue and modulus at each step.

The command-line interface supports --n=<value> and --skip-checkpoints.

Complexity Analysis

If the distinct prime factors are \(p_1,\dots,p_k\), then the number of DFS states is essentially

$$\prod_{i=1}^{k}\gcd(3,p_i-1).$$

Each transition performs only constant-time modular arithmetic at machine-word size, so the method is dominated by the number of local-root combinations.

Further Reading

  1. Problem page: https://projecteuler.net/problem=271
  2. Chinese Remainder Theorem: https://en.wikipedia.org/wiki/Chinese_remainder_theorem
  3. Primitive roots and roots of unity modulo primes: https://en.wikipedia.org/wiki/Multiplicative_group_of_integers_modulo_n

Problem 271 source code

C++

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

namespace {

using i64 = long long;
using u64 = std::uint64_t;
using i128 = __int128_t;

struct Options {
    u64 n = 13082761331670030ULL;
    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 = 0;
    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_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, "--n=", options.n)) {
            continue;
        }
        std::cerr << "Unknown argument: " << arg << '\n';
        return false;
    }
    return options.n > 2;
}

u64 mod_pow(u64 base, u64 exp, u64 mod) {
    u64 result = 1 % mod;
    u64 cur = base % mod;
    while (exp > 0) {
        if ((exp & 1ULL) != 0ULL) {
            result = static_cast<u64>((static_cast<i128>(result) * cur) % mod);
        }
        cur = static_cast<u64>((static_cast<i128>(cur) * cur) % mod);
        exp >>= 1U;
    }
    return result;
}

std::vector<u64> distinct_prime_factors(u64 n) {
    std::vector<u64> factors;
    if ((n & 1ULL) == 0ULL) {
        factors.push_back(2ULL);
        while ((n & 1ULL) == 0ULL) {
            n >>= 1U;
        }
    }
    for (u64 p = 3; p * p <= n; p += 2ULL) {
        if (n % p != 0ULL) {
            continue;
        }
        factors.push_back(p);
        while (n % p == 0ULL) {
            n /= p;
        }
    }
    if (n > 1ULL) {
        factors.push_back(n);
    }
    return factors;
}

i64 extended_gcd(const i64 a, const i64 b, i64& x, i64& y) {
    if (b == 0) {
        x = 1;
        y = 0;
        return a;
    }
    i64 x1 = 0;
    i64 y1 = 0;
    const i64 g = extended_gcd(b, a % b, x1, y1);
    x = y1;
    y = x1 - (a / b) * y1;
    return g;
}

u64 mod_inverse(const u64 a, const u64 mod) {
    i64 x = 0;
    i64 y = 0;
    const i64 g = extended_gcd(static_cast<i64>(a), static_cast<i64>(mod), x, y);
    if (g != 1) {
        return 0;
    }
    i64 r = x % static_cast<i64>(mod);
    if (r < 0) {
        r += static_cast<i64>(mod);
    }
    return static_cast<u64>(r);
}

u64 crt_pair(const u64 a1, const u64 m1, const u64 a2, const u64 m2) {
    const u64 inv = mod_inverse(m1 % m2, m2);
    const u64 t = static_cast<u64>((static_cast<i128>((a2 + m2 - (a1 % m2)) % m2) * inv) % m2);
    return static_cast<u64>(a1 + static_cast<i128>(m1) * t);
}

u64 solve(const u64 n) {
    const std::vector<u64> primes = distinct_prime_factors(n);

    std::vector<std::vector<u64>> roots_by_prime;
    roots_by_prime.reserve(primes.size());
    for (u64 p : primes) {
        std::vector<u64> roots;
        for (u64 x = 1; x < p; ++x) {
            if (mod_pow(x, 3, p) == 1ULL) {
                roots.push_back(x);
            }
        }
        roots_by_prime.push_back(std::move(roots));
    }

    u64 sum = 0;
    const auto dfs = [&](auto&& self, std::size_t idx, u64 residue, u64 modulus) -> void {
        if (idx == primes.size()) {
            if (residue > 1ULL && residue < n) {
                sum += residue;
            }
            return;
        }
        const u64 p = primes[idx];
        for (u64 root : roots_by_prime[idx]) {
            const u64 next_residue = crt_pair(residue, modulus, root, p);
            self(self, idx + 1, next_residue, modulus * p);
        }
    };
    dfs(dfs, 0, 0ULL, 1ULL);
    return sum;
}

bool run_checkpoints() {
    if (solve(91) != 363ULL) {
        std::cerr << "Checkpoint failed for n=91 sample" << '\n';
        return false;
    }
    if (solve(7) != 6ULL) {
        std::cerr << "Checkpoint failed for n=7" << '\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.n) << '\n';
    return 0;
}

Python

import math

def distinct_prime_factors(n):
    factors = []
    if n % 2 == 0:
        factors.append(2)
        while n % 2 == 0:
            n //= 2
    p = 3
    while p * p <= n:
        if n % p == 0:
            factors.append(p)
            while n % p == 0:
                n //= p
        p += 2
    if n > 1:
        factors.append(n)
    return factors

def extended_gcd(a, b):
    if b == 0:
        return a, 1, 0
    g, x1, y1 = extended_gcd(b, a % b)
    x = y1
    y = x1 - (a // b) * y1
    return g, x, y

def mod_inverse(a, mod):
    g, x, y = extended_gcd(a, mod)
    if g != 1: return 0
    return x % mod

def crt_pair(a1, m1, a2, m2):
    inv = mod_inverse(m1 % m2, m2)
    t = (((a2 - (a1 % m2)) % m2 + m2) % m2 * inv) % m2
    return a1 + m1 * t

def solve_for_n(n):
    primes = distinct_prime_factors(n)
    roots_by_prime = []
    for p in primes:
        roots = []
        for x in range(1, p):
            if pow(x, 3, p) == 1:
                roots.append(x)
        roots_by_prime.append(roots)

    ans = 0
    def dfs(idx, residue, modulus):
        nonlocal ans
        if idx == len(primes):
            if 1 < residue < n:
                ans += residue
            return
            
        p = primes[idx]
        for root in roots_by_prime[idx]:
            next_residue = crt_pair(residue, modulus, root, p)
            dfs(idx + 1, next_residue, modulus * p)

    dfs(0, 0, 1)
    return ans

def solve():
    n = 13082761331670030
    ans = solve_for_n(n)
    return str(ans)

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

Java

import java.math.BigInteger;
import java.util.ArrayList;
import java.util.List;

public class Euler271 {

    static List<Long> distinctPrimeFactors(long n) {
        List<Long> factors = new ArrayList<>();
        if ((n & 1) == 0) {
            factors.add(2L);
            while ((n & 1) == 0)
                n >>= 1;
        }
        for (long p = 3; p * p <= n; p += 2) {
            if (n % p == 0) {
                factors.add(p);
                while (n % p == 0)
                    n /= p;
            }
        }
        if (n > 1)
            factors.add(n);
        return factors;
    }

    static long[] extendedGcd(long a, long b) {
        if (b == 0)
            return new long[] { a, 1, 0 };
        long[] res = extendedGcd(b, a % b);
        long g = res[0];
        long x1 = res[1];
        long y1 = res[2];
        long x = y1;
        long y = x1 - (a / b) * y1;
        return new long[] { g, x, y };
    }

    static long modInverse(long a, long mod) {
        long[] res = extendedGcd(a, mod);
        if (res[0] != 1)
            return 0;
        long r = res[1] % mod;
        if (r < 0)
            r += mod;
        return r;
    }

    static long modPow(long base, long exp, long mod) {
        long res = 1 % mod;
        long cur = base % mod;
        while (exp > 0) {
            if ((exp & 1) == 1)
                res = multiplyMod(res, cur, mod);
            cur = multiplyMod(cur, cur, mod);
            exp >>= 1;
        }
        return res;
    }

    static long multiplyMod(long a, long b, long mod) {
        return BigInteger.valueOf(a).multiply(BigInteger.valueOf(b)).mod(BigInteger.valueOf(mod)).longValue();
    }

    static long crtPair(long a1, long m1, long a2, long m2) {
        long inv = modInverse(m1 % m2, m2);
        long diff = (a2 - (a1 % m2)) % m2;
        if (diff < 0)
            diff += m2;
        long t = multiplyMod(diff, inv, m2);
        return a1 + m1 * t;
    }

    static long sum = 0;
    static List<Long> primes;
    static List<List<Long>> rootsByPrime;
    static long N;

    static void dfs(int idx, long residue, long modulus) {
        if (idx == primes.size()) {
            if (residue > 1 && residue < N) {
                sum += residue;
            }
            return;
        }
        long p = primes.get(idx);
        for (long root : rootsByPrime.get(idx)) {
            long nextResidue = crtPair(residue, modulus, root, p);
            dfs(idx + 1, nextResidue, modulus * p);
        }
    }

    static long solve(long n) {
        N = n;
        primes = distinctPrimeFactors(n);
        rootsByPrime = new ArrayList<>();
        sum = 0;

        for (long p : primes) {
            List<Long> roots = new ArrayList<>();
            for (long x = 1; x < p; x++) {
                if (modPow(x, 3, p) == 1) {
                    roots.add(x);
                }
            }
            rootsByPrime.add(roots);
        }

        dfs(0, 0, 1);
        return sum;
    }

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