Problem 977: Iterated Functions

View on Project Euler

Project Euler Problem 977 Solution

EulerSolve provides an optimized solution for Project Euler Problem 977, Iterated Functions, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary For a fixed \(n\), we count sequences \((a_1,a_2,\dots,a_n)\) with each \(a_i \in \{1,2,\dots,n\}\) and with the iterated-function condition $$a_{i+1}=a_{a_i}\qquad (1 \le i < n).$$ Project Euler 977 asks for this count when \(n=10^6\), reported modulo \(10^9+7\). The brute-force search space has size \(n^n\), so the only viable route is to understand the structure forced by the recurrence and then sum the resulting combinatorial classes in closed form. Mathematical Approach The key idea is that the recurrence does not describe arbitrary sequences: it describes the forward orbit of a single self-map, and every valid sequence eventually becomes periodic. The solution classifies sequences by the first place where that periodic regime begins. From the recurrence to an orbit Define a function \(f:\{1,\dots,n\}\to\{1,\dots,n\}\) by \(f(i)=a_i\). Then $$a_{i+1}=a_{a_i}=f(a_i)=f(f(i)).$$ Inductively this gives $$a_i=f^i(1)\qquad (i \ge 1),$$ so the sequence is exactly the orbit of 1 under repeated application of \(f\). More generally, the orbit of index \(m\) is just a shifted suffix: $$f^t(m)=a_{m+t-1}\qquad (t \ge 1).$$ This has an important consequence: if two positions carry the same value, then their entire future tails coincide. In particular, once a suffix starts repeating with some period, everything after that point is rigid....

Detailed mathematical approach

Problem Summary

For a fixed \(n\), we count sequences \((a_1,a_2,\dots,a_n)\) with each \(a_i \in \{1,2,\dots,n\}\) and with the iterated-function condition

$$a_{i+1}=a_{a_i}\qquad (1 \le i < n).$$

Project Euler 977 asks for this count when \(n=10^6\), reported modulo \(10^9+7\). The brute-force search space has size \(n^n\), so the only viable route is to understand the structure forced by the recurrence and then sum the resulting combinatorial classes in closed form.

Mathematical Approach

The key idea is that the recurrence does not describe arbitrary sequences: it describes the forward orbit of a single self-map, and every valid sequence eventually becomes periodic. The solution classifies sequences by the first place where that periodic regime begins.

From the recurrence to an orbit

Define a function \(f:\{1,\dots,n\}\to\{1,\dots,n\}\) by \(f(i)=a_i\). Then

$$a_{i+1}=a_{a_i}=f(a_i)=f(f(i)).$$

Inductively this gives

$$a_i=f^i(1)\qquad (i \ge 1),$$

so the sequence is exactly the orbit of 1 under repeated application of \(f\). More generally, the orbit of index \(m\) is just a shifted suffix:

$$f^t(m)=a_{m+t-1}\qquad (t \ge 1).$$

This has an important consequence: if two positions carry the same value, then their entire future tails coincide. In particular, once a suffix starts repeating with some period, everything after that point is rigid.

Classify by the first periodic suffix

Let \(s\) be the first index such that the suffix

$$a_s,a_{s+1},\dots,a_n$$

is periodic from its first term. Write

$$N=n-s+1$$

for the suffix length, and let its exact period be \(l\). Then the suffix positions split into residue classes modulo \(l\):

$$S_u=\{\,s+u-1+ml : m \ge 0,\ s+u-1+ml \le n\,\}\qquad (1 \le u \le l).$$

If \(N=ql+r\) with \(0 \le r < l\), then the first \(r\) classes have size \(q+1\) and the remaining \(l-r\) classes have size \(q\).

Count one periodic suffix

Because the suffix has period \(l\), every position in the same class \(S_u\) carries the same value; call it \(c_u\). The recurrence forces a simple rule:

$$c_u \in S_{u+1}\quad (1 \le u < l),\qquad c_l \in S_1.$$

So choosing the suffix means choosing one element from the next residue class for each \(u\). The number of choices is therefore

