Problem 822: Square the Smallest

View on Project Euler

Project Euler Problem 822 Solution

EulerSolve provides an optimized solution for Project Euler Problem 822, Square the Smallest, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary Start with the list \(2,3,\dots,n\). Each operation finds the smallest current entry and replaces it by its square. After \(m\) operations we need the sum of all entries modulo \(1234567891\). Because \(m\) can be enormous, the solution does not simulate every step on the raw integers. Mathematical Approach For each initial base \(b\in\{2,\dots,n\}\), let \(t_b\) be the number of times that entry has been selected so far. Its current value is then $$b^{2^{t_b}}.$$ The solution separates two tasks: determining which entry is currently smallest, and evaluating the final values modulo the prime modulus. Step 1: Move the Ordering Problem into Logarithms The logarithm is strictly increasing, so comparing values is equivalent to comparing their logarithms. For the entry that started at \(b\), define $$\lambda_b=\log\left(b^{2^{t_b}}\right)=2^{t_b}\log b.$$ Choosing the smallest current number is therefore the same as choosing the smallest \(\lambda_b\). One squaring sends $$\lambda_b \longmapsto 2\lambda_b.$$ This lets the implementation maintain a min-priority structure on logarithmic keys instead of on astronomically large integers. When two logarithmic keys are equal, the values are equal as well; the implementations use a deterministic secondary order so the data structure stays stable....

Detailed mathematical approach

Problem Summary

Start with the list \(2,3,\dots,n\). Each operation finds the smallest current entry and replaces it by its square. After \(m\) operations we need the sum of all entries modulo \(1234567891\). Because \(m\) can be enormous, the solution does not simulate every step on the raw integers.

Mathematical Approach

For each initial base \(b\in\{2,\dots,n\}\), let \(t_b\) be the number of times that entry has been selected so far. Its current value is then

$$b^{2^{t_b}}.$$

The solution separates two tasks: determining which entry is currently smallest, and evaluating the final values modulo the prime modulus.

Step 1: Move the Ordering Problem into Logarithms

The logarithm is strictly increasing, so comparing values is equivalent to comparing their logarithms. For the entry that started at \(b\), define

$$\lambda_b=\log\left(b^{2^{t_b}}\right)=2^{t_b}\log b.$$

Choosing the smallest current number is therefore the same as choosing the smallest \(\lambda_b\). One squaring sends

$$\lambda_b \longmapsto 2\lambda_b.$$

This lets the implementation maintain a min-priority structure on logarithmic keys instead of on astronomically large integers. When two logarithmic keys are equal, the values are equal as well; the implementations use a deterministic secondary order so the data structure stays stable.

Step 2: Simulate Explicitly While the Smallest Entry Can Stay Competitive

Write the ordered logarithms at some moment as

$$\lambda_1\le \lambda_2\le \cdots \le \lambda_L,\qquad L=n-1.$$

Suppose the smallest one satisfies

$$2\lambda_1\le \lambda_L.$$

After squaring that smallest entry, its new logarithm is still no larger than the current maximum. So it may remain among the small values and can affect the immediate future ordering in a nontrivial way. During this phase the algorithm performs exact heap updates one operation at a time.

The number of such explicit updates is usually far smaller than \(m\), because each update doubles one logarithm and quickly pushes that entry upward.

Step 3: Prove That the Balanced Regime Becomes Cyclic

The explicit phase stops once

$$2\lambda_1>\lambda_L.$$

Now squaring the smallest entry sends it strictly past every untouched entry, because its new logarithm exceeds the current maximum. One step therefore transforms the sorted list into

$$\lambda_2,\lambda_3,\dots,\lambda_L,2\lambda_1.$$

The relative order of the untouched entries does not change, and the updated entry moves to the end. Repeating the same argument shows that the next selected entry must be \(\lambda_2\), then \(\lambda_3\), and so on. After one full round of \(L\) operations, every entry has been squared exactly once and the ordered list becomes

$$2\lambda_1,2\lambda_2,\dots,2\lambda_L.$$

