Problem 801: $x^y \equiv y^x$

View on Project Euler

Project Euler Problem 801 Solution

EulerSolve provides an optimized solution for Project Euler Problem 801, $x^y \equiv y^x$, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary For a prime \(p\), let \(R(p)\) be the number of ordered pairs \((x,y)\) with \(1\le x,y\le p(p-1)\) such that $$x^y \equiv y^x \pmod p.$$ The full task is to evaluate $$S(M,N)\equiv \sum_{\substack{M \le p \le N \\ p\text{ prime}}} R(p)\pmod{993353399}.$$ A direct search over all pairs is far too large, so the solution converts the congruence into a group-theoretic count that depends only on the factorization of \(p-1\). Mathematical Approach Fix one prime \(p\) and write \(n=p-1\). The multiplicative group \(\mathbb{F}_p^\times\) has order \(n\), and that is the structure the solution exploits. Step 1: Separate the Multiples of \(p\) The interval \(1,2,\dots,pn\) contains exactly \(n\) multiples of \(p\). If exactly one of \(x\) or \(y\) is divisible by \(p\), then one side of \(x^y \equiv y^x \pmod p\) is \(0\) and the other is nonzero, so such pairs never work. If both \(x\) and \(y\) are divisible by \(p\), then both sides are \(0\), so every such pair works. Therefore these pairs contribute $$n^2$$ solutions immediately. Everything else comes from pairs with \(p\nmid x\) and \(p\nmid y\). Step 2: Encode Nonzero Residues by a Primitive Root Choose a primitive root \(g\) modulo \(p\). Then every nonzero residue class can be written uniquely as \(g^a\) for some \(a\in \mathbb{Z}/n\mathbb{Z}\)....

Detailed mathematical approach

Problem Summary

For a prime \(p\), let \(R(p)\) be the number of ordered pairs \((x,y)\) with \(1\le x,y\le p(p-1)\) such that

$$x^y \equiv y^x \pmod p.$$

The full task is to evaluate

$$S(M,N)\equiv \sum_{\substack{M \le p \le N \\ p\text{ prime}}} R(p)\pmod{993353399}.$$

A direct search over all pairs is far too large, so the solution converts the congruence into a group-theoretic count that depends only on the factorization of \(p-1\).

Mathematical Approach

Fix one prime \(p\) and write \(n=p-1\). The multiplicative group \(\mathbb{F}_p^\times\) has order \(n\), and that is the structure the solution exploits.

Step 1: Separate the Multiples of \(p\)

The interval \(1,2,\dots,pn\) contains exactly \(n\) multiples of \(p\).

If exactly one of \(x\) or \(y\) is divisible by \(p\), then one side of \(x^y \equiv y^x \pmod p\) is \(0\) and the other is nonzero, so such pairs never work.

If both \(x\) and \(y\) are divisible by \(p\), then both sides are \(0\), so every such pair works. Therefore these pairs contribute

$$n^2$$

solutions immediately.

Everything else comes from pairs with \(p\nmid x\) and \(p\nmid y\).

Step 2: Encode Nonzero Residues by a Primitive Root

Choose a primitive root \(g\) modulo \(p\). Then every nonzero residue class can be written uniquely as \(g^a\) for some \(a\in \mathbb{Z}/n\mathbb{Z}\).

So for \(p\nmid x\) and \(p\nmid y\), write

$$x\equiv g^a \pmod p,\qquad y\equiv g^b \pmod p$$

with \(a,b\in \mathbb{Z}/n\mathbb{Z}\).

Now look at all integers in \(1\le x\le pn\) having the same residue modulo \(p\). They are

$$r,\ r+p,\ r+2p,\ \dots,\ r+(n-1)p.$$

Because \(p\equiv 1 \pmod n\), these numbers are congruent modulo \(n\) to

$$r,\ r+1,\ r+2,\ \dots,\ r+(n-1),$$