$$P(N,l)=\prod_{u=1}^{l}|S_u|=q^{\,l-r}(q+1)^r,$$

where

$$q=\left\lfloor \frac{N}{l}\right\rfloor,\qquad r=N \bmod l.$$

This is the basic factor that appears everywhere in the implementations.

Attach the non-periodic prefix

If the periodic part starts at the very beginning, so \(s=1\) and \(N=n\), there is no prefix to attach. Those sequences contribute

$$A(n)=\sum_{l=1}^{n} P(n,l).$$

Now assume \(s>1\). Then \(a_{s-1}\) must be chosen so that

$$a_s=a_{a_{s-1}}=c_1.$$

The positions whose value is \(c_1\) are exactly the elements of \(S_1\), so \(a_{s-1}\) must lie in \(S_1\). But one choice is forbidden: the special element \(c_l \in S_1\). If we took \(a_{s-1}=c_l\), then the same \(l\)-periodic pattern would already start at position \(s-1\), contradicting the minimality of \(s\).

Therefore the number of admissible attachments is

$$|S_1|-1=\left\lceil \frac{N}{l}\right\rceil - 1,$$

which is \(q-1\) when \(r=0\) and \(q\) when \(r>0\). Once \(a_{s-1}\) is fixed, every earlier position must be the tautological forward link

$$a_i=i+1\qquad (1 \le i \le s-2),$$

because any earlier nontrivial copy would make the periodic regime begin even sooner.

So the full count is

$$F(n)=\sum_{l=1}^{n} P(n,l)+\sum_{N=1}^{n-1}\sum_{l=1}^{N}\left(\left\lceil \frac{N}{l}\right\rceil-1\right)P(N,l).$$

Worked Example: \(n=7\), suffix length \(N=4\), period \(l=2\)

Take \(s=4\), so the periodic suffix is \(a_4,a_5,a_6,a_7\). The residue classes are

$$S_1=\{4,6\},\qquad S_2=\{5,7\}.$$

To build a 2-periodic suffix, choose \(c_1 \in S_2\) and \(c_2 \in S_1\). There are

$$P(4,2)=2\cdot 2=4$$

choices. For example, \(c_1=5\) and \(c_2=6\) give the suffix

$$5,6,5,6.$$

The previous term \(a_3\) must lie in \(S_1\), but it cannot equal \(c_2=6\), otherwise the 2-periodic pattern would already begin at position 3. So \(a_3=4\) is forced, and then the earlier prefix is the rigid chain \(a_1=2\), \(a_2=3\). One valid sequence is therefore

$$ (2,3,4,5,6,5,6). $$

All four suffix choices work in the same way, so this class contributes

$$\left(\left\lceil \frac{4}{2}\right\rceil-1\right)P(4,2)=1 \cdot 4=4$$

sequences, exactly as the formula predicts.

Regroup the double sum by quotient blocks

The direct formula above is already correct, and the slower validation routines evaluate it exactly. The fast solver reorganizes the second double sum. For fixed \(l\), write

$$N=ql+r,\qquad 0 \le r < l,$$

so \(q=\lfloor N/l \rfloor\). Then

$$P(N,l)=q^{\,l-r}(q+1)^r.$$

When \(N\) runs through the block \(ql,ql+1,\dots,(q+1)l-1\), the attachment factor is \(q-1\) at \(r=0\) and \(q\) for \(r \ge 1\). If

$$R=\min(l-1,n-1-ql),$$

the whole block contributes

$$B_{l,q}=(q-1)q^l+\sum_{r=1}^{R} q^{\,l+1-r}(q+1)^r.$$

The inner sum is telescoping:

$$\sum_{r=1}^{R} q^{\,l+1-r}(q+1)^r=q^{\,l+1-R}(q+1)^{R+1}-q^{\,l+1}(q+1).$$

This is exactly the closed form used by the production implementations.

How the Code Works

Power-table precomputation

The C++, Python, and Java implementations precompute all powers that can appear later. For a base \(b\), the largest exponent ever needed is at most

$$\left\lfloor \frac{n-1}{b-1}\right\rfloor + 1,$$