The key point is that the same inequality still holds after scaling the entire list by \(2\):

$$2(2\lambda_1)>2\lambda_L \iff 2\lambda_1>\lambda_L.$$

So once this regime begins, the process continues in identical cycles forever.

Step 4: Distribute the Remaining Operations in Bulk

Let \(r_{\mathrm{rem}}\) be the number of operations left after the explicit phase. Since the balanced regime proceeds in complete rounds over the \(L\) entries, write

$$r_{\mathrm{rem}}=qL+s,\qquad 0\le s<L.$$

Then every entry is squared \(q\) more times, and the first \(s\) entries in the current sorted logarithmic order are squared once additional time. Equivalently, if the balanced-order list is

$$\lambda_1\le \lambda_2\le \cdots \le \lambda_L,$$

then entries \(1,2,\dots,s\) receive \(q+1\) further squarings, while entries \(s+1,\dots,L\) receive \(q\). This replaces an astronomically large tail of the process by one division and one remainder.

Step 5: Recover the Final Values Modulo the Prime

If the current residue modulo \(M=1234567891\) is \(v\), then applying \(t\) further squarings produces

$$v^{2^t}\pmod{M}.$$

Therefore the bulk phase only needs the exponents \(2^q\) and \(2^{q+1}\). Because \(M\) is prime and the initial bases \(2,3,\dots,n\) are all nonzero modulo \(M\), every later residue is also nonzero modulo \(M\). Fermat's little theorem therefore reduces exponents modulo \(M-1\):

$$a^k\equiv a^{\,k\bmod (M-1)}\pmod{M}\qquad (a\not\equiv 0\pmod{M}).$$

So the implementation computes \(2^q\bmod (M-1)\), doubles it to obtain \(2^{q+1}\bmod (M-1)\), and then evaluates the required modular powers for the final sum.

Worked Example: \((n,m)=(5,3)\)

The initial list is \(2,3,4,5\). The first three operations are easy to follow directly:

$$[2,3,4,5]\to [4,3,4,5]\to [4,9,4,5]\to [16,9,4,5].$$

Therefore the required sum is

$$16+9+4+5=34,$$

which matches the checkpoint used by the implementations.

To see the bulk rule in isolation, suppose the ordered logarithms after the explicit phase are \(1.40,1.55,1.80,2.10\). Since \(2\cdot 1.40=2.80>2.10\), the next four selections are forced to occur in that same order. If seven operations remain, then

$$7=1\cdot 4+3,$$

so every entry is squared once more, and the first three receive one extra squaring.

How the Code Works

The C++, Python, and Java implementations store two pieces of information for each entry: a logarithmic key used only for ordering, and the current residue modulo \(1234567891\). They initialize a min-priority structure with the values \(2,3,\dots,n\) and also track the largest current logarithm.

While doubling the smallest logarithm does not jump beyond the current maximum, the implementation performs one exact update: remove the minimum, double its logarithmic key, square its residue modulo the prime, reinsert it, and refresh the current maximum if needed.

Once the balanced inequality \(2\lambda_{\min}>\lambda_{\max}\) holds, the remaining items are extracted and sorted by the same ordering rule. The implementation then computes the quotient \(q\) and remainder \(s\) of the remaining step count by \(L=n-1\), converts those counts into the exponents \(2^q\) and \(2^{q+1}\) modulo \(M-1\), and applies fast modular exponentiation to each stored residue. The first \(s\) sorted entries receive the larger exponent, the rest receive the smaller one, and the resulting residues are summed modulo \(M\).

Complexity Analysis

Let \(L=n-1\), and let \(u\) be the number of explicit priority-queue updates performed before the balanced regime begins. The implementations insert \(L\) initial items into the priority structure, which costs \(O(L\log L)\). The explicit phase then costs \(O(u\log L)\), because each update removes and reinserts one element.

After that, the remaining work is one sort of \(L\) items, one exponentiation to compute \(2^q\bmod(M-1)\), and \(L\) modular exponentiations for the final residues. This gives overall running time