so each residue class modulo \(n\) occurs exactly once. Hence every pair

$$\bigl(a,u\bigr)\in (\mathbb{Z}/n\mathbb{Z})^2$$

corresponds to a unique integer \(x\) with

$$x\equiv g^a \pmod p,\qquad x\equiv u \pmod n,$$

and similarly every \((b,v)\) determines a unique \(y\).

Step 3: Count Exponent Pairs for Fixed Logarithms

Since \(g\) has order \(n\), we have

$$x^y \equiv y^x \pmod p \iff g^{ay}\equiv g^{bx}\pmod p \iff ay\equiv bx\pmod n.$$

Only the classes of \(x\) and \(y\) modulo \(n\) matter in the exponents, so with \(u\equiv x\pmod n\) and \(v\equiv y\pmod n\), the condition becomes

$$av\equiv bu\pmod n.$$

For fixed \(a\) and \(b\), define

$$\Phi_{a,b}(u,v)=bu-av \pmod n.$$

This is a homomorphism from \((\mathbb{Z}/n\mathbb{Z})^2\) to \(\mathbb{Z}/n\mathbb{Z}\). If

$$d=\gcd(a,b,n),$$

then the image of \(\Phi_{a,b}\) is exactly the subgroup of multiples of \(d\), which has size \(n/d\). Therefore the kernel has size

$$\frac{n^2}{n/d}=nd.$$

So for each fixed pair \((a,b)\), the number of admissible pairs \((u,v)\) is

$$n\,\gcd(a,b,n).$$

Summing over all \(a,b\in \mathbb{Z}/n\mathbb{Z}\) gives the nonzero contribution

$$n\,H(n),\qquad H(n)=\sum_{a=0}^{n-1}\sum_{b=0}^{n-1}\gcd(a,b,n).$$

Step 4: Turn the GCD Sum into a Divisor Sum

Use the standard identity

$$\gcd(a,b,n)=\sum_{d\mid \gcd(a,b,n)} \varphi(d),$$

where \(\varphi\) is Euler's totient function.

Substituting this into \(H(n)\) and swapping the order of summation yields

$$H(n)=\sum_{d\mid n}\varphi(d)\left(\frac{n}{d}\right)^2.$$

Indeed, for a fixed divisor \(d\mid n\), there are exactly \(n/d\) residue classes modulo \(n\) divisible by \(d\), both for \(a\) and for \(b\).

This formula already shows that \(H(n)\) is multiplicative in \(n\).

Step 5: Evaluate the Prime-Power Factor

Let

$$n=\prod_{i=1}^t q_i^{e_i}.$$

Because \(H\) is multiplicative, it is enough to evaluate one prime power:

$$H(q^e)=\sum_{k=0}^{e}\varphi(q^k)\,q^{2e-2k}.$$

Now \(\varphi(q^0)=1\) and \(\varphi(q^k)=q^k-q^{k-1}\) for \(k\ge 1\), so

$$H(q^e)=q^{2e}+\sum_{k=1}^{e}(q^k-q^{k-1})q^{2e-2k}.$$

This simplifies to

$$H(q^e)=q^{2e}+q^{2e-1}-q^{e-1}=q^{e-1}(q^{e+1}+q^e-1).$$

Therefore

$$H(n)=\prod_{i=1}^{t} q_i^{e_i-1}(q_i^{e_i+1}+q_i^{e_i}-1),$$

and the final prime-local count is

$$\boxed{R(p)=n^2+n\prod_{q^e\parallel n} q^{e-1}(q^{e+1}+q^e-1),\qquad n=p-1.}$$

Worked Example: \(p=5\)

Here \(n=p-1=4\), so the search range is \(1\le x,y\le 20\).

There are \(4\) multiples of \(5\) in that range, namely \(5,10,15,20\). Pairs with both entries divisible by \(5\) contribute

$$4^2=16$$

solutions.

Next compute