so every required value \(b^e \bmod (10^9+7)\) can be read in \(O(1)\) time from a packed table. This avoids an enormous number of repeated modular exponentiations inside the main summation.

Two layers of counting

The implementations first evaluate

$$A(n)=\sum_{l=1}^{n}P(n,l),$$

which counts sequences whose periodic regime begins at position 1. They then add the correction term for later starts, but not as a naive triple loop over \((N,l,r)\). Instead, for each \(l\) they iterate over quotient plateaus \(q=\lfloor N/l \rfloor\), compute the block tail length \(R\), and add the closed block sum \(B_{l,q}\). That is why the code matches the mathematical formula but runs far faster.

Validation and parallel execution

Each implementation also contains small checkpoints: exhaustive enumeration confirms that the count for \(n=7\) is 174, and a direct evaluation of the ungathered double sum confirms that the count for \(n=100\) is 305741269. After that, the large computation for \(n=10^6\) uses the precomputed power table and the regrouped block formulas. The C++ and Java implementations split the \(l\)-range across worker threads, while the Python implementation performs the same arithmetic serially.

Complexity Analysis

The dominant work is not exponential anymore. Building the packed power tables costs

$$\sum_{b=2}^{n+1} O\!\left(\frac{n}{b-1}\right)=O(n \log n),$$

and the block summation for the correction term has the same order because

$$\sum_{l=1}^{n-1}\left\lfloor\frac{n-1}{l}\right\rfloor = O(n \log n).$$

So the fast method runs in \(O(n \log n)\) time and uses \(O(n \log n)\) memory for the power tables. In practice it is efficient because every block contribution is reduced to a handful of modular multiplications and table lookups.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=977
  2. Iterated function: Wikipedia - Iterated function
  3. Functional graph: Wikipedia - Functional graph
  4. Eventually periodic sequence: Wikipedia - Eventually periodic points
  5. Floor and ceiling functions: Wikipedia - Floor and ceiling functions
  6. Geometric series: Wikipedia - Geometric series
  7. Modular arithmetic: Wikipedia - Modular arithmetic

Problem 977 source code

C++

#include <algorithm>
#include <array>
#include <cassert>
#include <cmath>
#include <cstdint>
#include <cstdlib>
#include <future>
#include <iomanip>
#include <iostream>
#include <limits>
#include <map>
#include <numeric>
#include <queue>
#include <set>
#include <string>
#include <thread>
#include <tuple>
#include <unordered_map>
#include <unordered_set>
#include <utility>
#include <vector>
using namespace std;

static constexpr int MOD = 1'000'000'007;

static inline int addmod(int a, int b) {
    int s = a + b;
    if (s >= MOD) s -= MOD;
    return s;
}
static inline int submod(int a, int b) {
    int s = a - b;
    if (s < 0) s += MOD;
    return s;
}
static inline int mulmod(long long a, long long b) {
    return (int)((a * b) % MOD);
}

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

// Precomputed powers for bases 2..N+1, packed into one array.
static vector<int> pow_all;
static vector<int> pow_offset;
static vector<int> pow_len;

static void build_powers(int n) {
    int max_base = n + 1;
    pow_offset.assign(max_base + 1, 0);
    pow_len.assign(max_base + 1, 0);

    long long total_len = 0;
    for (int b = 2; b <= max_base; b++) {
        int max_exp = (n - 1) / (b - 1) + 1;
        pow_len[b] = max_exp + 1;
        pow_offset[b] = (int)total_len;
        total_len += pow_len[b];
    }

    pow_all.assign((size_t)total_len, 0);
    for (int b = 2; b <= max_base; b++) {
        int off = pow_offset[b];
        int len = pow_len[b];
        pow_all[off] = 1;
        for (int e = 1; e < len; e++) {
            pow_all[off + e] = mulmod(pow_all[off + e - 1], b);
        }
    }
}

static inline int pow_base(int base, int exp) {
    if (exp == 0 || base == 1) return 1;
    return pow_all[pow_offset[base] + exp];
}