$$O(u\log L+L\log L+L\log M+\log q),$$

which simplifies to \(O(u\log L+L\log L)\) when the modulus is treated as fixed. The memory usage is \(O(L)\). The important point is that the method avoids any linear dependence on the huge value of \(m\) after the short explicit phase.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=822
  2. Logarithm: Wikipedia — Logarithm
  3. Heap data structure: Wikipedia — Heap (data structure)
  4. Modular exponentiation: Wikipedia — Modular exponentiation
  5. Fermat's little theorem: Wikipedia — Fermat's little theorem

Problem 822 source code

C++

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

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

static constexpr u64 kMod = 1'234'567'891ULL;

struct Node {
    long double logv;
    u64 modv;
    int id;
};

struct MinCmp {
    bool operator()(const Node& a, const Node& b) const {
        if (a.logv != b.logv) return a.logv > b.logv;
        return a.id > b.id;
    }
};

static u64 mod_pow(u64 base, u64 exp, u64 mod) {
    u64 res = 1 % mod;
    base %= mod;
    while (exp > 0) {
        if (exp & 1ULL) {
            res = static_cast<u64>((static_cast<u128>(res) * base) % mod);
        }
        base = static_cast<u64>((static_cast<u128>(base) * base) % mod);
        exp >>= 1ULL;
    }
    return res;
}

static u64 S(u64 n, u64 m) {
    const int len = static_cast<int>(n - 1);

    std::priority_queue<Node, std::vector<Node>, MinCmp> pq;
    std::vector<Node> cur(len);
    long double max_log = 0.0L;

    for (int i = 0; i < len; ++i) {
        const u64 v = static_cast<u64>(i + 2);
        cur[i] = {std::log(static_cast<long double>(v)), v % kMod, i};
        pq.push(cur[i]);
        if (cur[i].logv > max_log) {
            max_log = cur[i].logv;
        }
    }

    u64 steps = 0;
    while (steps < m) {
        Node x = pq.top();
        if (2.0L * x.logv > max_log) {
            break;
        }
        pq.pop();
        x.logv *= 2.0L;
        x.modv = static_cast<u64>((static_cast<u128>(x.modv) * x.modv) % kMod);
        pq.push(x);
        if (x.logv > max_log) {
            max_log = x.logv;
        }
        ++steps;
    }

    std::vector<Node> arr;
    arr.reserve(len);
    while (!pq.empty()) {
        arr.push_back(pq.top());
        pq.pop();
    }

    if (steps == m) {
        u64 ans = 0;
        for (const auto& x : arr) {
            ans += x.modv;
            if (ans >= kMod) {
                ans -= kMod;
            }
        }
        return ans;
    }

    std::sort(arr.begin(), arr.end(), [](const Node& a, const Node& b) {
        if (a.logv != b.logv) return a.logv < b.logv;
        return a.id < b.id;
    });

    const u64 rem = m - steps;
    const u64 q = rem / static_cast<u64>(len);
    const u64 r = rem % static_cast<u64>(len);

    const u64 exp_q = mod_pow(2, q, kMod - 1);
    const u64 exp_q1 = static_cast<u64>((static_cast<u128>(exp_q) * 2ULL) % (kMod - 1));

    u64 ans = 0;
    for (int i = 0; i < len; ++i) {
        const u64 e = (static_cast<u64>(i) < r) ? exp_q1 : exp_q;
        const u64 add = mod_pow(arr[i].modv, e, kMod);
        ans += add;
        if (ans >= kMod) {
            ans -= kMod;
        }
    }

    return ans;
}

int main() {
    assert(S(5, 3) == 34);
    assert(S(10, 100) == 845339386ULL);
    std::cout << S(10'000, 10'000'000'000'000'000ULL) << '\n';
    return 0;
}

Python

import math
import heapq

kMod = 1234567891

class Node:
    def __init__(self, logv, modv, id):
        self.logv = logv
        self.modv = modv
        self.id = id
        
    def __lt__(self, other):
        if self.logv != other.logv:
            return self.logv < other.logv
        return self.id < other.id

