Problem 902: Permutation Powers

View on Project Euler

Project Euler Problem 902 Solution

EulerSolve provides an optimized solution for Project Euler Problem 902, Permutation Powers, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary Let \(m=100\) and \(n=1+2+\cdots+m=5050\). The permutation in the problem is obtained by starting with a block permutation whose disjoint cycles have lengths \(1,2,\dots,100\), and then relabeling the points by a bijection induced by multiplication by \(10^9+7\) modulo \(5050\). The task is to sum the 1-indexed lexicographic ranks of the first \(100!\) powers of that permutation and report the result modulo \(10^9+7\). A direct simulation would require handling \(100!\) different powers of a permutation on 5050 symbols, which is completely infeasible. The key observation is that lexicographic rank can be written as a weighted inversion count, and powers of a permutation behave periodically on each pair of cycles. Mathematical Approach Write the final permutation as \(\rho\), in one-line notation \((\rho(1),\rho(2),\dots,\rho(n))\). The entire solution comes from combining cycle structure, Lehmer-code weights, and a counting argument on pairs of cycles. Step 1: Identify the cycle structure The base permutation is a disjoint union of consecutive cycles of lengths \(1,2,\dots,100\). Conjugating by a relabeling permutation changes the labels written inside the cycles, but it does not change any cycle length. Therefore \(\rho\) still has exactly one cycle of each length from 1 through 100....

Detailed mathematical approach

Problem Summary

Let \(m=100\) and \(n=1+2+\cdots+m=5050\). The permutation in the problem is obtained by starting with a block permutation whose disjoint cycles have lengths \(1,2,\dots,100\), and then relabeling the points by a bijection induced by multiplication by \(10^9+7\) modulo \(5050\). The task is to sum the 1-indexed lexicographic ranks of the first \(100!\) powers of that permutation and report the result modulo \(10^9+7\).

A direct simulation would require handling \(100!\) different powers of a permutation on 5050 symbols, which is completely infeasible. The key observation is that lexicographic rank can be written as a weighted inversion count, and powers of a permutation behave periodically on each pair of cycles.

Mathematical Approach

Write the final permutation as \(\rho\), in one-line notation \((\rho(1),\rho(2),\dots,\rho(n))\). The entire solution comes from combining cycle structure, Lehmer-code weights, and a counting argument on pairs of cycles.

Step 1: Identify the cycle structure

The base permutation is a disjoint union of consecutive cycles of lengths \(1,2,\dots,100\). Conjugating by a relabeling permutation changes the labels written inside the cycles, but it does not change any cycle length. Therefore \(\rho\) still has exactly one cycle of each length from 1 through 100.

This immediately implies that the order of \(\rho\) divides \(100!\), because every cycle length divides \(100!\). So when we sum over the first \(100!\) powers, every cycle has completed an integer number of full turns, and every pair of cycles has completed an integer number of joint periods.

Step 2: Rewrite lexicographic rank as weighted inversions

For any permutation \(\alpha\) on \(\{1,\dots,n\}\), its 1-indexed lexicographic rank is

$$\operatorname{rank}(\alpha)=1+\sum_{i=1}^{n} c_i(\alpha)\,(n-i)!,$$

where

$$c_i(\alpha)=\#\{j \gt i:\alpha(j)\lt \alpha(i)\}.$$

This is exactly the factorial-number-system or Lehmer-code expansion. Summing over the first \(100!\) powers of \(\rho\) gives

$$S=\sum_{k=1}^{100!}\operatorname{rank}(\rho^k)=100!+\sum_{1\le i\lt j\le n}(n-i)!\,N_{i,j},$$

with

$$N_{i,j}=\#\{1\le k\le 100!:\rho^k(j)\lt \rho^k(i)\}.$$

So the whole problem reduces to one question: for each pair of positions \(i\lt j\), how often do their images appear in reversed numerical order as the exponent runs?

Step 3: Reduce one index pair to a pair of cycles

Suppose \(i\) belongs to a cycle

$$C=(c_0,c_1,\dots,c_{s-1}),$$

and \(j\) belongs to a cycle

