// 1004 high precision multiply: scalar 3-mod NTT (Montgomery), base 1e9, NTT 2^18
#include <sys/auxv.h>
#include <stdint.h>
#include <string.h>
typedef uint64_t u64;
typedef uint32_t u32;
typedef __uint128_t u128;
struct DuckInfo {
uint64_t abi_version;
const char *stdin_ptr; uint64_t stdin_size;
char *stdout_ptr; uint64_t stdout_limit; uint64_t stdout_size;
char *stderr_ptr; uint64_t stderr_limit; uint64_t stderr_size;
const char *IB_ptr; uint64_t IB_limit;
char *OB_ptr; uint64_t OB_limit;
uint64_t tsc_frequency;
} __attribute__((packed));
static const u32 MODS[3] = {998244353u, 1004535809u, 469762049u};
static const u32 GROOT[3] = {3u, 3u, 3u};
static const u32 NINV[3] = {998244351u, 1004535807u, 469762047u};
static const u32 R2[3] = {932051910u, 542374313u, 460175152u};
#define SIZE (1<<18)
static u32 A[SIZE], B[SIZE];
static u32 roots[SIZE], iroots[SIZE];
static u32 r0[SIZE], r1[SIZE], r2[SIZE];
static u32 limbsA[111112], limbsB[111112];
static u32 outlimbs[222224];
static inline u32 mont_mul(u32 a, u32 b, u32 mod, u32 ninv) {
u64 t = (u64)a * b;
u32 m = (u32)t * ninv;
u32 u = (u32)((t + (u64)m * mod) >> 32);
if (u >= mod) u -= mod;
return u;
}
static u32 powmod(u64 a, u64 e, u32 mod) {
u64 r = 1, b = a % mod;
while (e) { if (e & 1) r = r * b % mod; b = b * b % mod; e >>= 1; }
return (u32)r;
}
// forward DIF (natural in, bit-reversed out), values & roots in Montgomery form
static inline void ntt_fwd(u32* x, const u32* rts, u32 mod, u32 ninv) {
int n = SIZE;
for (int len = n; len > 1; len >>= 1) {
int half = len >> 1;
int step = n / len;
for (int i = 0; i < n; i += len) {
u32* y = x + i;
for (int j = 0; j < half; j++) {
u32 u = y[j];
u32 v = y[j + half];
u32 s = u + v; if (s >= mod) s -= mod;
u32 d = u - v; if (d > mod) d += mod;
y[j] = s;
y[j + half] = mont_mul(d, rts[j * step], mod, ninv);
}
}
}
}
// inverse DIT (bit-reversed in, natural out), values & roots in Montgomery form
static inline void ntt_inv(u32* x, const u32* rts, u32 mod, u32 ninv) {
int n = SIZE;
for (int len = 2; len <= n; len <<= 1) {
int half = len >> 1;
int step = n / len;
for (int i = 0; i < n; i += len) {
u32* y = x + i;
for (int j = 0; j < half; j++) {
u32 u = y[j];
u32 v = mont_mul(y[j + half], rts[j * step], mod, ninv);
u32 s = u + v; if (s >= mod) s -= mod;
u32 d = u - v; if (d > mod) d += mod;
y[j] = s;
y[j + half] = d;
}
}
}
}
static char tab3[1000][3];
static void build_tab3(void) {
for (int i = 0; i < 1000; i++) {
int v = i;
tab3[i][2] = '0' + v % 10; v /= 10;
tab3[i][1] = '0' + v % 10; v /= 10;
tab3[i][0] = '0' + v % 10;
}
}
#ifdef LOCAL_TEST
extern uintptr_t jd_getauxval(uintptr_t);
int jd_main() {
DuckInfo* di = (DuckInfo*)jd_getauxval(0x6b637564ull);
#else
int main() {
DuckInfo* di = (DuckInfo*)getauxval(0x6b637564ull);
#endif
const char* in = di->stdin_ptr;
u64 inlen = di->stdin_size;
const char* p = in;
const char* inend = in + inlen;
while (p < inend && (*p == '\n' || *p == '\r' || *p == ' ' || *p == '\t')) p++;
const char* a_start = p;
while (p < inend && *p >= '0' && *p <= '9') p++;
const char* a_end = p;
while (p < inend && (*p == '\n' || *p == '\r' || *p == ' ' || *p == '\t')) p++;
const char* b_start = p;
while (p < inend && *p >= '0' && *p <= '9') p++;
const char* b_end = p;
const char* sa = a_start;
while (sa < a_end - 1 && *sa == '0') sa++;
const char* sb = b_start;
while (sb < b_end - 1 && *sb == '0') sb++;
int na = 0, nb = 0;
{
const char* pos = a_end;
while (pos > sa) {
const char* st = pos - 9; if (st < sa) st = sa;
u32 v = 0;
for (const char* q = st; q < pos; q++) v = v * 10 + (u32)(*q - '0');
limbsA[na++] = v;
pos = st;
}
}
{
const char* pos = b_end;
while (pos > sb) {
const char* st = pos - 9; if (st < sb) st = sb;
u32 v = 0;
for (const char* q = st; q < pos; q++) v = v * 10 + (u32)(*q - '0');
limbsB[nb++] = v;
pos = st;
}
}
for (int mi = 0; mi < 3; mi++) {
u32 mod = MODS[mi], ninv = NINV[mi], r2v = R2[mi];
u64 g = GROOT[mi];
u64 w = powmod(g, (mod - 1) / SIZE, mod);
u64 invw = powmod(w, mod - 2, mod);
u64 cur = 1;
for (int i = 0; i < SIZE; i++) { roots[i] = mont_mul((u32)cur, r2v, mod, ninv); cur = cur * w % mod; }
cur = 1;
for (int i = 0; i < SIZE; i++) { iroots[i] = mont_mul((u32)cur, r2v, mod, ninv); cur = cur * invw % mod; }
memset(A, 0, sizeof(A));
memset(B, 0, sizeof(B));
for (int i = 0; i < na; i++) { u32 v = limbsA[i]; while (v >= mod) v -= mod; A[i] = mont_mul(v, r2v, mod, ninv); }
for (int i = 0; i < nb; i++) { u32 v = limbsB[i]; while (v >= mod) v -= mod; B[i] = mont_mul(v, r2v, mod, ninv); }
ntt_fwd(A, roots, mod, ninv);
ntt_fwd(B, roots, mod, ninv);
for (int i = 0; i < SIZE; i++) A[i] = mont_mul(A[i], B[i], mod, ninv);
ntt_inv(A, iroots, mod, ninv);
u32 ninvn = powmod(SIZE, mod - 2, mod);
u32* dst = r0;
if (mi == 1) dst = r1;
else if (mi == 2) dst = r2;
for (int i = 0; i < SIZE; i++) {
u32 v = mont_mul(A[i], 1, mod, ninv);
dst[i] = (u32)((u64)v * ninvn % mod);
}
}
// CRT reconstruction (Garner), carry in base 1e9
{
u64 M0 = MODS[0];
u64 M1 = MODS[1];
u64 INV01 = powmod(M0 % MODS[1], MODS[1] - 2, MODS[1]);
u128 M01 = (u128)M0 * M1;
u64 INV012 = powmod((u64)(M01 % MODS[2]), MODS[2] - 2, MODS[2]);
int outlen = na + nb - 1;
u128 carry = 0;
int nout = 0;
for (int i = 0; i < outlen; i++) {
u128 c = r0[i];
u64 t1 = (u64)(r1[i] - (u32)(c % MODS[1]) + MODS[1]) % MODS[1];
t1 = (u64)((u128)t1 * INV01 % MODS[1]);
c += (u128)t1 * M0;
u64 t2 = (u64)(r2[i] - (u32)(c % MODS[2]) + MODS[2]) % MODS[2];
t2 = (u64)((u128)t2 * INV012 % MODS[2]);
c += (u128)t2 * M01;
c += carry;
outlimbs[i] = (u32)(c % 1000000000u);
carry = c / 1000000000u;
}
while (carry) { outlimbs[outlen++] = (u32)(carry % 1000000000u); carry /= 1000000000u; }
int hi = outlen - 1;
while (hi > 0 && outlimbs[hi] == 0) hi--;
build_tab3();
char* out = di->stdout_ptr;
char* o = out;
{
u32 v = outlimbs[hi];
char tmp[10]; int t = 0;
do { tmp[t++] = '0' + (char)(v % 10); v /= 10; } while (v);
while (t > 0) *o++ = tmp[--t];
}
for (int i = hi - 1; i >= 0; i--) {
u32 v = outlimbs[i];
u32 g2 = v / 1000000u;
u32 g1 = (v / 1000u) % 1000u;
u32 g0 = v % 1000u;
const char* q;
q = tab3[g2]; *o++ = q[0]; *o++ = q[1]; *o++ = q[2];
q = tab3[g1]; *o++ = q[0]; *o++ = q[1]; *o++ = q[2];
q = tab3[g0]; *o++ = q[0]; *o++ = q[1]; *o++ = q[2];
}
*o++ = '\n';
di->stdout_size = (u64)(o - out);
}
#ifdef LOCAL_TEST
return 0;
#else
asm volatile("mov $60, %%eax; xor %%edi, %%edi; syscall" ::: "rax", "rdi", "rcx", "r11", "memory");
__builtin_unreachable();
#endif
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 125.483 ms | 10 MB + 628 KB | Accepted | Score: 100 | 显示更多 |