// 1004a 高精度乘法2 - 3-prime NTT, base 1e9
#include <stdio.h>
#include <stdint.h>
#include <string.h>
#include <stdlib.h>
typedef uint32_t u32;
typedef uint64_t u64;
typedef __uint128_t u128;
static const u64 MODS[3] = {998244353ULL, 1004535809ULL, 469762049ULL};
static const u64 GROOT[3] = {3ULL, 3ULL, 3ULL};
static inline u64 mulmod(u64 a, u64 b, u64 m){ return (u64)((u128)a*b % m); }
static u64 powmod(u64 a, u64 e, u64 m){
u64 r=1; a%=m;
while(e){ if(e&1) r=(u64)((u128)r*a%m); a=(u64)((u128)a*a%m); e>>=1; }
return r;
}
#define NMAX 4096
static u32 rev[NMAX];
static u32 ra[NMAX], rb[NMAX], r0[NMAX], r1[NMAX], r2[NMAX];
// in-place iterative NTT, n power of 2, root = primitive n-th root mod m
static void ntt(u32* a, int n, int inv, u64 m, u64 root){
for(int i=0;i<n;i++) if(i<rev[i]){ u32 t=a[i]; a[i]=a[rev[i]]; a[rev[i]]=t; }
for(int len=2; len<=n; len<<=1){
u64 wlen = powmod(root, (m-1)/len, m);
if(inv) wlen = powmod(wlen, m-2, m);
for(int i=0;i<n;i+=len){
u64 w=1;
int half=len>>1;
for(int j=0;j<half;j++){
u64 u=a[i+j];
u64 v=(u64)a[i+j+half]*w % m;
u64 x = u+v; if(x>=m) x-=m;
u64 y = u+m-v; if(y>=m) y-=m;
a[i+j]=(u32)x; a[i+j+half]=(u32)y;
w = w*wlen % m;
}
}
}
if(inv){
u64 ninv = powmod(n, m-2, m);
for(int i=0;i<n;i++) a[i]=(u32)((u64)a[i]*ninv % m);
}
}
static char outbuf[20016];
static int outpos=0;
// table for 3-digit decimal strings
static const char* dec3(unsigned v){
static char t[4];
t[0] = '0' + v/100;
t[1] = '0' + (v/10)%10;
t[2] = '0' + v%10;
t[3] = 0;
return t;
}
int main(){
int n = NMAX;
for(int i=1;i<n;i++) rev[i] = (rev[i>>1]>>1) | ((i&1)? (n>>1):0);
// read input (two whitespace-separated tokens)
static char bufA[10016], bufB[10016];
int la=0, lb=0;
{
int c;
while((c=getchar())==' '||c=='\n'||c=='\r'||c=='\t'){}
while(c!=EOF && c!=' '&&c!='\n'&&c!='\r'&&c!='\t'){ if(la<10008) bufA[la++]=c; c=getchar(); }
while(c!=EOF && (c==' '||c=='\n'||c=='\r'||c=='\t')) c=getchar();
while(c!=EOF && c!=' '&&c!='\n'&&c!='\r'&&c!='\t'){ if(lb<10008) bufB[lb++]=c; c=getchar(); }
}
// strip leading zeros (keep at least 1 digit)
int sa=0; while(sa<la-1 && bufA[sa]=='0') sa++;
int sb=0; while(sb<lb-1 && bufB[sb]=='0') sb++;
// parse into base 1e9 limbs (little endian)
static u32 A[1112], B[1112];
int na=0, nb=0;
{
int pos = sa + la;
while(pos > sa){
int start = pos-9; if(start < sa) start = sa;
u32 v=0;
for(int i=start;i<pos;i++) v = v*10 + (bufA[i]-'0');
A[na++] = v; pos = start;
}
}
{
int pos = sb + lb;
while(pos > sb){
int start = pos-9; if(start < sb) start = sb;
u32 v=0;
for(int i=start;i<pos;i++) v = v*10 + (bufB[i]-'0');
B[nb++] = v; pos = start;
}
}
// CRT constants
u64 INV01 = powmod(MODS[0]%MODS[1], MODS[1]-2, MODS[1]);
u64 M01 = (u64)((u128)MODS[0]*MODS[1] % MODS[2]);
u64 INV012 = powmod(M01, MODS[2]-2, MODS[2]);
u128 M01full = (u128)MODS[0]*MODS[1];
for(int p=0;p<3;p++){
u64 m = MODS[p], rt = GROOT[p];
for(int i=0;i<n;i++){ ra[i]=0; rb[i]=0; }
for(int i=0;i<na;i++) ra[i] = A[i] % m;
for(int i=0;i<nb;i++) rb[i] = B[i] % m;
ntt(ra, n, 0, m, rt);
ntt(rb, n, 0, m, rt);
for(int i=0;i<n;i++) ra[i] = (u32)((u64)ra[i]*rb[i] % m);
ntt(ra, n, 1, m, rt);
if(p==0) memcpy(r0, ra, n*sizeof(u32));
else if(p==1) memcpy(r1, ra, n*sizeof(u32));
else memcpy(r2, ra, n*sizeof(u32));
}
// reconstruct coefficients and carry in base 1e9
int outlen = na + nb - 1; // max meaningful limbs
// full coefficient array (u128)
// carry: process each coefficient
u128 carry = 0;
// We'll generate output limbs into a buffer, most significant first later.
// Instead build decimal digits directly from LSD.
static u64 outlimbs[2224];
int nout = 0;
for(int i=0;i<outlen;i++){
u128 c = r0[i];
u128 t = ((u128)(r1[i] - (u64)(c % MODS[1]) + MODS[1]) % MODS[1]) * INV01 % MODS[1];
c += t * MODS[0];
u128 t2 = ((u128)(r2[i] - (u64)(c % MODS[2]) + MODS[2]) % MODS[2]) * INV012 % MODS[2];
c += t2 * M01full;
c += carry;
outlimbs[i] = (u64)(c % 1000000000ULL);
carry = c / 1000000000ULL;
}
while(carry){ outlimbs[outlen++] = (u64)(carry % 1000000000ULL); carry /= 1000000000ULL; }
// outlimbs[i] is LSD-first now with outlen limbs
// find highest non-zero limb
int hi = outlen-1;
while(hi>0 && outlimbs[hi]==0) hi--;
// output
// print most significant limb without padding
char tmp[16];
int tl = sprintf(tmp, "%llu", (unsigned long long)outlimbs[hi]);
fwrite(tmp, 1, tl, stdout);
for(int i=hi-1;i>=0;i--){
u64 v = outlimbs[i];
// 9 digits, 3 groups of 3
u32 g0 = v % 1000;
u32 g1 = (v/1000) % 1000;
u32 g2 = (v/1000000);
// write g2,g1,g0 as 3 digits each
putchar('0'+g2/100); putchar('0'+(g2/10)%10); putchar('0'+g2%10);
putchar('0'+g1/100); putchar('0'+(g1/10)%10); putchar('0'+g1%10);
putchar('0'+g0/100); putchar('0'+(g0/10)%10); putchar('0'+g0%10);
}
putchar('\n');
return 0;
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 3.923 ms | 172 KB | Accepted | Score: 100 | 显示更多 |