$$D=(d_0,d_1,\dots,d_{t-1}).$$

If \(i\) starts at position \(u\) in \(C\) and \(j\) starts at position \(v\) in \(D\), then after \(k\) applications of \(\rho\) we have

$$\rho^k(i)=c_{u+k \bmod s},\qquad \rho^k(j)=d_{v+k \bmod t}.$$

Now define

$$g=\gcd(s,t),\qquad L=\operatorname{lcm}(s,t)=\frac{st}{g}.$$

The ordered pair of cycle positions repeats every \(L\) steps. Since \(L\mid 100!\), it is enough to count favorable exponents in one \(L\)-step period and then multiply by \(100!/L\).

Step 4: Use residue classes modulo the gcd

The relative phase between the two cycles matters only modulo \(g\). Set

$$\delta\equiv v-u \pmod g.$$

Partition each cycle according to position modulo \(g\):

$$C_r=\{c_a:a\equiv r \pmod g\},\qquad D_r=\{d_b:b\equiv r \pmod g\}.$$

During one full \(L\)-step period, the pair \((\rho^k(i),\rho^k(j))\) visits exactly the pairs of positions whose residues differ by \(\delta\). Each admissible position pair occurs once. Therefore the number of favorable exponents in one period is

$$Q_\delta(C,D)=\sum_{r=0}^{g-1}\#\{(x,y)\in C_r\times D_{r+\delta}: y\lt x\}.$$

Hence

$$N_{i,j}=Q_\delta(C,D)\cdot\frac{100!}{L}.$$

The comparison problem is now finite and static: sort each residue class once, and for every phase \(\delta\) count how many elements of the second class are smaller than each element of the first.

Step 5: Assemble the final summation

Substituting the cycle-pair count into the weighted inversion formula yields

$$\boxed{S=100!+\sum_{1\le i\lt j\le n}(n-i)!\,Q_{\delta(i,j)}(C(i),C(j))\frac{100!}{\operatorname{lcm}(|C(i)|,|C(j)|)} \pmod{10^9+7}.}$$

Here \(C(i)\) and \(C(j)\) are the cycles containing \(i\) and \(j\), and \(\delta(i,j)\) is the phase difference of their positions modulo the relevant gcd. This is exactly the quantity accumulated by the implementations.

Worked Example: One Cycle Pair

Take two sample cycles

$$C=(9,1,8,2),\qquad D=(6,3,7,4,5,10).$$

Then \(s=4\), \(t=6\), so

$$g=\gcd(4,6)=2,\qquad L=\operatorname{lcm}(4,6)=12.$$

Split them by position parity:

$$C_0=\{9,8\},\ C_1=\{1,2\},\qquad D_0=\{6,7,5\},\ D_1=\{3,4,10\}.$$

After sorting,

$$C_0=\{8,9\},\ C_1=\{1,2\},\qquad D_0=\{5,6,7\},\ D_1=\{3,4,10\}.$$

If the starting phase difference is \(\delta=1\), then

$$Q_1(C,D)=\#\{(x,y)\in C_0\times D_1:y\lt x\}+\#\{(x,y)\in C_1\times D_0:y\lt x\}=4+0=4.$$

So there are exactly 4 favorable exponents in each 12-step joint period. Across the first \(100!\) powers this becomes

$$N=4\cdot\frac{100!}{12}=\frac{100!}{3}.$$

This miniature example is the same mechanism used for every pair of indices in the full problem.

How the Code Works

The C++, Python, and Java implementations build the relabeled permutation, decompose it into cycles, and record for every value both its cycle index and its position inside that cycle. They also precompute the factorial weights \((n-i)!\) used by the lexicographic-rank formula, together with \(100! \bmod (10^9+7)\).

For each ordered pair of cycles, the implementation computes \(g=\gcd(s,t)\), \(L=\operatorname{lcm}(s,t)\), and the phase tables \(Q_\delta(C,D)\) for \(\delta=0,1,\dots,g-1\). Because the modulus is prime and \(L<10^9+7\), the factor \(100!/L\) is represented modulo \(10^9+7\) as \(100!\cdot L^{-1}\).