$$H(4)=\sum_{d\mid 4}\varphi(d)\left(\frac{4}{d}\right)^2=\varphi(1)\cdot 4^2+\varphi(2)\cdot 2^2+\varphi(4)\cdot 1^2=16+4+2=22.$$

So the nonzero part contributes

$$4\cdot 22=88,$$

and hence

$$R(5)=16+88=104.$$

A direct brute-force check over \(1\le x,y\le 20\) gives the same value.

How the Code Works

The C++, Python, and Java implementations scan the interval \([M,N]\), handle the even prime separately, and test odd candidates for primality. Only primes contribute to the final sum.

For each such prime \(p\), the implementation factors \(n=p-1\). Fast modular exponentiation is used throughout, and Pollard-Rho is used to split composite factors until the full prime factorization is known. Equal prime factors are then grouped to obtain the exponents \(e\).

Once the factorization is available, the implementation evaluates the local prime-power factors

$$q^{e-1}(q^{e+1}+q^e-1)\pmod{993353399}$$

and multiplies them to obtain \(H(n)\) modulo \(993353399\). Finally it adds

$$n^2+nH(n)\pmod{993353399}$$

to the running interval sum. The program never searches over \((x,y)\); all counting is done through the closed formula above.

Complexity Analysis

If \(W=N-M+1\), the interval scan examines \(O(W)\) candidates, skipping even numbers after \(2\). For each prime \(p\), the dominant cost is factoring \(p-1\); once the factorization is known, evaluating the product formula only needs a number of modular multiplications proportional to the number of distinct prime factors. In practice the method is fast because it replaces a quadratic search over pairs by arithmetic on the much smaller factorization of \(p-1\). Memory usage per prime stays small, essentially the temporary factor list and recursion stack.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=801
  2. Primitive roots: Wikipedia - Primitive root modulo n
  3. Cyclic groups: Wikipedia - Cyclic group
  4. Chinese remainder theorem: Wikipedia - Chinese remainder theorem
  5. Euler's totient function: Wikipedia - Euler's totient function
  6. Miller-Rabin primality test: Wikipedia - Miller-Rabin primality test
  7. Pollard's rho algorithm: Wikipedia - Pollard's rho algorithm

Problem 801 source code

C++

#include <algorithm>
#include <cassert>
#include <chrono>
#include <cstdint>
#include <iostream>
#include <map>
#include <numeric>
#include <random>
#include <vector>

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

static constexpr i64 MOD = 993'353'399LL;

static inline u64 mul_mod_u64(u64 a, u64 b, u64 mod) {
    return static_cast<u64>((static_cast<u128>(a) * static_cast<u128>(b)) % static_cast<u128>(mod));
}

static u64 pow_mod_u64(u64 a, u64 e, u64 mod) {
    u64 r = 1 % mod;
    a %= mod;
    while (e > 0) {
        if (e & 1ULL) {
            r = mul_mod_u64(r, a, mod);
        }
        a = mul_mod_u64(a, a, mod);
        e >>= 1ULL;
    }
    return r;
}

static bool is_prime_u64(u64 n) {
    if (n < 2) {
        return false;
    }
    for (u64 p : {2ULL, 3ULL, 5ULL, 7ULL, 11ULL, 13ULL, 17ULL, 19ULL, 23ULL, 29ULL, 31ULL, 37ULL}) {
        if (n == p) {
            return true;
        }
        if (n % p == 0) {
            return false;
        }
    }

    u64 d = n - 1;
    int s = 0;
    while ((d & 1ULL) == 0) {
        d >>= 1ULL;
        ++s;
    }

    static constexpr u64 WITNESSES[] = {2ULL, 325ULL, 9'375ULL, 28'178ULL, 450'775ULL, 9'780'504ULL, 1'795'265'022ULL};

    for (u64 a : WITNESSES) {
        if (a % n == 0) {
            continue;
        }
        u64 x = pow_mod_u64(a, d, n);
        if (x == 1 || x == n - 1) {
            continue;
        }
        bool comp = true;
        for (int r = 1; r < s; ++r) {
            x = mul_mod_u64(x, x, n);
            if (x == n - 1) {
                comp = false;
                break;
            }
        }
        if (comp) {
            return false;
        }
    }
    return true;
}

