萌新求助NTT
查看原帖
萌新求助NTT
577796
prokali楼主2023/7/11 10:35

RT,全部WA,求大佬帮助

#include<bits/stdc++.h>
#define ll long long
#define db double
using namespace std;
const ll MOD=998244353;
ll len;
string ab;
ll n_4[4000005],n_3[4000005],n_2[4000005],n_1[4000005],a[4000005];
ll rev[4000005];
long long ans[4000005],ta[4000005];
ll qpow(ll a,ll b)
{
	ll ans=1;
	while(b)
	{
		if(b&1ll)ans=ans*a%MOD;
		a=a*a%MOD;b>>=1;
	}
	return ans;
}
void div(ll *a,ll len,ll pm)
{
	for(ll i=0;i<len;i++)
		if(i<rev[i])
			swap(a[i],a[rev[i]]);
	for(ll mid=1;mid<len;mid*=2ll)
	{
		ll wn=qpow(3ll,(MOD-1)/(mid*2ll));
		if(pm==-1)wn=qpow(wn,MOD-2);
		for(ll j=0;j<len;j+=2ll*mid)
		{
			ll w=1;
			for(ll k=0;k<mid;k++,w=w*wn%MOD)
			{
				ll x=a[k+j],y=w*a[j+k+mid]%MOD;
				a[k+j]=(x+y)%MOD;
				a[j+k+mid]=(x-y+MOD)%MOD;
			}
		}
	}
}
int main()
{
	ios::sync_with_stdio(0);cin.tie(0);cout.tie(0);
	cin>>ab;
	ll n=ab.size();
	n_4[0]=(ab[n-1ll]-'0'+4ll);
	n_3[0]=(ab[n-1ll]-'0'+3ll);
	n_2[0]=(ab[n-1ll]-'0'+2ll);
	n_1[0]=(ab[n-1ll]-'0'+1ll);
	for(ll i=n-2;i>-1;--i)
		n_4[n-1ll-i]=(ab[i]-'0'),n_3[n-1ll-i]=(ab[i]-'0'),n_2[n-1ll-i]=(ab[i]-'0'),n_1[n-1ll-i]=(ab[i]-'0');
	//++n;
	len=1;
	ll sum=0;
	while(len<=4*n)len<<=1ll,++sum;
	for(int i=0;i<=len;i++)
		rev[i]=(rev[i>>1ll]>>1ll|(i&1ll)<<(sum-1ll));
	div(n_4,len,1);
	div(n_3,len,1);
	div(n_2,len,1);
	div(n_1,len,1);
	//-----
	for(ll i=0;i<=len;++i)
		a[i]=1,a[i]*=n_4[i],a[i]%=MOD;
	div(a,len,-1);
	for(ll i=0;i<=len;++i)
		a[i]=(a[i]/len);
	for(ll i=0;i<=len;++i)
	if(a[i]>=10)
		a[i+1]=(a[i]/10ll+a[i+1]),a[i]=(a[i]%10ll);
	//-----
	div(a,len,1);
	for(ll i=0;i<=len;++i)
		a[i]*=n_3[i],a[i]%=MOD;
	div(a,len,-1);
	for(ll i=0;i<=len;++i)
		a[i]=(a[i]/len);
	for(ll i=0;i<=len;++i)
	if(a[i]>=10)
		a[i+1]=(a[i]/10ll+a[i+1]),a[i]=(a[i]%10ll);
	//-----
	div(a,len,1);
	for(ll i=0;i<=len;++i)
		a[i]*=n_2[i],a[i]%=MOD;
	div(a,len,-1);
	for(ll i=0;i<=len;++i)
		a[i]=((a[i])/len);
	for(ll i=0;i<=len;++i)
	if(a[i]>=10)
		a[i+1]=(a[i]/10ll+a[i+1]),a[i]=(a[i]%10ll);
	div(a,len,1);
	for(ll i=0;i<=len;++i)
		a[i]*=n_1[i],a[i]%=MOD;
	div(a,len,-1);n*=4;
	for(ll i=0;i<=n;++i)
		ans[i]=(a[i]/len); 
	//cout<<ans[0]<<"????\n";
	for(ll i=0;i<n;++i)
		ans[i+1]+=ans[i]/10,ans[i]%=10;
	while(ans[n]>=10)
	{
		++n;
		ans[n]=ans[n-1]/10;
		ans[n-1]%=10;
	}
	bool flag=0;
	ll wz;
	for(ll i=n;i>-1;--i)
	{
	if(ans[i]!=0)
	{
		wz=i;break;
	}
	}
	for(ll i=wz;i>-1;--i)
	{
		ta[i]=ans[i]/24;
        ans[i]%=24;
        ans[i-1]+=ans[i]*10;
	}
	while(wz&&ta[wz]==0)--wz;
	for(ll i=wz;i>-1;--i)cout<<ta[i];
	
} 
2023/7/11 10:35
加载中...