搞半天是NTT写挂了,变量写错了。。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
#include <bits/stdc++.h>
/*
#include<ext/pb_ds/assoc_container.hpp>
#include<ext/pb_ds/hash_policy.hpp>
*/
using namespace std;

const double eps = 1e-10;
const double pi = 3.1415926535897932384626433832795;
const double eln = 2.718281828459045235360287471352;
#define int long long
#define f(i, a, b) for (int i = a; i <= b; i++)
#define scan(x) scanf("%d", &x)
#define mp make_pair
#define pb push_back
#define lowbit(x) (x&(-x))

#define fi first
#define se second
#define SZ(x) int((x).size())
#define all(x) x.begin(), x.end()
#define rall(x) x.rbegin(), x.rend()
#define summ(a) (accumulate(all(a), 0ll))

typedef unsigned long long ull;
typedef pair<int,int> pii;
typedef vector<int> vi;

using ll=long long;
const int mod=998244353,g=3,gi=332748118,N=1e5+9;
int tt,n,x,k,fac[N],inv[N],a[N],dp[N];

int rev[N],a0[N],a1[N],bit,lim;
int powmod(int a,int b){
int res=1;
for(;b>0;b>>=1,a=1ll*a*a%mod)if(b&1)res=1ll*res*a%mod;
return res;
}

void init(int n){
fac[0]=1,inv[0]=1;
for(int i=1;i<=n;++i)fac[i]=1ll*fac[i-1]*i%mod;
inv[n]=powmod(fac[n],mod-2);
for(int i=n-1;i;--i)inv[i]=1ll*inv[i+1]*(i+1)%mod;
}

void NTT(int *aa,int type){
for(int i=0;i<lim;++i)if(i<rev[i])swap(aa[i],aa[rev[i]]);
for (int k = 1; k < lim; k <<= 1){
int wn = powmod((type == 1 ? g : gi), (mod - 1) / (k << 1));
for (int i = 0; i < lim; i += (k << 1)){
int w=1;
for(int j=0;j<k;++j,w=1ll*w*wn%mod){
ll x=aa[i+j],y=1ll*w * aa[i + j + k]%mod;
aa[i + j] = (x + y)%mod;
aa[i + j + k] = ((x - y)%mod+mod)%mod;
}
}
}
if (type == 1)
return;
int invn = powmod(lim,mod-2);
for (int i = 0; i < lim; ++i)
aa[i] = 1ll*aa[i] * invn%mod;
}

void solve(){
cin>>n>>k>>x;
for(int i=0;i<=k;++i)dp[i]=0;
lim=1,bit=0;
while(lim<=2*k)lim<<=1,bit+=1;
for (int i = 0; i < lim; ++i)
rev[i] = (rev[i >> 1] >> 1) | ((i & 1) << (bit - 1));
for(int i=1,m;i<=n;++i){
m=1;
cin>>a[i];
for(int j=0;j<=k;++j){
a0[j]=1ll*dp[j]*inv[j]%mod;
a1[j]=1ll*m*inv[j]%mod;
m=1ll*m*a[i]%mod;
}
for(int j=k+1;j<=lim;++j)a0[j]=a1[j]=0;
NTT(a0,1),NTT(a1,1);
for(int i=0;i<lim;++i)a0[i]=1ll*a0[i]*a1[i]%mod;
NTT(a0,0);
m=1;
for(int j=0;j<=k;++j){
dp[j]=(1ll*dp[j]+1ll*m*x%mod+1ll*a0[j]*x%mod*fac[j]%mod)%mod;
m=1ll*m*a[i]%mod;
}
}
cout<<dp[k]<<"\n";
}
signed main()
{
ios::sync_with_stdio(false);
cin.tie(0);
init(1e4+100);
cin>>tt;
f(sb,1,tt)solve();
return 0;
}

只想到了二分图的做法

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
#include <bits/stdc++.h>
/*
#include<ext/pb_ds/assoc_container.hpp>
#include<ext/pb_ds/hash_policy.hpp>
*/
using namespace std;

const double eps = 1e-10;
const double pi = 3.1415926535897932384626433832795;
const double eln = 2.718281828459045235360287471352;

#define f(i, a, b) for (int i = a; i <= b; i++)
#define scan(x) scanf("%d", &x)
#define mp make_pair
#define pb push_back
#define lowbit(x) (x&(-x))