// Slow O(n^2) solver for validation.
static int solve_slow(int n) {
    long long total = 0;
    for (int t = 0; t < n; t++) {
        int N = n - t;
        for (int l = 1; l <= N; l++) {
            int q = N / l;
            int r = N % l;
            long long P = (modpow(q, l - r) * modpow(q + 1, r)) % MOD;
            if (t == 0) {
                total += P;
            } else {
                total += ((r == 0 ? q - 1 : q) * P) % MOD;
            }
            if (total >= (1LL << 62)) total %= MOD;
        }
    }
    return (int)(total % MOD);
}

// Brute force by enumerating sequences a_1..a_n and checking a_{k+1} = a_{a_k}.
static long long brute_count(int n) {
    vector<int> a(n + 1, 0);
    long long cnt = 0;
    function<void(int)> dfs = [&](int idx) {
        if (idx > n) {
            for (int k = 1; k < n; k++) {
                if (a[k + 1] != a[a[k]]) return;
            }
            cnt++;
            return;
        }
        for (int v = 1; v <= n; v++) {
            a[idx] = v;
            dfs(idx + 1);
        }
    };
    dfs(1);
    return cnt;
}

static int compute_B_range(int n, int l_start, int l_end) {
    long long sum = 0;
    int n1 = n - 1;
    for (int l = l_start; l < l_end; l++) {
        int max_q = n1 / l;
        for (int q = 1; q <= max_q; q++) {
            int N0 = q * l;
            int N1 = (q + 1) * l - 1;
            if (N1 > n1) N1 = n1;
            int R = N1 - N0;

            int powqL = pow_base(q, l);
            int term0 = mulmod(q - 1, powqL);
            int sum_block = term0;

            if (R >= 1) {
                int powqL1 = pow_base(q, l + 1);
                int powqL1minusR = pow_base(q, l + 1 - R);
                int powq1R1 = pow_base(q + 1, R + 1);
                int sum_r = submod(mulmod(powqL1minusR, powq1R1),
                                   mulmod(powqL1, q + 1));
                sum_block = addmod(sum_block, sum_r);
            }

            sum += sum_block;
            if (sum >= (1LL << 62)) sum %= MOD;
        }
    }
    return (int)(sum % MOD);
}

static int solve_fast(int n, int threads) {
    if (n <= 0) return 0;

    int sumA = 0;
    for (int l = 1; l <= n; l++) {
        int q = n / l;
        int r = n % l;
        int P = mulmod(pow_base(q, l - r), pow_base(q + 1, r));
        sumA = addmod(sumA, P);
    }

    if (n == 1) return sumA;

    int n1 = n - 1;
    if (threads <= 1 || n1 < 50'000) {
        int sumB = compute_B_range(n, 1, n);
        return addmod(sumA, sumB);
    }

    threads = min(threads, n1);
    vector<int> partial(threads, 0);
    vector<thread> th;
    th.reserve(threads);

    for (int t = 0; t < threads; t++) {
        int L = 1 + (long long)n1 * t / threads;
        int R = 1 + (long long)n1 * (t + 1) / threads;
        th.emplace_back([&, t, L, R]() {
            partial[t] = compute_B_range(n, L, R);
        });
    }
    for (auto& tt : th) tt.join();

    int sumB = 0;
    for (int v : partial) sumB = addmod(sumB, v);
    return addmod(sumA, sumB);
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    const int N = 1'000'000;

    if (brute_count(7) != 174) {
        cerr << "[FATAL] brute_count(7) validation failed.\n";
        return 1;
    }
    if (solve_slow(100) != 305741269) {
        cerr << "[FATAL] solve_slow(100) validation failed.\n";
        return 1;
    }

    build_powers(N);

    if (solve_fast(7, 1) != 174) {
        cerr << "[FATAL] solve_fast(7) validation failed.\n";
        return 1;
    }

    int threads = (int)thread::hardware_concurrency();
    if (threads <= 0) threads = 1;

    int ans = solve_fast(N, threads);
    cout << ans << "\n";
    return 0;
}

Python

import sys
import multiprocessing

MOD = 1_000_000_007

def modpow(a, e):
    return pow(a, e, MOD)

