RE+WA 60pts 求助
查看原帖
RE+WA 60pts 求助
203008
山田リョウ楼主2022/8/19 21:58
// Problem: P5273 【模板】多项式幂函数(加强版)
// Contest: Luogu
// URL: https://www.luogu.com.cn/problem/P5273
// Memory Limit: 128 MB
// Time Limit: 1000 ms

#include<stdio.h>
#include<ctype.h>
#include<random>
namespace fasti{
	char buf[1<<21],*p1=buf,*p2=buf;
	inline char getc(){return p1==p2&&(p2=(p1=buf)+fread(buf,1,1<<21,stdin),p1==p2)?EOF:*(p1++);}
	inline void read(int&x){
    	char c=getc(),f=0;
    	for(;!isdigit(c);c=getc())f^=!(c^'-');
    	for(x=0;isdigit(c);c=getc())x=x*10+(c^48);
    	if(f)x=-x;
	}
	inline void read(char*s){
		char c=getc();int n=0;
		for(;!isdigit(c);c=getc());
		for(;s[n++]=c,isdigit(c=getc()););
		s[n]='\0';
	}
	template<typename... Args>
	inline void read(int&x,Args&...args){read(x),read(args...);}
}
using fasti::getc;
using fasti::read;
namespace fasto{
	char obuf[1<<21],stk[20];int o_p=-1,tp;
	inline void flush(){fwrite(obuf,1,o_p+1,stdout),o_p=-1;}
	inline void print(char c){if(o_p==(1<<21)-1)flush();obuf[++o_p]=c;}
	void print(int x){
		for(stk[tp=0]=x%10;x>9;stk[++tp]=x%10)x/=10;
		for(;~tp;)print((char)(stk[tp--]|48));
	}
	template<typename T,typename... Args>
	void print(T x,Args... args){print(x),print(args...);}
}
using fasto::flush;
using fasto::print;
const int p=998244353,g=3,maxn=4194304;
inline int sum(int x,int y,int P=p){return (P-x)>y?x+y:y-(P-x);}
inline int dif(int x,int y,int P=p){return x<y?P-(y-x):x-y;}
inline int pro(int x,int y,int P=p){return (long long)x*y%P;}
inline int pow(int x,int y,int P=p){int res=1;for(;y;y>>=1,x=pro(x,x,P))if(y&1)res=pro(res,x,P);return res;}
inline void swap(int&a,int&b){int t=a;a=b;b=t;}
int sqr(int n){
    static std::mt19937 myrand(114514);
    int a,b=1,x=1,y=0;for(;a=myrand()%p,pow(dif(pro(a,a),n),p-1>>1)==1;);
    int res=dif(pro(a,a),n);
    auto mult=[p,res](int&a1,int&a2,int b1,int b2)->void{
    	int t1=sum(pro(a1,b1),pro(res,pro(a2,b2))),t2=sum(pro(a1,b2),pro(a2,b1));
    	a1=t1,a2=t2;
    };
    for(int i=(p+1>>1);i;i>>=1,mult(a,b,a,b))if(i&1)mult(x,y,a,b);
    return x>(p-x)?(p-x):x;
}
void cpy(int*a,int*b,int n){for(int i=0;i<n;++i)a[i]=b[i];}
void clr(int*a,int n){for(int i=0;i<n;++i)a[i]=0;}
int rev[maxn];
inline void init(int n){
	int l=__builtin_ctz(n);
	for(int i=0;i<n;++i)rev[i]=rev[i>>1]>>1|(i&1)<<l-1;
}
void NTT(int*a,int n,int op){
	init(n);
	for(int i=1;i<n;++i)if(i<rev[i])swap(a[i],a[rev[i]]);
	for(int i=1;i<n;i<<=1){
		int wn=pow(g,op?p-1-((p-1)/i>>1):(p-1)/i>>1);
		for(int j=0;j<n;j+=(i<<1)){
			int w=1;
			for(int k=j;k<i+j;++k,w=pro(w,wn)){
				int x=a[k],y=pro(a[i|k],w);
				a[k]=sum(x,y),a[i|k]=dif(x,y);
			}
		}
	}
	if(op){int inv=pro((p-1)/n,p-1);for(int i=0;i<n;++i)a[i]=pro(a[i],inv);}
}
inline void padd(int*a,int*b,int n,int m){
	for(int i=0;i<n&&i<m;++i)a[i]=sum(a[i],b[i]);
	for(;n<m;++n)a[n]=b[n];
}
inline void pdif(int*a,int*b,int n,int m){
	for(int i=0;i<n&&i<m;++i)a[i]=dif(a[i],b[i]);
	for(;n<m;++n)a[n]=dif(0,b[n]);
}
void pmul(int*a,int*b,int n,int m){
	int t=1;
	for(;t<(n+m-1);t<<=1);
	clr(a+n,t-n),clr(b+m,t-m);
	NTT(a,t,0),NTT(b,t,0);
	for(int i=0;i<t;++i)a[i]=pro(a[i],b[i]);
	NTT(a,t,1),NTT(b,t,1);
}
void pinv(int*a,int*b,int n){
	static int c[maxn],d[maxn];
	b[0]=pow(a[0],p-2);int i=2;
	for(;i<(n<<1);i<<=1){
		clr(b+(i>>1),i>>1),cpy(c,b,i>>1),cpy(d,a,i);
		NTT(c,i,0),NTT(d,i,0);
		for(int j=0;j<i;++j)c[j]=pro(c[j],d[j]);
		NTT(c,i,1);
		clr(c,i>>1),c[0]=1,cpy(d,b,i);
		NTT(d,i,0),NTT(c,i,0);
		for(int j=0;j<i;++j)c[j]=pro(c[j],d[j]);
		NTT(c,i,1);
		for(int j=(i>>1);j<i;++j)b[j]=dif(sum(b[j],b[j]),c[j]);
	}
	clr(c,i),clr(d,i);
}
void psqr(int*a,int*b,int n){
	static int c[maxn],d[maxn];
	b[0]=sqr(a[0]);
	for(int i=2;i<(n<<1);i<<=1){
		clr(b+(i>>1),i>>1),pinv(b,c,i),cpy(d,a,i);
		clr(c+i,i),clr(d+i,i);
		NTT(c,i<<1,0),NTT(d,i<<1,0);
		for(int j=0;j<(i<<1);++j)c[j]=pro(c[j],d[j]);
		NTT(c,i<<1,1);
		for(int j=(i>>1);j<i;++j)b[j]=pro(c[j],p+1>>1);
	}
}
void pdiv(int*f,int*g,int*q,int*r,int n,int m){
	static int a[maxn],b[maxn];
	int t=n-m+1;
	for(int i=0;i<m&&i<t;++i)a[i]=g[m-1-i];
	pinv(a,b,t);
	for(int i=0;i<t;++i)a[i]=f[n-1-i];
	pmul(a,b,t,t);
	for(int i=0;i<t;++i)q[i]=a[t-1-i];
	pmul(g,q,m,t);
	for(int i=0;i<m-1;++i)r[i]=dif(f[i],g[i]);
}
inline void pder(int*a,int*b,int n){
	for(int i=1;i<n;++i)b[i-1]=pro(i,a[i]);
	b[n-1]=0;
}
inline void pint(int*a,int n,int c=0){
	static int inv[maxn];
	inv[1]=1;for(int i=2;i<n;++i)inv[i]=pro(inv[p%i],dif(0,p/i));
	for(int i=n;--i;)a[i]=pro(inv[i],a[i-1]);
	a[0]=c;
}
void pln(int*a,int*b,int n){
	static int c[maxn];
	pder(a,c,n),pinv(a,b,n);
	pmul(b,c,n,n),clr(b+n,n),pint(b,n);
}
void pexp(int*a,int*b,int n){
	static int c[maxn];
	b[0]=1;
	for(int i=2;i<(n<<1);i<<=1){
		clr(b+(i>>1),i>>1),pln(b,c,i),pdif(c,a,i,i);
		for(int j=0;j<i;++j)c[j]=dif(!j,c[j]);
		pmul(b,c,i>>1,i),clr(b+i,i>>1);
	}
}
void _ppow(int*a,int*b,int k,int n){
    static int c[maxn];
    pln(a,c,n);
    for(int i=0;i<n;++i)c[i]=pro(c[i],k);
    pexp(c,b,n);
}
void ppow(int*a,int*b,long long k,int n){
	int cnt=0;for(;cnt<n&&a[cnt]==0;++cnt);
	if(k*cnt>=n)for(int i=0;i<n;++i)b[i]=0;
	else{
		for(int i=0;i<n;++i)a[i]=(i+cnt<n?a[i+cnt]:0);
		int x=a[0],y=pow(x,p-2),z=pow(x,k%(p-1));
		for(int i=0;i<n;++i)a[i]=pro(a[i],y);
		_ppow(a,b,k%p,n);
		for(int i=0;i<n;++i)a[i]=pro(a[i],x);
		for(int i=0;i<n;++i)b[i]=pro(b[i],z);
		for(int i=n;i--;)a[i]=(i-cnt<0?0:a[i-cnt]);
		cnt*=k;
		for(int i=n;i--;)b[i]=(i-cnt<0?0:b[i-cnt]);
	}
}
void ppow(int*a,int*b,char*k,int n){
	int cnt=0;for(;cnt<n&&a[cnt]==0;++cnt);
	int k1=0,k2=0,k3=0;
	for(int i=0;k[i]!='\0';++i)k1=sum(pro(k1,10),k[i]^48),k2=sum(pro(k2,10,p-1),k[i]^10,p-1);
	for(int i=0;i<8&&k[i]!='\0';++i)k3=k3*10+(k[i]^48);
	if(k3*cnt>=n)for(int i=0;i<n;++i)b[i]=0;
	else{
		for(int i=0;i<n;++i)a[i]=(i+cnt<n?a[i+cnt]:0);
		int x=a[0],y=pow(x,p-2),z=pow(x,k2);
		for(int i=0;i<n;++i)a[i]=pro(a[i],y);
		_ppow(a,b,k1,n);
		for(int i=0;i<n;++i)a[i]=pro(a[i],x);
		for(int i=0;i<n;++i)b[i]=pro(b[i],z);
		for(int i=n;i--;)a[i]=(i-cnt<0?0:a[i-cnt]);
		cnt*=k3;
		for(int i=n;i--;)b[i]=(i-cnt<0?0:b[i-cnt]);
	}
}
char k[100001];
int a[maxn],b[maxn];
int main(){
	int n;
	read(n,k);
	for(int i=0;i<n;++i)read(a[i]);
	ppow(a,b,k,n);
	for(int i=0;i<n;++i)print(b[i],' ');
	flush();
	return 0;
}
2022/8/19 21:58
加载中...