#define fi first
#define se second
#define SZ(x) int((x).size())
#define all(x) x.begin(), x.end()
#define rall(x) x.rbegin(), x.rend()
#define summ(a) (accumulate(all(a), 0ll))

typedef unsigned long long ull;
typedef pair<int,int> pii;
typedef vector<int> vi;

using ll=long long;

int n,m,k,ans,tot=1,hd[(int)2e5+9],cur[(int)2e5+9],s,t,dep[(int)2e5+9];
vector<pii> v;
map<pii,int> exi;
pii dir[4]={{1,0},{0,1},{-1,0},{0,-1}};

map<pii,int> cov;
struct Edge{
int to,nxt,v;
}e[(int)1e7+9];
void add(int u,int v){
e[++tot]={v,hd[u],1},hd[u]=tot;
e[++tot]={u,hd[v],0},hd[v]=tot;
}

bool bfs(int st,int en){
for(int i=0;i<=t;++i)dep[i]=0;
queue<int>q;
memcpy(cur,hd,sizeof(hd));
q.push(st);
dep[st]=1;
while(!q.empty()){
auto u=q.front();
q.pop();
for(int eg=hd[u];eg;eg=e[eg].nxt){
if(!dep[e[eg].to]&&e[eg].v>0){
dep[e[eg].to]=dep[u]+1;
q.push(e[eg].to);
}
}
}
return !!dep[en];
}

int dfs(int u,int en,int flow){
if(u==en)return flow;
int r=flow;
for(int eg=cur[u];eg&&r;eg=e[eg].nxt){
cur[u]=eg;
if(e[eg].v>0&&dep[e[eg].to]==dep[u]+1){
int c=dfs(e[eg].to,en,min(r,e[eg].v));
r-=c;
e[eg].v-=c,e[eg^1].v+=c;
}
}
return flow-r;
}

int dinic(int st,int en){
int ret=0;
while(bfs(st,en)){
ret+=dfs(st,en,0x3f3f3f3f);
}
return ret;
}

int main()
{
ios::sync_with_stdio(false);
cin.tie(0);
cin>>n>>m>>k;
v.resize(k+1);
for(int i=1,x,y;i<=k;++i){
cin>>x>>y;
cov[{x,y}]=i,v[i]={x,y},exi[{x,y}]=1;
}
s=2*k+1,t=2*k+2;
for(int i=1;i<=k;++i){
auto it=v[i];
auto p=cov[it];
add(s,p),add(p+k,t);
for(auto [aa,bb]:dir){
aa+=it.first,bb+=it.second;
//cerr<<it.first<<"--"<<it.second<<":::"<<aa<<"--"<<bb<<"\n";
if(exi.count({aa,bb})){
auto pp=cov[{aa,bb}];
add(p,pp+k);
}
}
}
bool ok=(dinic(s,t)==k);
cout<<(ok?"YES\n":"NO\n");
return 0;
}

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
#include <bits/stdc++.h>
/*
#include<ext/pb_ds/assoc_container.hpp>
#include<ext/pb_ds/hash_policy.hpp>
*/
using namespace std;

const double eps = 1e-10;
const double pi = 3.1415926535897932384626433832795;
const double eln = 2.718281828459045235360287471352;

#define f(i, a, b) for (int i = a; i <= b; i++)
#define scan(x) scanf("%d", &x)
#define mp make_pair
#define pb push_back
#define lowbit(x) (x&(-x))

#define fi first
#define se second
#define SZ(x) int((x).size())
#define all(x) x.begin(), x.end()
#define rall(x) x.rbegin(), x.rend()
#define summ(a) (accumulate(all(a), 0ll))

typedef unsigned long long ull;
typedef pair<int,int> pii;
typedef vector<int> vi;

using ll=long long;

ll n,a[1<<24],b[1<<24],rev[1<<24],bit,len,inv[1<<24],fac[1<<24];
const int mod=167772161,g=3,gi=55924054;
ll powmod(ll a,ll b){
ll res=1;
for(;b>0;b>>=1,a=a*a%mod)if(b&1)res=res*a%mod;
return res;
}