static u64 pollard_rho(u64 n, std::mt19937_64& rng) {
    if ((n & 1ULL) == 0ULL) {
        return 2;
    }

    std::uniform_int_distribution<u64> dist(2, n - 2);

    while (true) {
        const u64 c = dist(rng);
        u64 x = dist(rng);
        u64 y = x;
        u64 d = 1;

        auto f = [&](u64 v) {
            return (mul_mod_u64(v, v, n) + c) % n;
        };

        while (d == 1) {
            x = f(x);
            y = f(f(y));
            const u64 diff = x > y ? x - y : y - x;
            d = std::gcd(diff, n);
        }

        if (d != n) {
            return d;
        }
    }
}

static void factor_rec(u64 n, std::vector<u64>& out, std::mt19937_64& rng) {
    if (n == 1) {
        return;
    }
    if (is_prime_u64(n)) {
        out.push_back(n);
        return;
    }
    const u64 d = pollard_rho(n, rng);
    factor_rec(d, out, rng);
    factor_rec(n / d, out, rng);
}

static i64 mod_pow_i64(i64 a, i64 e) {
    i64 r = 1 % MOD;
    a %= MOD;
    if (a < 0) {
        a += MOD;
    }
    while (e > 0) {
        if (e & 1LL) {
            r = static_cast<i64>((static_cast<u128>(r) * static_cast<u128>(a)) % static_cast<u128>(MOD));
        }
        a = static_cast<i64>((static_cast<u128>(a) * static_cast<u128>(a)) % static_cast<u128>(MOD));
        e >>= 1LL;
    }
    return r;
}

static i64 f_prime_mod(u64 p, std::mt19937_64& rng) {
    const u64 n = p - 1;
    std::vector<u64> fac;
    fac.reserve(16);
    factor_rec(n, fac, rng);
    std::sort(fac.begin(), fac.end());

    std::map<u64, int> expo;
    for (u64 q : fac) {
        ++expo[q];
    }

    i64 h = 1;
    for (const auto& [q, e] : expo) {
        const i64 qm = static_cast<i64>(q % static_cast<u64>(MOD));
        const i64 q_e_minus_1 = mod_pow_i64(qm, e - 1);
        const i64 q_e = mod_pow_i64(qm, e);
        const i64 q_e_plus_1 = mod_pow_i64(qm, e + 1);

        i64 t = q_e_plus_1 + q_e - 1;
        t %= MOD;
        if (t < 0) {
            t += MOD;
        }

        const i64 term = static_cast<i64>((static_cast<u128>(q_e_minus_1) * static_cast<u128>(t)) % static_cast<u128>(MOD));
        h = static_cast<i64>((static_cast<u128>(h) * static_cast<u128>(term)) % static_cast<u128>(MOD));
    }

    const i64 nm = static_cast<i64>(n % static_cast<u64>(MOD));
    const i64 n2 = static_cast<i64>((static_cast<u128>(nm) * static_cast<u128>(nm)) % static_cast<u128>(MOD));
    const i64 nh = static_cast<i64>((static_cast<u128>(nm) * static_cast<u128>(h)) % static_cast<u128>(MOD));

    i64 f = n2 + nh;
    f %= MOD;
    if (f < 0) {
        f += MOD;
    }
    return f;
}

static std::vector<u64> primes_in_range(u64 lo, u64 hi) {
    std::vector<u64> out;
    if (hi < 2 || lo > hi) {
        return out;
    }

    if (lo <= 2 && 2 <= hi) {
        out.push_back(2);
    }

    u64 start = (lo <= 3 ? 3 : lo);
    if ((start & 1ULL) == 0ULL) {
        ++start;
    }

    for (u64 x = start; x <= hi; x += 2) {
        if (is_prime_u64(x)) {
            out.push_back(x);
        }
    }

    return out;
}

