peiwenjun's blog 没有知识的荒原

P5824 十二重计数法 题解

题目描述

将 \(n\) 个球装进 \(m\) 个盒子,对如下 \(12\) 种限制条件,分别计算方案数,对 \(998244353\) 取模。

  1. 球互不相同,盒子互不相同。
  2. 球互不相同,盒子互不相同,每个盒子至多装一个球。
  3. 球互不相同,盒子互不相同,每个盒子至少装一个球。
  4. 球互不相同,盒子全部相同。
  5. 球互不相同,盒子全部相同,每个盒子至多装一个球。
  6. 球互不相同,盒子全部相同,每个盒子至少装一个球。
  7. 球全部相同,盒子互不相同。
  8. 球全部相同,盒子互不相同,每个盒子至多装一个球。
  9. 球全部相同,盒子互不相同,每个盒子至少装一个球。
  10. 球全部相同,盒子全部相同。
  11. 球全部相同,盒子全部相同,每个盒子至多装一个球。
  12. 球全部相同,盒子全部相同,每个盒子至少装一个球。

数据范围

  • \(1\le n,m\le 2\cdot 10^5\) 。

时间限制 \(\texttt{2s}\) ,空间限制 \(\texttt{256MB}\) 。

分析

一、球不同,盒子不同,无其他限制

每个球有 \(m\) 种选法,并且不同的球之间互相独立,答案为 \(m^n\) 。

二、球不同,盒子不同,每个盒子至多装一个球

第 \(i\) 个球有 \(m-i+1\) 种选法,答案为 \(m^\underline n\) 。

三、球不同,盒子不同,盒子非空

容斥,枚举 \(k\) 个盒子为空,答案为 \(\sum_{k=0}^m(-1)^k\binom mk(m-k)^n\) 。

四、球不同,盒子相同,无其他限制

枚举 \(k\) 个盒子非空,根据第二类斯特林数的定义,答案为 \(\sum_{k=0}^m\begin{Bmatrix}n\\k\\\end{Bmatrix}\) 。

五、球不同,盒子相同,每个盒子至多装一个球

每个球无论装入哪个盒子效果都一样,答案为 \([n\le m]\) 。

六、球不同,盒子相同,盒子非空

根据第二类斯特林数的定义,答案为 \(\begin{Bmatrix}n\\m\\\end{Bmatrix}\) 。

七、球相同,盒子不同,无其他限制

等价于求 \(\sum_{i=1}^mx_i=n\) 的非负整数解个数,答案为 \(\binom{n+m-1}{m-1}\) 。

八、球相同,盒子不同,每个盒子至多装一个球

从 \(m\) 个盒子中选 \(n\) 个放球即可,答案为 \(\binom mn\) 。

九、球相同,盒子不同,盒子非空

等价于求 \(\sum_{i=1}^mx_i=n\) 的正整数解个数,答案为 \(\binom{n-1}{m-1}\) 。

十、球相同,盒子相同,无其他限制

等价于求将 \(n\) 划分为 \(m\) 个自然数之和的方案数,记为 \(p_{n,m}\) 。

维护一个可重集,每次 push 一个零或者给集合中全体数加一,可得转移方程:

\[p_{n,m}=p_{n,m-1}+p_{n-m,m}\\ \]

将第一维压入生成函数,记 \(f_i(x)=\sum_{n=0}^\infty p_{n,i}\cdot x^i\) 。

容易得出递推式:

\[f_i(x)=f_{i-1}(x)\cdot(1+x^i+x^{2i}+\cdots)\\ \]

类比付公主的背包,继续推式子:

\[\begin{aligned} f_m(x)&=\prod_{i=1}^m\frac1{1-x^i}\\ &=\exp(\sum_{i=1}^m-\ln(1-x^i))\\ &=\exp(\sum_{i=1}^m\sum_{k=1}^\infty\frac1kx^{ki})\\ \end{aligned} \]

最后输出 \([x^n]f_m(x)\) 即可。

十一、球相同,盒子相同,每个盒子至多装一个球

同第五种情况,答案为 \([n\le m]\) 。

十二、球相同,盒子相同,盒子非空

从每个盒子中删掉一个球,再套用第十种情况的结论,答案为 \([x^{n-m}]f_m(x)\) 。

时间复杂度 \(\mathcal O(n\log n+n\log m)\) 。

