Problem 977: Iterated Functions
View on Project EulerProject 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
- Problem page: https://projecteuler.net/problem=977
- Iterated function: Wikipedia - Iterated function
- Functional graph: Wikipedia - Functional graph
- Eventually periodic sequence: Wikipedia - Eventually periodic points
- Floor and ceiling functions: Wikipedia - Floor and ceiling functions
- Geometric series: Wikipedia - Geometric series
- 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());
}
}