def S(n, m):
    len_n = n - 1
    pq = []
    max_log = 0.0
    
    for i in range(len_n):
        v = i + 2
        logv = math.log(v)
        node = Node(logv, v % kMod, i)
        heapq.heappush(pq, node)
        if logv > max_log:
            max_log = logv
            
    steps = 0
    while steps < m:
        x = pq[0]
        if 2.0 * x.logv > max_log:
            break
        heapq.heappop(pq)
        x.logv *= 2.0
        x.modv = (x.modv * x.modv) % kMod
        heapq.heappush(pq, x)
        if x.logv > max_log:
            max_log = x.logv
        steps += 1
        
    if steps == m:
        ans = 0
        for x in pq:
            ans = (ans + x.modv) % kMod
        return ans
        
    arr = list(pq)
    arr.sort(key=lambda x: (x.logv, x.id))
    
    rem = m - steps
    q = rem // len_n
    r = rem % len_n
    
    exp_q = pow(2, q, kMod - 1)
    exp_q1 = (exp_q * 2) % (kMod - 1)
    
    ans = 0
    for i in range(len_n):
        e = exp_q1 if i < r else exp_q
        add = pow(arr[i].modv, e, kMod)
        ans = (ans + add) % kMod
        
    return ans

def solve():
    ans = S(10000, 10000000000000000)
    return str(ans)

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

Java

import java.util.PriorityQueue;
import java.util.Comparator;
import java.util.ArrayList;
import java.util.Collections;

public class Euler822 {

    static final long kMod = 1234567891L;

    static class Node {
        double logv;
        long modv;
        int id;

        Node(double logv, long modv, int id) {
            this.logv = logv;
            this.modv = modv;
            this.id = id;
        }
    }

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

    static long S(long n, long m) {
        int len = (int) (n - 1);

        PriorityQueue<Node> pq = new PriorityQueue<>(new Comparator<Node>() {
            @Override
            public int compare(Node a, Node b) {
                if (a.logv != b.logv) {
                    return Double.compare(a.logv, b.logv);
                }
                return Integer.compare(a.id, b.id);
            }
        });

        double maxLog = 0.0;

        for (int i = 0; i < len; ++i) {
            long v = i + 2;
            Node node = new Node(Math.log(v), v % kMod, i);
            pq.add(node);
            if (node.logv > maxLog) {
                maxLog = node.logv;
            }
        }

        long steps = 0;
        while (steps < m) {
            Node x = pq.peek();
            if (2.0 * x.logv > maxLog) {
                break;
            }
            pq.poll();
            x.logv *= 2.0;
            x.modv = (x.modv * x.modv) % kMod;
            pq.add(x);
            if (x.logv > maxLog) {
                maxLog = x.logv;
            }
            ++steps;
        }

        ArrayList<Node> arr = new ArrayList<>(pq);

        if (steps == m) {
            long ans = 0;
            for (Node x : arr) {
                ans += x.modv;
                if (ans >= kMod) {
                    ans -= kMod;
                }
            }
            return ans;
        }

        Collections.sort(arr, new Comparator<Node>() {
            @Override
            public int compare(Node a, Node b) {
                if (a.logv != b.logv) {
                    return Double.compare(a.logv, b.logv);
                }
                return Integer.compare(a.id, b.id);
            }
        });

        long rem = m - steps;
        long q = rem / len;
        long r = rem % len;

        long expQ = modPow(2, q, kMod - 1);
        long expQ1 = (expQ * 2L) % (kMod - 1);

        long ans = 0;
        for (int i = 0; i < len; ++i) {
            long e = (i < r) ? expQ1 : expQ;
            long add = modPow(arr.get(i).modv, e, kMod);
            ans += add;
            if (ans >= kMod) {
                ans -= kMod;
            }
        }

        return ans;
    }

    public static String solve() {
        return Long.toString(S(10000, 10000000000000000L));
    }

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