def pow_base(b, e):
    if e == 0 or b == 1:
        return 1
    return pow_all[pow_offset[b] + e]

pow_all = []
pow_offset = []

def build_powers(n):
    global pow_all, pow_offset
    max_base = n + 1
    pow_offset = [0] * (max_base + 1)
    pow_len = [0] * (max_base + 1)
    
    total_len = 0
    for b in range(2, max_base + 1):
        max_exp = (n - 1) // (b - 1) + 1
        pow_len[b] = max_exp + 1
        pow_offset[b] = total_len
        total_len += pow_len[b]
        
    pow_all = [0] * total_len
    for b in range(2, max_base + 1):
        off = pow_offset[b]
        length = pow_len[b]
        pow_all[off] = 1
        for e in range(1, length):
            pow_all[off + e] = (pow_all[off + e - 1] * b) % MOD

def compute_B_range(n, l_start, l_end):
    sum_val = 0
    n1 = n - 1
    for l in range(l_start, l_end):
        max_q = n1 // l
        for q in range(1, max_q + 1):
            N0 = q * l
            N1 = min((q + 1) * l - 1, n1)
            R = N1 - N0
            
            powqL = pow_base(q, l)
            term0 = ((q - 1) * powqL) % MOD
            sum_block = term0
            
            if R >= 1:
                powqL1 = pow_base(q, l + 1)
                powqL1minusR = pow_base(q, l + 1 - R)
                powq1R1 = pow_base(q + 1, R + 1)
                sum_r = (powqL1minusR * powq1R1 - powqL1 * (q + 1)) % MOD
                if sum_r < 0:
                    sum_r += MOD
                sum_block = (sum_block + sum_r) % MOD
                
            sum_val = (sum_val + sum_block) % MOD
            
    return sum_val

def compute_B_chunk(args):
    n, l_start, l_end = args
    return compute_B_range(n, l_start, l_end)

def solve_fast(n):
    if n <= 0: return 0
    
    sumA = 0
    for l in range(1, n + 1):
        q = n // l
        r = n % l
        P = (pow_base(q, l - r) * pow_base(q + 1, r)) % MOD
        sumA = (sumA + P) % MOD
        
    if n == 1: return sumA
    
    n1 = n - 1
    
    sumB = compute_B_range(n, 1, n)
    return (sumA + sumB) % MOD

def brute_count(n):
    a = [0] * (n + 1)
    cnt = 0
    
    def dfs(idx):
        nonlocal cnt
        if idx > n:
            for k in range(1, n):
                if a[k + 1] != a[a[k]]:
                    return
            cnt += 1
            return
            
        for v in range(1, n + 1):
            a[idx] = v
            dfs(idx + 1)
            
    dfs(1)
    return cnt

def solve_slow(n):
    total = 0
    for t in range(n):
        N_val = n - t
        for l in range(1, N_val + 1):
            q = N_val // l
            r = N_val % l
            P = (modpow(q, l - r) * modpow(q + 1, r)) % MOD
            if t == 0:
                total += P
            else:
                multiplier = q - 1 if r == 0 else q
                total += (multiplier * P) % MOD
            total %= MOD
    return total

def solve():
    N = 1000000
    build_powers(N)
    return str(solve_fast(N))

def run_checkpoints():
    assert brute_count(7) == 174
    assert solve_slow(100) == 305741269
    build_powers(7)
    assert solve_fast(7) == 174

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

Java

import java.util.stream.IntStream;

public class Euler977 {
    static final long MOD = 1_000_000_007;

    static long modpow(long a, long e) {
        long r = 1 % MOD;
        a %= MOD;
        while (e > 0) {
            if ((e & 1) != 0)
                r = (r * a) % MOD;
            a = (a * a) % MOD;
            e >>= 1;
        }
        return r;
    }

    static int[] powAll;
    static int[] powOffset;