inline void NTT(ll *a,int type){
for(int i=0;i<len;++i){
if(i<rev[i])swap(a[i],a[rev[i]]);
}
for(int i=1;i<len;i<<=1){
ll gn=powmod(type?g:gi,(mod-1)/(i<<1));
for(int j=0;j<len;j+=(i<<1)){
ll g0=1;
for(int k=0;k<i;++k,g0=g0*gn%mod){
ll x=a[j+k],y=g0*a[i+j+k]%mod;
a[j+k]=(x+y)%mod;
a[i+j+k]=(x-y+mod)%mod;
}
}
}
if(type==1)return ;
ll invl=powmod(len,mod-2);
for(int i=0;i<len;++i)a[i]=a[i]*invl%mod;
}

int main()
{
ios::sync_with_stdio(false);
cin.tie(0);
cin>>n;
fac[1]=fac[0]=1,inv[0]=1;
for(int i=2;i<=n;++i)fac[i]=fac[i-1]*i%mod;
inv[n]=powmod(fac[n],mod-2);
for(int i=n-1;i;--i)inv[i]=inv[i+1]*(i+1)%mod;
for(int i=0;i<=n;++i)a[i]=(i&1)?(mod-inv[i]):inv[i],b[i]=powmod(i,n)*inv[i]%mod;
while((1<<bit)<=(n<<1))bit+=1;
len=1<<bit;
for(int i=0;i<len;++i)rev[i]=(rev[i>>1]>>1)|((i&1)<<(bit-1));
NTT(a,1),NTT(b,1);
for(int i=0;i<len;++i)a[i]=a[i]*b[i]%mod;
NTT(a,0);
for(int i=0;i<=n;++i)cout<<a[i]<<" \n"[i==n];
return 0;
}

抄的代码,实在是不会组合数学:):):)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
#include <bits/stdc++.h>
/*
#include<ext/pb_ds/assoc_container.hpp>
#include<ext/pb_ds/hash_policy.hpp>
*/
using namespace std;

const double eps = 1e-10;
const double pi = 3.1415926535897932384626433832795;
const double eln = 2.718281828459045235360287471352;

#define f(i, a, b) for (int i = a; i <= b; i++)
#define scan(x) scanf("%d", &x)
#define mp make_pair
#define pb push_back
#define lowbit(x) (x&(-x))

#define fi first
#define se second
#define SZ(x) int((x).size())
#define all(x) x.begin(), x.end()
#define rall(x) x.rbegin(), x.rend()
#define summ(a) (accumulate(all(a), 0ll))

typedef unsigned long long ull;
typedef pair<int,int> pii;
typedef vector<int> vi;

using ll=long long;

template <class T> inline T read(const T sample) {
T x=0; int f=1; char s;
while((s=getchar())>'9'||s<'0') if(s=='-') f=-1;
while(s>='0'&&s<='9') x=(x<<1)+(x<<3)+(s^48),s=getchar();
return x*f;
}

const int N=1e7+9,mod=998244353;
int n,m,k,f[N],ans;
namespace math{
int id[N],pct,inv[N];
inline int qpow(int n,int k){
int res=1;
for(;k>0;k>>=1,n=1ll*n*n%mod)if(k&1)res=1ll*n*res%mod;
return res;
}
inline void fmod(int &x){x-=mod;x+=x>>31&mod;}
void initmath(const int &n=N-1){
id[1]=1,id[0]=0;
for(int i=2;i<=n;++i){
if(!inv[i])f[++pct]=i,id[i]=qpow(i,k);
for(int j=1;i*f[j]<=n&&j<=pct;++j){
inv[i*f[j]]=1,id[i*f[j]]=1ll*id[i]*id[f[j]]%mod;
if(i%f[j]==0)break;
}
}
inv[1]=1;
for(int i=2;i<=n;++i)inv[i]=1ll*inv[mod%i]*(mod-mod/i)%mod;
}
}
using namespace math;

namespace solve1{
void main(){
f[0]=1;
for(int i=1;i<=n;++i)f[i]=1ll*(mod+1-m)*f[i-1]%mod;
for(int i=0,j=1,y=1;i<=n;++i){
fmod(ans+=1ll*j*y%mod*f[n-i]%mod*id[i]%mod);
j=1ll*j*(n-i)%mod*inv[i+1]%mod,y=1ll*y*m%mod;
}
printf("%d\n",ans);
exit(0);
}
}

