P1405 苦恼的小明

· · 题解

解题思路

本题需要求 a_1^{a_2^{a_3^{...^{a_n}}}} \bmod 10007 的值。

首先,由于模数很小,可以采用快速幂的方法计算:

int qpow(int a, int b, int p) {
    int ans = 1;
    while (b) {
        if (b & 1) ans = 1ll * ans * a % p;
        a = 1ll * a * a % p;
        b >>= 1;
    }
    return ans;
}

但是直接使用快速幂会导致数据爆炸,因此需要使用欧拉函数和扩展欧拉定理。

首先,需要求出模 10007 意义下的欧拉函数的值 phi,由于 10007 是质数,因此有 phi(10007)=10006。

接下来使用扩展欧拉定理,即:

a_1^{a_2^{a_3^{...^{a_n}}}}\equiv \begin{cases} a_1^{a_2^{a_3^{...^{a_n}} \bmod \varphi(p)}+\varphi(p)} &\textrm{gcd}(a_1,p)=1 \\ a_1^{a_2^{a_3^{...^{a_n}}} \bmod (p-1)+p-1} &\textrm{gcd}(a_1,p)\neq 1 \end{cases} \pmod p

其中 p=10007,\varphi(p)=10006。

代码如下:

int phi(int n) {
    int res = n;
    for (int i = 2; i * i <= n; i++) {
        if (n % i == 0) {
            res = res / i * (i - 1);
            while (n % i == 0) n /= i;
        }
    }
    if (n > 1) res = res / n * (n - 1);
    return res;
}

long long solve(long long i, long long p) {
    if (i == n) 
      return a[i]=a[i] % p;
    long long v=phi(p),t = solve(i + 1, v);
    if(gcd(a[i],p)==1)
      return a[i]=qpow(a[i],t,p);
    else{
        if(a[i+1]<=v)
            return a[i]=qpow(a[i],a[i+1],p);
        return a[i]=qpow(a[i],t+v,p);
    }
}

时间复杂度

求欧拉函数的时间复杂度为 O(\sqrt n),使用递归求解的时间复杂度为 O(n)(因为每次需要求解欧拉函数),最终时间复杂度为 O(n\sqrt n)。

完整代码

#include <iostream>

using namespace std;

const long long N = 1e6 + 7, P = 10007;

long long qpow(long long a, long long b, long long p) {
    long long ans = 1;
    while (b) {
        if (b & 1) ans = 1ll * ans * a % p;
        a = 1ll * a * a % p;
        b >>= 1;
    }
    return ans;
}

long long phi(long long x) {
    long long res = x;
    for (long long i = 2; i * i <= x; i++) {
        if (x % i == 0) {
            res = res / i * (i - 1);
            while (x % i == 0) x /= i;
        }
    }
    if (x > 1) 
      res = res / x * (x - 1);
    return res;
}
long long gcd(long long a,long long b){
    if(b==0) 
      return a;
    return gcd(b,a%b);
}
long long n, a[N];
long long solve(long long i, long long p) {
    if (i == n) 
      return a[i]=a[i] % p;
    long long v=phi(p),t = solve(i + 1, v);
    if(gcd(a[i],p)==1)
      return a[i]=qpow(a[i],t,p);
    else{
        if(a[i+1]<=v)
            return a[i]=qpow(a[i],a[i+1],p);
        return a[i]=qpow(a[i],t+v,p);
    }
}
int main() {
    cin >> n;
    for (long long i=1;i<=n;i++)
      cin>>a[i];
    cout<<solve(1,P)<<'\n';
    return 0;
}