atc abc473G 思路分享(概率 dp,生成函数,ntt)

题目链接

题意

\(n\) 张卡片,编号从 \(1\)\(n\),随机打乱排列,进行如下游戏:

初始 \(x=1\),每一轮随机翻出一张卡片,直到 \(x\gt n\)

  • 若卡片编号等于 \(x\),丢弃这张卡片,\(x \leftarrow x+1\).

  • 否则放回卡片,你会永远记住这张卡的编号.

你想要用最少操作次数丢弃所有卡片,并且你永远会做最优的决策,求操作次数恰好为 \(k\) 的概率,模 \(998244353\).

\(1\le n \le 5\times 10^5\)\(n\le k \le 10^9\).

思路

首先不难发现,每张卡最多被翻 \(2\) 次,最少被翻一次,因此 \(k\gt 2n\) 概率为 \(0\).

考虑 \(dp\),令 \(dp[i][j]\) 为还剩 \(i\) 张未知卡,此时额外操作了 \(j\) 次的概率(即不算基础的 \(n\) 次),转移方程:

\[dp[i][j] = \frac{1}{i}dp[i-1][j]+\frac{i-1}{i}dp[i-1][j-1] \]

尝试构造生成函数 \(F_i(x) = \sum_{j}{dp[i][j] \cdot x^j}\),将转移方程带入:

\[F_i(x) = \sum_{j}{\frac{1}{i}dp[i-1][j] \cdot x^j}+\sum_{j}{\frac{i-1}{i}dp[i-1][j-1]\cdot x^j} \]

\[\Rightarrow F_i(x) = \frac{1}{i}F_{i-1}(x)+\frac{i-1}{i}x \cdot F_{i-1}(x) \]

\[\Rightarrow F_i(x) = (\frac{1}{i}+\frac{i-1}{i}x)F_{i-1}(x) \]

\[\Rightarrow F_i(x) = \prod_{j=1}^{i}{(\frac{1}{j}+\frac{j-1}{j}x)} \]

于是 \(F_n(x)\) 实际上是 \(n\) 个一次多项式相乘,使用分治 \(ntt\)(多项式乘积树)计算即可,注意截断到第 \(k-n\) 项.

时间复杂度 \(\mathcal{O}(n\log^2 n)\).

代码

//author:kzssCCC

#include <bits/stdc++.h>
using namespace std;
using ll = long long;

template<int MOD>
struct modint {
    int val;
    
    modint() : val(0) {}
    modint(long long v) {
        val = v % MOD;
        if (val < 0) val += MOD;
    }
    
    modint& operator++() { val = (val + 1 == MOD ? 0 : val + 1); return *this; }
    modint& operator--() { val = (val == 0 ? MOD - 1 : val - 1); return *this; }
    modint operator++(int) { modint res = *this; ++*this; return res; }
    modint operator--(int) { modint res = *this; --*this; return res; }
    
    modint& operator+=(const modint& o) { val += o.val; if (val >= MOD) val -= MOD; return *this; }
    modint& operator-=(const modint& o) { val -= o.val; if (val < 0) val += MOD; return *this; }
    modint& operator*=(const modint& o) { val = 1LL * val * o.val % MOD; return *this; }
    modint& operator/=(const modint& o) { return *this *= o.inv(); }
    
    friend modint operator+(modint a, const modint& b) { return a += b; }
    friend modint operator-(modint a, const modint& b) { return a -= b; }
    friend modint operator*(modint a, const modint& b) { return a *= b; }
    friend modint operator/(modint a, const modint& b) { return a /= b; }

    friend bool operator==(const modint& a, const modint& b) { return a.val == b.val; }
    friend bool operator!=(const modint& a, const modint& b) { return a.val != b.val; }
    modint operator-() const { modint res = *this; res.val = (res.val == 0 ? 0 : MOD - res.val); return res;};
    modint operator+() const { return *this; };
    
    modint qpow(long long p) const {
        modint res = 1, a = *this;
        while (p > 0) {
            if (p & 1) res *= a;
            a *= a;
            p >>= 1;
        }
        return res;
    }

    modint inv() const {
        return qpow(MOD - 2);
    }
    
    friend std::ostream& operator<<(std::ostream& os, const modint& m) { return os << m.val; }
    friend std::istream& operator>>(std::istream& is, modint& m) { long long v; is >> v; m = modint(v); return is; }
};

const int MOD = 998244353;
using Z = modint<MOD>;

const Z g = 3;
const Z ig = g.inv();

void ntt(vector<Z>& vec,int op){
    int n = vec.size();
    int B = 31-__builtin_clz(n);

    for (int i=0;i<n;i++){
        int rev = 0;
        for (int j=0;j<B;j++){
            rev <<= 1;
            rev |= i>>j&1;
        }
        if (i<rev){
            swap(vec[i],vec[rev]);
        }
    }

    for (int len=2;len<=n;len<<=1){
        Z W = op==0?g.qpow((MOD-1)/len):ig.qpow((MOD-1)/len);
        for (int i=0;i<n;i+=len){
            Z w = 1;
            for (int j=0;j<len/2;j++){
                Z u = vec[i+j];
                Z v = vec[i+j+len/2];
                vec[i+j] = u+w*v;
                vec[i+j+len/2] = u-w*v;
                w *= W;
            }
        }
    }
    
    if (op==1){
        Z temp = Z(n).inv();
        for (int i=0;i<n;i++){
            vec[i] *= temp;
        }
    }
}

vector<Z> mul(vector<Z>& A,vector<Z>& B){
    int n = A.size();
    int m = B.size();

    int N = 1;
    while (N<n+m-1){
        N <<= 1;
    }

    A.resize(N);
    B.resize(N);
    ntt(A,0);
    ntt(B,0);
    
    vector<Z> C(N);
    for (int i=0;i<N;i++){
        C[i] = A[i]*B[i];   
    }

    ntt(C,1);
    return C;
}

void solve(){
    int n,k;
    cin >> n >> k;

    k -= n;
    if (k>n){
        cout << 0 << '\n';
        return;
    }

    vector<vector<Z>> vec(n+1);
    for (int i=1;i<=n;i++){
        vec[i] = {Z(1)/i,Z(i-1)/i};
    }

    function<vector<Z>(int,int)> work = [&](int l,int r){
        if (l==r) return vec[l];
            
        int mid = l+r >> 1;
        auto A = work(l,mid);
        auto B = work(mid+1,r);
        auto C = mul(A,B);
        C.resize(min(k+1,(int)C.size()));
        return C;
    };

    auto C = work(1,n);
    cout << C[k] << '\n';
}

int main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    
    int t = 1;
    // cin >> t;
    while (t--) solve();

    return 0;
}
posted @ 2026-09-01 18:49  kzssCCC  阅读(15)  评论(0)    收藏  举报