    static void buildPowers(int n) {
        int maxBase = n + 1;
        powOffset = new int[maxBase + 1];
        int[] powLen = new int[maxBase + 1];

        long totalLen = 0;
        for (int b = 2; b <= maxBase; b++) {
            int maxExp = (n - 1) / (b - 1) + 1;
            powLen[b] = maxExp + 1;
            powOffset[b] = (int) totalLen;
            totalLen += powLen[b];
        }

        powAll = new int[(int) totalLen];
        for (int b = 2; b <= maxBase; b++) {
            int off = powOffset[b];
            int len = powLen[b];
            powAll[off] = 1;
            for (int e = 1; e < len; e++) {
                long prev = powAll[off + e - 1];
                powAll[off + e] = (int) ((prev * b) % MOD);
            }
        }
    }

    static int powBase(int base, int exp) {
        if (exp == 0 || base == 1)
            return 1;
        return powAll[powOffset[base] + exp];
    }

    static long bruteCount(int n) {
        int[] a = new int[n + 1];
        return dfs(1, a, n);
    }

    static long dfs(int idx, int[] a, int n) {
        if (idx > n) {
            for (int k = 1; k < n; k++) {
                if (a[k + 1] != a[a[k]])
                    return 0;
            }
            return 1;
        }
        long cnt = 0;
        for (int v = 1; v <= n; v++) {
            a[idx] = v;
            cnt += dfs(idx + 1, a, n);
        }
        return cnt;
    }

    static long solveSlow(int n) {
        long total = 0;
        for (int t = 0; t < n; t++) {
            int N = n - t;
            for (int l = 1; l <= N; l++) {
                int q = N / l;
                int r = N % l;
                long P = (modpow(q, l - r) * modpow(q + 1, r)) % MOD;
                if (t == 0) {
                    total = (total + P) % MOD;
                } else {
                    long multiplier = (r == 0) ? (q - 1) : q;
                    total = (total + multiplier * P) % MOD;
                }
            }
        }
        return total;
    }

    static long computeBRange(int n, int lStart, int lEnd) {
        long sum = 0;
        int n1 = n - 1;
        for (int l = lStart; l < lEnd; l++) {
            int maxQ = n1 / l;
            for (int q = 1; q <= maxQ; q++) {
                int n0 = q * l;
                int n1Limit = (q + 1) * l - 1;
                if (n1Limit > n1)
                    n1Limit = n1;
                int R = n1Limit - n0;

                long powqL = powBase(q, l);
                long term0 = ((q - 1) * powqL) % MOD;
                long sumBlock = term0;

                if (R >= 1) {
                    long powqL1 = powBase(q, l + 1);
                    long powqL1minusR = powBase(q, l + 1 - R);
                    long powq1R1 = powBase(q + 1, R + 1);
                    long sumR = (powqL1minusR * powq1R1 - powqL1 * (q + 1)) % MOD;
                    if (sumR < 0)
                        sumR += MOD;
                    sumBlock = (sumBlock + sumR) % MOD;
                }

                sum = (sum + sumBlock) % MOD;
            }
        }
        return sum;
    }

    static long solveFast(int n) {
        if (n <= 0)
            return 0;

        long sumA = 0;
        for (int l = 1; l <= n; l++) {
            int q = n / l;
            int r = n % l;
            long P = ((long) powBase(q, l - r) * powBase(q + 1, r)) % MOD;
            sumA = (sumA + P) % MOD;
        }

        if (n == 1)
            return sumA;

        long sumB = IntStream.range(0, 16)
                .parallel()
                .mapToLong(t -> {
                    int L = 1 + (n - 1) * t / 16;
                    int R = 1 + (n - 1) * (t + 1) / 16;
                    return computeBRange(n, L, R);
                })
                .reduce(0, (a, b) -> (a + b) % MOD);

        return (sumA + sumB) % MOD;
    }

    public static String solve() {
        int N = 1000000;
        buildPowers(N);
        return Long.toString(solveFast(N));
    }

    public static void main(String[] args) {
        if (bruteCount(7) != 174) {
            System.out.println("Validation failed");
            return;
        }
        if (solveSlow(100) != 305741269) {
            System.out.println("Validation failed");
            return;
        }
        buildPowers(7);
        if (solveFast(7) != 174) {
            System.out.println("Validation failed");
            return;
        }

        System.out.println(solve());
    }
}