商的常数项和余数不对,改了好久也没发现问题
我自己造的数据:
输入:
6 2
43 26 52 33 67 15 25
10 1 5
正确输出:
4 2 3 2 5
3 2
我的输出:
135761236 2 3 2 5
638876389 862483123
解释:
(5x4+2x3+3x2+2x+4)(5x2+x+10)+2x+3
=25x6+15x5+67x4+33x3+52x2+26x+43
C++代码
码风清奇,不喜勿喷
#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const ll maxn=2097153;
const double pi=acos(-1.0);
const ll mod=998244353;
const ll g=3;
const ll gi=332748118;
ll qpow(ll a,ll b,ll p) {
ll res=1;
for (; b; a=a*a%p,b>>=1) if (b&1) res=res*a%p;
return res;
}
ll a[maxn],b[maxn],c[maxn];
ll rev[maxn],len,lim=1;
void NTT(ll*a,ll n,ll op) {
for(ll i=0; i<=n; i++) rev[i]=(rev[i>>1]>>1)|((i&1)<<(len-1));
for(ll i=0; i<=n; i++) if(i<rev[i]) swap(a[i],a[rev[i]]);
for(ll i=1; i<=n; i<<=1) {
ll w=qpow((op==1)?g:gi,(mod-1)/i,mod);
for(ll j=0,p=i/2; j+i-1<=n; j+=i) {
ll wn=1;
for(ll k=j; k<j+p; k++,wn=wn*w%mod) {
ll u=a[k],v=wn*a[k+p]%mod;
a[k]=(u+v)%mod;
a[k+p]=(u-v+mod)%mod;
}
}
}
if(op==-1) {
ll inv=qpow(lim,mod-2,mod);
for(ll i=0; i<=n; i++) a[i]=a[i]*inv%mod;
}
}
void inv(ll*a,ll*b,ll n) {
if(n==1) {
b[0]=qpow(a[0],mod-2,mod);
return;
}
inv(a,b,n>>1);
len=0,lim=1;
while(lim<=(n<<1)) lim<<=1,len++;
for(ll i=1; i<=lim; i++) rev[i]=(rev[i>>1]>>1)|((i&1)<<(len-1));
for(ll i=0; i<=n; i++) c[i]=a[i];
for(ll i=n+1; i<=lim; i++) c[i]=0;
NTT(c,lim,1),NTT(b,lim,1);
for(ll i=0; i<=lim; i++) b[i]=(2ll-c[i]*b[i]%mod+mod)%mod*b[i]%mod;
NTT(b,lim,-1);
for(ll i=n+1; i<=lim; i++) b[i]=0;
}
void mul(ll*a,ll*b,ll n,ll m) {
len=0,lim=1;
while(lim<=n+m) lim<<=1,len++;
NTT(a,lim,1),NTT(b,lim,1);
for(ll i=0; i<=lim; i++) a[i]=a[i]*b[i]%mod;
NTT(a,lim,-1);
}
ll FR[maxn],GR[maxn],GI[maxn];
void div(ll *a,ll *b,ll n,ll m,ll *Q,ll* R) {
for (ll i=0; i<=n; i++) FR[n-i]=a[i];
for (ll i=0; i<=m; i++) GR[m-i]=b[i];
for (int i=n-m+1;i<=m;i++) GR[i]=0;
inv(GR,GI,n-m);
mul(FR,GI,n,n-m);
for (ll i=0;i<=n-m;i++) Q[i]=FR[n-m-i];
for (ll i=0;i<=n-m;i++) cout<<Q[i]<<" ";
cout<<"\n";
mul(b,Q,m,n-m);
for (int i=0;i<m;i++) R[i]=(a[i]-b[i]+mod)%mod;
for (int i=0;i<m;i++) cout<<R[i]<<" ";
}
ll Q[maxn],R[maxn];
ll n,m;
int main() {
ios::sync_with_stdio(false);
cin>>n>>m;
for (int i=0; i<=n; i++) cin>>a[i];
for (int i=0; i<=m; i++) cin>>b[i];
div(a,b,n,m,Q,R);
return 0;
}
各位dalao帮帮我吧