题解 P7884 【模板】Meissel–Lehmer 算法

· · 题解

题目传送门

前言

这篇题解提供一个时间复杂度为 O(\dfrac{n^{\frac{2}{3}}}{\log ^{2}n}) 的较易实现的方法。

流程

考虑朴素的 DP :

S(n,a)=S(n,a-1)-S(\dfrac{n}{p_a},a-1)+S(p_{a-1},a-1)

对于 p \le n^{\frac{1}{8}} ,我们使用以上式子进行转移

对于 p > n^{\frac{1}{8}} ,我们考虑用树状数组维护 k 以下的数。注意到含有小于 n^{\frac{1}{8}} 的素因数的数已被筛去,树状数组只需进行 O(\dfrac{k}{\log n}) 次修改,时间复杂度 O(k)

如果我们每次都使用树状数组查询 k 以下的值,时间复杂度将超出我们的预期。于是我们考虑在每一次筛完素数后,将不会再更改的值存储于静态数组中。 这一部分的时间复杂度为 O(\sum\limits_{i=1}^{n/k}\dfrac{\sqrt{n/i}}{\log n}) = O(\dfrac{n}{\sqrt{k}\log n})

事实上,我们没有必要对每个 \frac{n}{i} 的数进行转移,在筛完某个素数 p 后,会影响结果的只有满足最小素因数大于 pi ,我们考虑在朴素 dp 后进行这一优化 。于是这部分的时间复杂度优化为 O(\dfrac{n}{\sqrt{k}\log ^{2}n})

然而我们并没有必要转移 p\sqrt{n},将 p 转移至 n^{\frac{1}{4}} 后可借助容斥得到答案。于是我们再次优化为O(\dfrac{n\pi(n^{\frac{1}{4}})}{k\log n}) = O(\dfrac{n^{\frac{5}{4}}}{k\log ^{2}n})。取 k = O(\dfrac{n^{\frac{5}{8}}}{\log n}),得 O(\dfrac{n^{\frac{5}{8}}}{\log n}) 的时间复杂度。

容斥部分的时间复杂度为 O(\sum\limits_{n^{\frac{1}{4}} <p\le n^{\frac{1}{3}}}\pi(\sqrt{\dfrac{n}{p}})) = O(\dfrac{n^{\frac{2}{3}}}{\log ^{2}n})

因此总体时间复杂度为 O(\dfrac{n^{\frac{2}{3}}}{\log ^{2}n})

实现

代码没做过多优化,以供阅读参考。

#include<bits/stdc++.h>
using namespace std;
using i64 = long long;
i64 count_pi(i64 N) {
    if(N <= 1) return 0;
    int v = sqrt(N + 0.5);
    int n_4 = sqrt(v + 0.5);
    int T = min((int)sqrt(n_4) * 2, n_4);
    int K = pow(N, 0.625) / log(N) * 2;
    K = max(K, v);
    K = min<i64>(K, N);
    int B = N / K;
    B = N / (N / B);
    B = min<i64>(N / (N / B), K);

    vector<i64> l(v + 1);
    vector<int> s(K + 1);
    vector<bool> e(K + 1);
    vector<int> w(K + 1);
    for (int i = 1; i <= v; ++i) l[i] = N / i - 1;
    for (int i = 1; i <= v; ++i) s[i] = i - 1;

    const auto div = [] (i64 n, int d) -> int { return double(n) / d; };
    int p;
    for (p = 2; p <= T; ++p) 
        if (s[p] != s[p - 1]) {
            i64 M = N / p;
            int t = v / p, t0 = s[p - 1];
            for (int i = 1; i <= t; ++i) l[i] -= l[i * p] - t0;
            for (int i = t + 1; i <= v; ++i) l[i] -= s[div(M, i)] - t0;
            for (int i = v, j = t; j >= p; --j)
                for (int l = j * p; i >= l; --i)
                    s[i] -= s[j] - t0;
            for (int i = p * p; i <= K; i += p) e[i] = 1;
        }
    e[1] = 1;
    int cnt = 1;
    vector<int> roughs(B + 1);
    for (int i = 1; i <= B; ++i)
        if(!e[i]) roughs[cnt++] = i;
    roughs[cnt] = 0x7fffffff;
    for (int i = 1; i <= K; ++i) w[i] = e[i] + w[i - 1];
    for (int i = 1; i <= K; ++i) s[i] = w[i] - w[i - (i & -i)];

    const auto query = [&] (int x) -> int {
        int sum = x;
        while(x) sum -= s[x], x ^= x & -x;
        return sum;
    };
    const auto add = [&] (int x) -> void {
        e[x] = 1;
        while(x <= K) ++s[x], x += x & -x;
    };
    cnt = 1;
    for (; p <= n_4; ++p) 
        if(!e[p]) {
            i64 q = i64(p) * p, M = N / p;
            while(cnt < q) w[cnt] = query(cnt), cnt++;
            int t1 = B / p, t2 = min<i64>(B, M / q), t0 = query(p - 1);
            int id = 1, i = 1;
            for (; i <= t1; i = roughs[++id]) l[i] -= l[i * p] - t0;
            for (; i <= t2; i = roughs[++id]) l[i] -= query(div(M, i)) - t0;
            for (; i <= B; i = roughs[++id]) l[i] -= w[div(M, i)] - t0;
            for (int i = q; i <= K; i += p)
                if(!e[i]) add(i);
        }
    while(cnt <= v) w[cnt] = query(cnt), cnt++;

    vector<int> primes;
    primes.push_back(1);
    for (int i = 2; i <= v; ++i)
        if(!e[i]) primes.push_back(i);
    l[1] += i64(w[v] + w[n_4] - 1) * (w[v] - w[n_4]) / 2;
    for (int i = w[n_4] + 1; i <= w[B]; ++i) l[1] -= l[primes[i]];
    for (int i = w[B] + 1; i <= w[v]; ++i) l[1] -= query(N / primes[i]);
    for (int i = w[n_4] + 1; i <= w[v]; ++i) {
        int q = primes[i];
        i64 M = N / q;
        int e = w[M / q];
        if (e <= i)  break;
        l[1] += e - i;
        i64 t = 0;
        int m = w[sqrt(M + 0.5)];
        for (int k = i + 1; k <= m; ++k) t += w[div(M, primes[k])];
        l[1] += 2 * t - (i + m) * (m - i);
    }
    return l[1];
}
int main() {
    i64 n;
    cin >> n;
    cout << count_pi(n);
    return 0;
}