题解 P8104 【「LCOI2021」Cow Function】

· · 题解

前置芝士:Min_25 筛

考虑算贡献。对于 g(m) > 0,求出 g(m) = \displaystyle\sum_{i = 0}^n [\omega(i) = m] 后可得 ans_i = \displaystyle\sum_{i \equiv j \pmod 8} 3^j g(j)

考虑先执行 Min_25 筛第一部分求出 \pi 的块筛,再执行类似于 Min_25 筛第二部分的限制质因数幂次的 dfs。

时间复杂度不会算(

具体细节见代码注释。

代码:

#include <stdio.h>
#include <math.h>

typedef long long ll;

const int N = 1e5 + 7, M = 2e5 + 7, K = 8;
ll sqrt_n;
int max_prime_cnt[M], id1[N], id2[N];
ll prime[N], number[M], g[M], ans[K + 7];
bool p[N];

inline ll sqrt(ll n){
    ll ans = sqrt((double)n);
    while (ans * ans <= n) ans++;
    return ans - 1;
}

inline ll max(ll a, ll b){
    return a > b ? a : b;
}

inline int get_id(ll n, ll m){
    return n <= sqrt_n ? id1[n] : id2[m / n];
}

inline int init(ll n){
    int cnt = 0, sqrt_cnt, id = 0;
    ll limit;
    sqrt_n = sqrt(n);
    if (sqrt_n == 1) sqrt_cnt = 0;
    limit = max(sqrt_n, 3);
    p[0] = p[1] = true;
    for (register ll i = 2; i <= limit; i++){
        if (!p[i]) prime[++cnt] = i;
        if (i == sqrt_n) sqrt_cnt = cnt;
        for (register int j = 1; j <= cnt && i * prime[j] <= limit; j++){
            p[i * prime[j]] = true;
            if (i % prime[j] == 0) break;
        }
    }
    for (register ll i = 1, j; i <= n; i = j + 1){
        ll tn = n / i;
        j = n / tn;
        id++;
        number[id] = tn;
        g[id] = tn - 1;
        if (tn <= sqrt_n){
            id1[tn] = id;
        } else {
            id2[j] = id;
        }
    }
    for (register int i = 1; i <= sqrt_cnt; i++){
        ll t = prime[i] * prime[i];
        for (register int j = 1; j <= id && number[j] >= t; j++){
            g[j] -= g[get_id(number[j] / prime[i], n)] - (i - 1);
        }
    }
    return cnt;
}

inline int get_max_prime_cnt(ll n, int cnt){
    int ans = 0;
    for (register ll i = 1; i <= n && ans <= cnt; i *= prime[ans]) ans++;
    return ans - 1;
}

inline ll quick_pow(ll x, ll p){
    ll ans = 1;
    while (p){
        if (p & 1) ans *= x;
        x *= x;
        p >>= 1;
    }
    return ans;
}

inline ll log(ll n, ll m){
    ll ans = log2(n) / log2(m);
    while (quick_pow(m, ans) <= n) ans++;
    return ans - 1;
}

ll solve(ll n, int m, int k, ll x, int cnt){
    if (m == 0) return 1;
    if (quick_pow(prime[k], m) > n) return 0; // 剪枝
    if (m == 1){
        ll ans = g[get_id(n, x)] - k;
        for (register int i = k + 1; i <= cnt && prime[i] * prime[i] <= n; i++){
            ans += log(n, prime[i]) - 1; // 质数的 > 1 次幂
        }
        return ans;
    }
    int md = m - 1;
    ll ans = 0;
    for (register int i = k + 1; i <= cnt && prime[i] * prime[i] <= n; i++){
        for (register ll j = prime[i]; j <= n; j *= prime[i]){
            ans += solve(n / j, md, i, x, cnt);
        }
    }
    return ans;
}

int main(){
    int cnt, m;
    ll n;
    scanf("%lld", &n);
    cnt = init(n);
    m = get_max_prime_cnt(n, cnt);
    for (register int i = 0, j = 0, k = 1; i <= m; i++, j = (j + 1) % K, k *= 3){
        ans[j] += k * solve(n, i, 0, n, cnt);
    }
    for (register int i = 0; i < K; i++){
        printf("%lld\n", ans[i]);
    }
    return 0;
}