static i64 S_mod(u64 M, u64 N) {
    auto primes = primes_in_range(M, N);

    std::mt19937_64 rng(0xC0FFEEULL);
    i64 ans = 0;
    for (u64 p : primes) {
        ans += f_prime_mod(p, rng);
        ans %= MOD;
    }

    if (ans < 0) {
        ans += MOD;
    }
    return ans;
}

int main() {
    assert(S_mod(1, 100) == 7'381'000LL);
    assert(S_mod(1, 100'000) == 701'331'986LL);

    std::cout << S_mod(10'000'000'000'000'000ULL, 10'000'000'000'000'000ULL + 1'000'000ULL) << '\n';
    return 0;
}

Python

import random
import math

MOD = 993353399

def is_prime_u64(n):
    if n < 2:
        return False
    for p in [2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37]:
        if n == p:
            return True
        if n % p == 0:
            return False
            
    d = n - 1
    s = 0
    while d % 2 == 0:
        d >>= 1
        s += 1
        
    witnesses = [2, 325, 9375, 28178, 450775, 9780504, 1795265022]
    for a in witnesses:
        if a % n == 0:
            continue
        x = pow(a, d, n)
        if x == 1 or x == n - 1:
            continue
        comp = True
        for _ in range(1, s):
            x = (x * x) % n
            if x == n - 1:
                comp = False
                break
        if comp:
            return False
    return True

def pollard_rho(n):
    if n % 2 == 0:
        return 2
    while True:
        c = random.randint(2, n - 2)
        x = random.randint(2, n - 2)
        y = x
        d = 1
        
        f = lambda v: (v * v + c) % n
        
        while d == 1:
            x = f(x)
            y = f(f(y))
            diff = x - y if x > y else y - x
            d = math.gcd(diff, n)
            
        if d != n:
            return d

def factor_rec(n, out):
    if n == 1:
        return
    if is_prime_u64(n):
        out.append(n)
        return
    d = pollard_rho(n)
    factor_rec(d, out)
    factor_rec(n // d, out)

def f_prime_mod(p):
    n = p - 1
    fac = []
    factor_rec(n, fac)
    fac.sort()
    
    expo = {}
    for q in fac:
        expo[q] = expo.get(q, 0) + 1
        
    h = 1
    for q, e in expo.items():
        qm = q % MOD
        q_e_minus_1 = pow(qm, e - 1, MOD)
        q_e = pow(qm, e, MOD)
        q_e_plus_1 = pow(qm, e + 1, MOD)
        
        t = (q_e_plus_1 + q_e - 1) % MOD
        if t < 0:
            t += MOD
            
        term = (q_e_minus_1 * t) % MOD
        h = (h * term) % MOD
        
    nm = n % MOD
    n2 = (nm * nm) % MOD
    nh = (nm * h) % MOD
    
    f = (n2 + nh) % MOD
    if f < 0:
        f += MOD
    return f

def primes_in_range(lo, hi):
    out = []
    if hi < 2 or lo > hi:
        return out
    if lo <= 2 <= hi:
        out.append(2)
        
    start = 3 if lo <= 3 else lo
    if start % 2 == 0:
        start += 1
        
    for x in range(start, hi + 1, 2):
        if is_prime_u64(x):
            out.append(x)
    return out

def S_mod(M, N):
    primes = primes_in_range(M, N)
    ans = 0
    for p in primes:
        ans = (ans + f_prime_mod(p)) % MOD
    return ans

def solve():
    return str(S_mod(10000000000000000, 10000000000000000 + 1000000))

if __name__ == "__main__":
    random.seed(0xC0FFEE)
    print(solve())

Java

import java.math.BigInteger;
import java.util.ArrayList;
import java.util.Collections;
import java.util.HashMap;
import java.util.Map;
import java.util.concurrent.ThreadLocalRandom;

public class Euler801 {
    static final long MOD = 993353399L;

    static boolean isPrime(long n) {
        if (n < 2)
            return false;
        long[] smallPrimes = { 2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37 };
        for (long p : smallPrimes) {
            if (n == p)
                return true;
            if (n % p == 0)
                return false;
        }
        return BigInteger.valueOf(n).isProbablePrime(10);
    }

    static long pollardRho(long n) {
        if (n % 2 == 0)
            return 2;
        BigInteger bn = BigInteger.valueOf(n);
        ThreadLocalRandom rng = ThreadLocalRandom.current();
        while (true) {
            BigInteger c = BigInteger.valueOf(rng.nextLong(n - 3) + 2);
            BigInteger x = BigInteger.valueOf(rng.nextLong(n - 3) + 2);
            BigInteger y = x;
            long d = 1;

            while (d == 1) {
                x = x.multiply(x).add(c).mod(bn);
                y = y.multiply(y).add(c).mod(bn);
                y = y.multiply(y).add(c).mod(bn);
                BigInteger diff = x.subtract(y).abs();
                d = diff.gcd(bn).longValue();
            }
            if (d != n)
                return d;
        }
    }

    static void factorRec(long n, ArrayList<Long> out) {
        if (n == 1)
            return;
        if (isPrime(n)) {
            out.add(n);
            return;
        }
        long d = pollardRho(n);
        factorRec(d, out);
        factorRec(n / d, out);
    }

    static long modPow(long a, long e) {
        long r = 1 % MOD;
        a %= MOD;
        if (a < 0)
            a += MOD;
        while (e > 0) {
            if ((e & 1L) == 1) {
                r = (r * a) % MOD;
            }
            a = (a * a) % MOD;
            e >>= 1L;
        }
        return r;
    }

    static long fPrimeMod(long p) {
        long n = p - 1;
        ArrayList<Long> fac = new ArrayList<>();
        factorRec(n, fac);
        Collections.sort(fac);

        HashMap<Long, Integer> expo = new HashMap<>();
        for (long q : fac) {
            expo.put(q, expo.getOrDefault(q, 0) + 1);
        }

        long h = 1;
        for (Map.Entry<Long, Integer> entry : expo.entrySet()) {
            long q = entry.getKey();
            int e = entry.getValue();

            long qm = q % MOD;
            long qEMinus1 = modPow(qm, e - 1);
            long qE = modPow(qm, e);
            long qEPlus1 = modPow(qm, e + 1);

            long t = (qEPlus1 + qE - 1) % MOD;
            if (t < 0)
                t += MOD;

            long term = (qEMinus1 * t) % MOD;
            h = (h * term) % MOD;
        }

        long nm = n % MOD;
        long n2 = (nm * nm) % MOD;
        long nh = (nm * h) % MOD;

        long f = (n2 + nh) % MOD;
        if (f < 0)
            f += MOD;
        return f;
    }

    static ArrayList<Long> primesInRange(long lo, long hi) {
        ArrayList<Long> out = new ArrayList<>();
        if (hi < 2 || lo > hi)
            return out;
        if (lo <= 2 && 2 <= hi)
            out.add(2L);

        long start = lo <= 3 ? 3 : lo;
        if ((start & 1L) == 0L)
            start++;

        for (long x = start; x <= hi; x += 2) {
            if (isPrime(x))
                out.add(x);
        }
        return out;
    }

    static long SMod(long M, long N) {
        ArrayList<Long> primes = primesInRange(M, N);
        long ans = 0;
        for (long p : primes) {
            ans = (ans + fPrimeMod(p)) % MOD;
        }
        if (ans < 0)
            ans += MOD;
        return ans;
    }

    public static String solve() {
        return Long.toString(SMod(10000000000000000L, 10000000000000000L + 1000000L));
    }

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