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];
}