namespace solve2{
void main(){
f[k]=1;
for(int i=k-1,x=0,z=1,y=1;i>=0;--i){
++x,z=1ll*z*(x+n-k-1)%mod*inv[x]%mod,y=1ll*y*(mod-m)%mod;
fmod(f[i]=1ll*(mod+1-m)*f[i+1]%mod),fmod(f[i]+=1ll*y*z%mod);
}
for(int i=0,j=1,y=1;i<=k;++i){
fmod(ans+=1ll*id[i]*j%mod*y%mod*f[i]%mod);
j=1ll*j*(n-i)%mod*inv[i+1]%mod,y=1ll*y*m%mod;
}
printf("%d\n",ans);
exit(0);
}
}
int main()
{
n=read(1),m=qpow(read(1),mod-2),k=read(1);
initmath(min(n,k)+3);
if(n<=k)solve1::main();
else solve2::main();
}

T半天竟然是引用耗时大 :)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
#include <bits/stdc++.h>
/*
#include<ext/pb_ds/assoc_container.hpp>
#include<ext/pb_ds/hash_policy.hpp>
*/
using namespace std;

const double eps = 1e-10;
const double pi = 3.1415926535897932384626433832795;
const double eln = 2.718281828459045235360287471352;

#define f(i, a, b) for (int i = a; i <= b; i++)
#define scan(x) scanf("%d", &x)
#define mp make_pair
#define pb push_back
#define lowbit(x) (x&(-x))

#define fi first
#define se second
#define SZ(x) int((x).size())
#define all(x) x.begin(), x.end()
#define rall(x) x.rbegin(), x.rend()
#define summ(a) (accumulate(all(a), 0ll))

typedef unsigned long long ull;
typedef pair<int,int> pii;
typedef vector<int> vi;

using ll=long long;
const int N=3e3+9;
int tc,n,m1,m2,s,t,ss,tt,hd[N],cur[N],flow[N],dep[N],sc[N],tot,totp,sum;
struct Edge{
int to,nxt,v;
}e[4*N];
inline void add(int u,int v,int w){
e[++tot]={v,hd[u],w},hd[u]=tot;
e[++tot]={u,hd[v],0},hd[v]=tot;
}

bool bfs(int st,int en){
for(int i=0;i<=totp;++i)dep[i]=0;
queue<int>q;
memcpy(cur,hd,sizeof(hd));
q.push(st);
dep[st]=1;
while(!q.empty()){
auto u=q.front();
q.pop();
for(int eg=hd[u];eg;eg=e[eg].nxt){
if(!dep[e[eg].to]&&e[eg].v>0){
dep[e[eg].to]=dep[u]+1;
q.push(e[eg].to);
}
}
}
return !!dep[en];
}

int dfs(int u,int en,int flow){
if(u==en)return flow;
int r=flow;
for(int eg=cur[u];eg&&r;eg=e[eg].nxt){
cur[u]=eg;
if(e[eg].v>0&&dep[e[eg].to]==dep[u]+1){
int c=dfs(e[eg].to,en,min(r,e[eg].v));
r-=c;
e[eg].v-=c,e[eg^1].v+=c;
}
}
return flow-r;
}

int dinic(int st,int en){
int ret=0;
while(bfs(st,en)){
ret+=dfs(st,en,0x3f3f3f3f);
}
return ret;
}

bool solve(){
tot=1,sum=0;
cin>>n>>m1>>m2;
s=n+m2+1,t=s+1,tt=t+1,totp=tt;
for(int i=0;i<=totp;++i)hd[i]=0,sc[i]=0,flow[i]=0;
for(int i=1,x,y,z;i<=m1;++i){
cin>>x>>y>>z;
if(z==1)sc[x]+=1;
else sc[y]+=1;
}
for(int i=1,x,y;i<=m2;++i){
cin>>x>>y;
if(x==1||y==1)sc[1]+=1;
else{
add(s,n+i,1);
flow[s]-=1,flow[n+i]+=1;
add(n+i,x,1),add(n+i,y,1);
}
}
for(int i=2;i<=n;++i)if(sc[i]>sc[1])return false;
for(int i=2;i<=n;++i){
add(i,t,sc[1]-sc[i]);
}
for(int i=n+1;i<tt;++i){
if(flow[i]>0){
sum+=flow[i];
add(ss,i,flow[i]);
}else if(flow[i]<0)add(i,tt,-flow[i]);
}
add(t,s,0x3f3f3f3f);
return sum==dinic(ss,tt);
}

int main()
{
ios::sync_with_stdio(false);
cin.tie(0);
cin>>tc;
f(sb,1,tc){
cout<<(solve()?"YES\n":"NO\n");
}
return 0;
}
0%