#include<bits/stdc++.h>
using namespace std;
const int maxn=1<<19,mod=998244353;
int m,n;
int a[maxn],b[maxn],s[maxn],fac[maxn],inv[maxn];
inline int qpow(int a,int k)
{
    int res=1;
    for(;k;a=1ll*a*a%mod,k>>=1) if(k&1) res=1ll*res*a%mod;
    return res;
}
inline int add(int x,int y) {return x+y>=mod?x+y-mod:x+y;}
inline int dec(int x,int y) {return x-y<0?x-y+mod:x-y;}
namespace Poly
{
    int p[maxn],q[maxn],r[maxn],w[maxn];
    int inum[maxn];
    inline int extend(int n) {return n!=1?1<<(__lg(n-1)+1):1;}
    inline void get_r(int n) {for(int i=0;i<n;i++) r[i]=(r[i>>1]>>1)|(i&1?n>>1:0);}
    static auto init=[]()
    {
        for(int k=2,m=1;k<=maxn;k<<=1,m<<=1)
        {
            w[m]=1;
            for(int i=m+1,x=qpow(3,(mod-1)/k);i<k;i++) w[i]=1ll*w[i-1]*x%mod;
        }
        for(int i=1;i<maxn;i++) inum[i]=qpow(i,mod-2);
        return 0;
    }();
    inline void print(int *a,int n) {for(int i=0;i<n;i++) printf("%d%c",a[i]," \n"[i==n-1]);}
    void ntt(int *a,int n,int x)
    {
        for(int i=0;i<n;i++) if(i<r[i]) swap(a[i],a[r[i]]);
        for(int k=2,m=1;k<=n;k<<=1,m<<=1)
            for(int i=0;i<n;i+=k)
                for(int j=i,*x=w+m;j<i+m;j++,x++)
                {
                    int v=1ll*a[j+m]**x%mod;
                    a[j+m]=dec(a[j],v),a[j]=add(a[j],v);
                }
        if(x==-1)
        {
            reverse(a+1,a+n);
            for(int i=0,v=qpow(n,mod-2);i<n;i++) a[i]=1ll*a[i]*v%mod;
        }
    }
    void mul(int *a,int *b,int n,int m)
    {
        if(!n||!m) return ;
        int len=extend(n+m-1);
        for(int i=0;i<len;i++) p[i]=i<n?a[i]:0,q[i]=i<m?b[i]:0;
        get_r(len),ntt(p,len,1),ntt(q,len,1);
        for(int i=0;i<len;i++) a[i]=1ll*p[i]*q[i]%mod;
        ntt(a,len,-1);
    }
    void inv(int *a,int *b,int n)
    {
        static int c[maxn],d[maxn];
        n=extend(n),memset(b,0,4*n),b[0]=qpow(a[0],mod-2);
        for(int k=2;k<=n;k<<=1)
        {
            for(int i=0;i<k<<1;i++) c[i]=i<k?a[i]:0,d[i]=i<k>>1?b[i]:0;
            get_r(k),ntt(d,k,1);
            for(int i=0;i<k;i++) d[i]=1ll*d[i]*d[i]%mod;
            ntt(d,k,-1),get_r(k<<1),ntt(c,k<<1,1),ntt(d,k<<1,1);
            for(int i=0;i<k<<1;i++) c[i]=1ll*c[i]*d[i]%mod;
            ntt(c,k<<1,-1);
            for(int i=0;i<k;i++) b[i]=(2ll*b[i]-c[i]+mod)%mod;
        }
    }
    void diff(int *a,int *b,int n)
    {
        for(int i=1;i<n;i++) b[i-1]=1ll*i*a[i]%mod;
        b[n-1]=0;
    }
    void integ(int *a,int *b,int n)
    {
        for(int i=1;i<n;i++) b[i]=1ll*inum[i]*a[i-1]%mod;
        b[0]=0;
    }
    void ln(int *a,int *b,int n)
    {
        static int c[maxn],d[maxn];
        assert(a[0]==1);
        n=extend(n),inv(a,c,n),diff(a,d,n),mul(c,d,n,n),integ(c,b,n);
    }
    void exp(int *a,int *b,int n)
    {
        static int c[maxn];
        assert(a[0]==0);
        n=extend(n),memset(b,0,4*n),memset(c,0,4*n),b[0]=1;
        for(int k=2;k<=n;k<<=1)
        {
            ln(b,c,k);
            for(int i=0;i<k;i++) c[i]=dec(a[i],c[i]);
            c[0]++,mul(b,c,k,k);
        }
    }
}
void init(int n)
{
    fac[0]=1;
    for(int i=1;i<=n;i++) fac[i]=1ll*fac[i-1]*i%mod;
    inv[n]=qpow(fac[n],mod-2);
    for(int i=n;i>=1;i--) inv[i-1]=1ll*inv[i]*i%mod;
}
inline int c(int n,int m)
{
    return n>=m?1ll*fac[n]*inv[m]%mod*inv[n-m]%mod:0ll;
}
int main()
{
    scanf("%d%d",&n,&m),init(n+m);
    for(int i=0;i<=n;i++) a[i]=(i&1?mod-1ll:1)*inv[i]%mod,b[i]=1ll*qpow(i,n)*inv[i]%mod;
    Poly::mul(a,b,n+1,n+1),memcpy(s,a,4*(n+1));///第二类斯特林数
    memset(a,0,sizeof(a));
    for(int i=1;i<=m;i++) for(int j=i;j<=n;j+=i) a[j]=add(a[j],Poly::inum[j/i]);
    Poly::exp(a,b,n);///f_m(x)
    for(int x=1;x<=12;x++)
    {
        if(x==1) printf("%d\n",qpow(m,n));
        if(x==2) printf("%d\n",n<=m?1ll*fac[m]*inv[m-n]%mod:0);
        if(x==3)
        {
            int res=0;
            for(int i=0;i<=m;i++) res=(res+(i&1?-1ll:1ll)*c(m,i)*qpow(m-i,n))%mod;
            printf("%d\n",add(res,mod));
        }
        if(x==4)
        {
            int res=0;
            for(int i=0;i<=m;i++) res=add(res,s[i]);
            printf("%d\n",res);
        }
        if(x==5) printf("%d\n",n<=m);
        if(x==6) printf("%d\n",s[m]);
        if(x==7) printf("%d\n",c(n+m-1,m-1));
        if(x==8) printf("%d\n",c(m,n));
        if(x==9) printf("%d\n",c(n-1,m-1));
        if(x==10) printf("%d\n",b[n]);
        if(x==11) printf("%d\n",n<=m);
        if(x==12) printf("%d\n",n>=m?b[n-m]:0);
    }
    return 0;
}

posted on 2023-04-07 10:57  peiwenjun  阅读(17)  评论(0)    收藏  举报

导航