After that preprocessing, the final accumulation is straightforward: for each \(i<j\), look up the two cycles, recover the correct phase \(\delta\), read the precomputed table entry, multiply by \((n-i)!\) and by the modular version of \(100!/L\), and add the result to the running sum. The C++ and Java implementations parallelize this last double loop across several threads; the Python implementation performs the same arithmetic serially.

Complexity Analysis

The permutation size is fixed at \(n=5050\), so the dominant cost is the sweep over all \(\binom{5050}{2}\) index pairs in the final accumulation. In asymptotic terms, if one writes \(n=1+2+\cdots+m=\Theta(m^2)\), then the main phase is \(O(n^2)=O(m^4)\).

The cycle-pair preprocessing is smaller: there are only 100 cycles, and each pair of cycles needs a table of size \(\gcd(s,t)\) plus sorted residue classes. Memory usage is \(O(n)\) for the permutation metadata and factorial weights, plus the cycle-pair tables, which are much smaller than the final pair sweep. In practice the algorithm is fast because it replaces gigantic exponentiation over \(100!\) powers by a single precompute-and-count pass.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=902
  2. Permutation: Wikipedia - Permutation
  3. Cycle notation and decomposition: Wikipedia - Cycle notation
  4. Lehmer code: Wikipedia - Lehmer code
  5. Lexicographic order: Wikipedia - Lexicographic order
  6. Greatest common divisor: Wikipedia - Greatest common divisor
  7. Least common multiple: Wikipedia - Least common multiple

Problem 902 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 const long long MOD = 1000000007LL;
static const long long A = 1000000007LL;

struct PairData {
    int g = 1;
    long long mul = 0;
    vector<int> counts;
};

bool validate_permutation(const vector<int> &p, int n, const string &name) {
    vector<char> seen(n + 1, 0);
    for (int i = 1; i <= n; ++i) {
        int v = p[i];
        if (v < 1 || v > n) {
            cerr << "Validation failed: " << name << " has out-of-range value.\n";
            return false;
        }
        if (seen[v]) {
            cerr << "Validation failed: " << name << " has duplicate value.\n";
            return false;
        }
        seen[v] = 1;
    }
    return true;
}

vector<int> compute_counts(const vector<int> &Acyc, const vector<int> &Bcyc, int g) {
    const int s = static_cast<int>(Acyc.size());
    const int t = static_cast<int>(Bcyc.size());
    vector<vector<int>> A_r(g), B_r(g);
    for (int i = 0; i < s; ++i) {
        A_r[i % g].push_back(Acyc[i]);
    }
    for (int i = 0; i < t; ++i) {
        B_r[i % g].push_back(Bcyc[i]);
    }
    for (int r = 0; r < g; ++r) {
        sort(A_r[r].begin(), A_r[r].end());
        sort(B_r[r].begin(), B_r[r].end());
    }
    vector<vector<int>> f(g, vector<int>(g, 0));
    for (int r = 0; r < g; ++r) {
        for (int rp = 0; rp < g; ++rp) {
            int cnt = 0;
            const auto &a = A_r[r];
            const auto &b = B_r[rp];
            int j = 0;
            for (int x : a) {
                while (j < static_cast<int>(b.size()) && b[j] < x) {
                    ++j;
                }
                cnt += j;
            }
            f[r][rp] = cnt;
        }
    }
    // Count for any shift depends only on the offset modulo g.
    vector<int> counts(g, 0);
    for (int rd = 0; rd < g; ++rd) {
        int total = 0;
        for (int r = 0; r < g; ++r) {
            total += f[r][(r + rd) % g];
        }
        counts[rd] = total;
    }
    return counts;
}

int main() {
    const int m = 100;
    const int n = m * (m + 1) / 2;

    vector<int> tau(n + 1, 0), tau_inv(n + 1, 0), sigma(n + 1, 0), pi(n + 1, 0);
    for (int i = 1; i <= n; ++i) {
        tau[i] = static_cast<int>((A * i) % n) + 1;
    }
    for (int i = 1; i <= n; ++i) {
        tau_inv[tau[i]] = i;
    }
    for (int i = 1; i <= n; ++i) {
        if (tau_inv[tau[i]] != i) {
            cerr << "Validation failed: tau inverse mismatch.\n";
            return 0;
        }
    }

    for (int i = 1; i < n; ++i) {
        sigma[i] = i + 1;
    }
    int prev = 0;
    for (int k = 1; k <= m; ++k) {
        int t = k * (k + 1) / 2;
        sigma[t] = prev + 1;
        prev = t;
    }

    for (int i = 1; i <= n; ++i) {
        pi[i] = tau_inv[sigma[tau[i]]];
    }
    if (!validate_permutation(pi, n, "pi")) {
        return 0;
    }

    vector<vector<int>> cycles;
    vector<int> cycle_id(n + 1, -1), pos_in(n + 1, -1);
    vector<char> visited(n + 1, 0);
    for (int i = 1; i <= n; ++i) {
        if (visited[i]) {
            continue;
        }
        vector<int> cyc;
        int cur = i;
        while (!visited[cur]) {
            visited[cur] = 1;
            pos_in[cur] = static_cast<int>(cyc.size());
            cycle_id[cur] = static_cast<int>(cycles.size());
            cyc.push_back(cur);
            cur = pi[cur];
        }
        cycles.push_back(move(cyc));
    }

    int total_len = 0;
    vector<int> len_count(m + 1, 0);
    for (const auto &cyc : cycles) {
        total_len += static_cast<int>(cyc.size());
        if (static_cast<int>(cyc.size()) >= 1 && static_cast<int>(cyc.size()) <= m) {
            len_count[cyc.size()]++;
        }
    }
    if (total_len != n) {
        cerr << "Validation failed: cycle coverage mismatch.\n";
        return 0;
    }
    for (int len = 1; len <= m; ++len) {
        if (len_count[len] != 1) {
            cerr << "Validation failed: unexpected cycle length distribution.\n";
            return 0;
        }
    }

    vector<long long> fact(n + 1, 1), weight(n + 1, 0);
    for (int i = 1; i <= n; ++i) {
        fact[i] = (fact[i - 1] * i) % MOD;
    }
    for (int i = 1; i <= n; ++i) {
        weight[i] = fact[n - i];
    }

    long long m_fact = 1;
    for (int i = 1; i <= m; ++i) {
        m_fact = (m_fact * i) % MOD;
    }

    const int num_cycles = static_cast<int>(cycles.size());
    vector<int> cycle_len(num_cycles, 0);
    for (int i = 0; i < num_cycles; ++i) {
        cycle_len[i] = static_cast<int>(cycles[i].size());
    }

    int maxL = 1;
    for (int i = 0; i < num_cycles; ++i) {
        for (int j = 0; j < num_cycles; ++j) {
            int g = std::gcd(cycle_len[i], cycle_len[j]);
            int L = cycle_len[i] / g * cycle_len[j];
            maxL = max(maxL, L);
        }
    }
    vector<long long> inv(maxL + 1, 1);
    for (int i = 2; i <= maxL; ++i) {
        inv[i] = MOD - (MOD / i) * inv[MOD % i] % MOD;
    }

    vector<vector<PairData>> pair_data(num_cycles, vector<PairData>(num_cycles));
    for (int ci = 0; ci < num_cycles; ++ci) {
        for (int cj = 0; cj < num_cycles; ++cj) {
            int s = cycle_len[ci];
            int t = cycle_len[cj];
            int g = std::gcd(s, t);
            int L = s / g * t;
            PairData pd;
            pd.g = g;
            pd.mul = (m_fact * inv[L]) % MOD;
            pd.counts = compute_counts(cycles[ci], cycles[cj], g);
            pair_data[ci][cj] = move(pd);
        }
    }

    long long total = m_fact % MOD;
    unsigned int threads = std::thread::hardware_concurrency();
    if (threads == 0) {
        threads = 1;
    }
    if (threads > 8) {
        threads = 8;
    }
    vector<long long> partial(threads, 0);
    vector<thread> workers;

    int block = n / static_cast<int>(threads);
    int start = 1;
    for (unsigned int t = 0; t < threads; ++t) {
        int end = (t == threads - 1) ? n : (start + block - 1);
        workers.emplace_back([&, start, end, t]() {
            long long local = 0;
            for (int i = start; i <= end; ++i) {
                long long wi = weight[i];
                int ci = cycle_id[i];
                int pos_i = pos_in[i];
                for (int j = i + 1; j <= n; ++j) {
                    int cj = cycle_id[j];
                    int pos_j = pos_in[j];
                    const PairData &pd = pair_data[ci][cj];
                    int g = pd.g;
                    int r = pos_j - pos_i;
                    r %= g;
                    if (r < 0) {
                        r += g;
                    }
                    int cnt = pd.counts[r];
                    long long add = (wi * (static_cast<long long>(cnt) * pd.mul % MOD)) % MOD;
                    local += add;
                    if (local >= MOD) {
                        local -= MOD;
                    }
                }
            }
            partial[t] = local;
        });
        start = end + 1;
    }
    for (auto &th : workers) {
        th.join();
    }
    for (long long v : partial) {
        total += v;
        if (total >= MOD) {
            total -= MOD;
        }
    }

    cout << (total % MOD) << "\n";
    return 0;
}

Python

import math
from typing import List

MOD = 1000000007
A = 1000000007

def compute_counts(Acyc: List[int], Bcyc: List[int], g: int) -> List[int]:
    s = len(Acyc)
    t = len(Bcyc)
    A_r = [[] for _ in range(g)]
    B_r = [[] for _ in range(g)]
    for i in range(s):
        A_r[i % g].append(Acyc[i])
    for i in range(t):
        B_r[i % g].append(Bcyc[i])
        
    for r in range(g):
        A_r[r].sort()
        B_r[r].sort()
        
    f = [[0] * g for _ in range(g)]
    for r in range(g):
        for rp in range(g):
            cnt = 0
            a = A_r[r]
            b = B_r[rp]
            j = 0
            b_len = len(b)
            for x in a:
                while j < b_len and b[j] < x:
                    j += 1
                cnt += j
            f[r][rp] = cnt
            
    counts = [0] * g
    for rd in range(g):
        total = 0
        for r in range(g):
            total += f[r][(r + rd) % g]
        counts[rd] = total
    return counts

def solve():
    m = 100
    n = m * (m + 1) // 2
    
    tau = [0] * (n + 1)
    tau_inv = [0] * (n + 1)
    sigma = [0] * (n + 1)
    pi = [0] * (n + 1)
    
    for i in range(1, n + 1):
        tau[i] = ((A * i) % n) + 1
        
    for i in range(1, n + 1):
        tau_inv[tau[i]] = i
        
    for i in range(1, n):
        sigma[i] = i + 1
        
    prev = 0
    for k in range(1, m + 1):
        t = k * (k + 1) // 2
        sigma[t] = prev + 1
        prev = t
        
    for i in range(1, n + 1):
        pi[i] = tau_inv[sigma[tau[i]]]
        
    cycles = []
    cycle_id = [-1] * (n + 1)
    pos_in = [-1] * (n + 1)
    visited = [False] * (n + 1)
    
    for i in range(1, n + 1):
        if visited[i]:
            continue
        cyc = []
        cur = i
        while not visited[cur]:
            visited[cur] = True
            pos_in[cur] = len(cyc)
            cycle_id[cur] = len(cycles)
            cyc.append(cur)
            cur = pi[cur]
        cycles.append(cyc)
        
    fact = [1] * (n + 1)
    weight = [0] * (n + 1)
    for i in range(1, n + 1):
        fact[i] = (fact[i - 1] * i) % MOD
    for i in range(1, n + 1):
        weight[i] = fact[n - i]
        
    m_fact = 1
    for i in range(1, m + 1):
        m_fact = (m_fact * i) % MOD
        
    num_cycles = len(cycles)
    cycle_len = [len(c) for c in cycles]
    
    maxL = 1
    for i in range(num_cycles):
        for j in range(num_cycles):
            g = math.gcd(cycle_len[i], cycle_len[j])
            L = cycle_len[i] // g * cycle_len[j]
            if L > maxL: maxL = L
            
    inv = [1] * (maxL + 1)
    for i in range(2, maxL + 1):
        inv[i] = MOD - (MOD // i) * inv[MOD % i] % MOD
        
    pair_data = [[None] * num_cycles for _ in range(num_cycles)]
    for ci in range(num_cycles):
        for cj in range(num_cycles):
            s = cycle_len[ci]
            t = cycle_len[cj]
            g = math.gcd(s, t)
            L = s // g * t
            mul = (m_fact * inv[L]) % MOD
            counts = compute_counts(cycles[ci], cycles[cj], g)
            pair_data[ci][cj] = {'g': g, 'mul': mul, 'counts': counts}
            
    total = m_fact % MOD
    
    local = 0
    for i in range(1, n + 1):
        wi = weight[i]
        ci = cycle_id[i]
        pos_i = pos_in[i]
        for j in range(i + 1, n + 1):
            cj = cycle_id[j]
            pos_j = pos_in[j]
            pd = pair_data[ci][cj]
            g = pd['g']
            r = (pos_j - pos_i) % g
            if r < 0:
                r += g
            cnt = pd['counts'][r]
            add = (wi * (cnt * pd['mul'] % MOD)) % MOD
            local += add
            if local >= MOD:
                local -= MOD
                
    total += local
    total %= MOD
    return str(total)

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

Java

import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;

public class Euler902 {
    static final long MOD = 1000000007L;
    static final long A = 1000000007L;

    static class PairData {
        int g;
        long mul;
        int[] counts;
    }

    static int gcd(int a, int b) {
        return b == 0 ? a : gcd(b, a % b);
    }

    static int[] computeCounts(List<Integer> Acyc, List<Integer> Bcyc, int g) {
        int s = Acyc.size();
        int t = Bcyc.size();

        List<List<Integer>> aR = new ArrayList<>(g);
        List<List<Integer>> bR = new ArrayList<>(g);
        for (int i = 0; i < g; i++) {
            aR.add(new ArrayList<>());
            bR.add(new ArrayList<>());
        }
        for (int i = 0; i < s; i++)
            aR.get(i % g).add(Acyc.get(i));
        for (int i = 0; i < t; i++)
            bR.get(i % g).add(Bcyc.get(i));

        int[][] aArr = new int[g][];
        int[][] bArr = new int[g][];
        for (int r = 0; r < g; r++) {
            aArr[r] = aR.get(r).stream().mapToInt(Integer::intValue).toArray();
            bArr[r] = bR.get(r).stream().mapToInt(Integer::intValue).toArray();
            Arrays.sort(aArr[r]);
            Arrays.sort(bArr[r]);
        }

        int[][] f = new int[g][g];
        for (int r = 0; r < g; r++) {
            for (int rp = 0; rp < g; rp++) {
                int cnt = 0;
                int[] a = aArr[r];
                int[] b = bArr[rp];
                int j = 0;
                for (int x : a) {
                    while (j < b.length && b[j] < x) {
                        j++;
                    }
                    cnt += j;
                }
                f[r][rp] = cnt;
            }
        }

        int[] counts = new int[g];
        for (int rd = 0; rd < g; rd++) {
            int total = 0;
            for (int r = 0; r < g; r++) {
                total += f[r][(r + rd) % g];
            }
            counts[rd] = total;
        }
        return counts;
    }

    public static String solve() {
        int m = 100;
        int n = m * (m + 1) / 2;

        int[] tau = new int[n + 1];
        int[] tauInv = new int[n + 1];
        int[] sigma = new int[n + 1];
        int[] pi = new int[n + 1];

        for (int i = 1; i <= n; ++i) {
            tau[i] = (int) ((A * i) % n) + 1;
        }
        for (int i = 1; i <= n; ++i) {
            tauInv[tau[i]] = i;
        }

        for (int i = 1; i < n; ++i) {
            sigma[i] = i + 1;
        }
        int prev = 0;
        for (int k = 1; k <= m; ++k) {
            int t = k * (k + 1) / 2;
            sigma[t] = prev + 1;
            prev = t;
        }

        for (int i = 1; i <= n; ++i) {
            pi[i] = tauInv[sigma[tau[i]]];
        }

        List<List<Integer>> cycles = new ArrayList<>();
        int[] cycleId = new int[n + 1];
        int[] posIn = new int[n + 1];
        boolean[] visited = new boolean[n + 1];

        for (int i = 1; i <= n; ++i) {
            if (visited[i])
                continue;
            List<Integer> cyc = new ArrayList<>();
            int cur = i;
            while (!visited[cur]) {
                visited[cur] = true;
                posIn[cur] = cyc.size();
                cycleId[cur] = cycles.size();
                cyc.add(cur);
                cur = pi[cur];
            }
            cycles.add(cyc);
        }

        long[] fact = new long[n + 1];
        long[] weight = new long[n + 1];
        fact[0] = 1;
        for (int i = 1; i <= n; ++i)
            fact[i] = (fact[i - 1] * i) % MOD;
        for (int i = 1; i <= n; ++i)
            weight[i] = fact[n - i];

        long mFact = 1;
        for (int i = 1; i <= m; ++i)
            mFact = (mFact * i) % MOD;

        int numCycles = cycles.size();
        int[] cycleLen = new int[numCycles];
        for (int i = 0; i < numCycles; ++i)
            cycleLen[i] = cycles.get(i).size();

        int maxL = 1;
        for (int i = 0; i < numCycles; ++i) {
            for (int j = 0; j < numCycles; ++j) {
                int g = gcd(cycleLen[i], cycleLen[j]);
                int L = cycleLen[i] / g * cycleLen[j];
                maxL = Math.max(maxL, L);
            }
        }

        long[] inv = new long[maxL + 1];
        inv[1] = 1;
        for (int i = 2; i <= maxL; ++i) {
            inv[i] = MOD - (MOD / i) * inv[(int) (MOD % i)] % MOD;
        }

        PairData[][] pairData = new PairData[numCycles][numCycles];
        for (int ci = 0; ci < numCycles; ++ci) {
            for (int cj = 0; cj < numCycles; ++cj) {
                int s = cycleLen[ci];
                int t = cycleLen[cj];
                int g = gcd(s, t);
                int L = s / g * t;
                PairData pd = new PairData();
                pd.g = g;
                pd.mul = (mFact * inv[L]) % MOD;
                pd.counts = computeCounts(cycles.get(ci), cycles.get(cj), g);
                pairData[ci][cj] = pd;
            }
        }

        long total = mFact % MOD;
        int threads = Math.max(1, Runtime.getRuntime().availableProcessors());
        if (threads > 8)
            threads = 8;
        long[] partial = new long[threads];
        Thread[] workers = new Thread[threads];

        int block = (n + threads - 1) / threads;

        for (int th = 0; th < threads; th++) {
            final int tIdx = th;
            final int start = th * block + 1;
            final int end = Math.min(n, (th + 1) * block);

            workers[th] = new Thread(() -> {
                long local = 0;
                if (start <= end) {
                    for (int i = start; i <= end; ++i) {
                        long wi = weight[i];
                        int ci = cycleId[i];
                        int posI = posIn[i];
                        for (int j = i + 1; j <= n; ++j) {
                            int cj = cycleId[j];
                            int posJ = posIn[j];
                            PairData pd = pairData[ci][cj];
                            int g = pd.g;
                            int r = (posJ - posI) % g;
                            if (r < 0)
                                r += g;
                            int cnt = pd.counts[r];
                            long add = (wi * (cnt * pd.mul % MOD)) % MOD;
                            local += add;
                            if (local >= MOD)
                                local -= MOD;
                        }
                    }
                }
                partial[tIdx] = local;
            });
            workers[th].start();
        }

        try {
            for (Thread t : workers) {
                if (t != null)
                    t.join();
            }
        } catch (InterruptedException e) {
            e.printStackTrace();
        }

        for (long v : partial) {
            total = (total + v) % MOD;
        }

        return Long.toString(total);
    }

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