// NTT mod 998244353, radix-8 DIF/DIT, AVX2 Montgomery, contiguous twiddles
#include <immintrin.h>
#pragma GCC target("avx2")
typedef unsigned long long u64;
typedef unsigned u32;
const u32 MOD = 998244353u;
const u32 NINV = 998244351u;
const u32 ONE = 301989884u;
const u32 R2 = 932051910u;
const u32 IROOT = 911660635u; // sqrt(-1)
const u32 ZETA = 372528824u; // primitive 8th root
const u32 ZETA3 = 488723995u; // zeta^3
const u32 ZETA5 = 625715529u; // zeta^5 = -zeta
const u32 ZETA7 = 509520358u; // zeta^7 = -zeta^3
const u32 IMINUS = 86583718u; // -i
const u32 I_p = 3922439030u; // floor(i*2^32/p)
const u32 IM_p = 372528265u; // floor(-i*2^32/p)
const u32 Z1_p = 1602813089u; // floor(zeta*2^32/p)
const u32 Z3_p = 2102745253u; // floor(zeta^3*2^32/p)
const u32 Z5_p = 2692154206u; // floor(zeta^5*2^32/p)
const u32 Z7_p = 2192222042u; // floor(zeta^7*2^32/p)
const int MAXL = 1 << 21;
const int MAXTW = (1 << 21) + 64;
static u32 A[MAXL];
static u32 B[MAXL];
static u32 roots[MAXL];
static u32 roots_inv[MAXL];
static u32 TW_fwd[MAXTW];
static u32 TW_inv[MAXTW];
static inline __m256i load8(const u32* p) { return _mm256_loadu_si256((const __m256i*)p); }
static inline void store8(u32* p, __m256i v) { _mm256_storeu_si256((__m256i*)p, v); }
static inline __m256i hi32_mul(__m256i a, __m256i b) {
__m256i t0 = _mm256_mul_epu32(a, b);
__m256i t1 = _mm256_mul_epu32(_mm256_srli_si256(a, 4), _mm256_srli_si256(b, 4));
return _mm256_or_si256(_mm256_srli_epi64(t0, 32), _mm256_slli_si256(_mm256_srli_epi64(t1, 32), 4));
}
static inline __m256i mont_mul8(__m256i x, __m256i y) {
const __m256i ninv = _mm256_set1_epi32(NINV);
const __m256i p = _mm256_set1_epi32(MOD);
const __m256i pminus1 = _mm256_set1_epi32(MOD - 1);
const __m256i one = _mm256_set1_epi32(1);
const __m256i zero = _mm256_setzero_si256();
__m256i lo = _mm256_mullo_epi32(x, y);
__m256i m = _mm256_mullo_epi32(lo, ninv);
__m256i res = _mm256_add_epi32(hi32_mul(x, y), hi32_mul(m, p));
__m256i eqz = _mm256_cmpeq_epi32(lo, zero);
res = _mm256_add_epi32(res, _mm256_andnot_si256(eqz, one));
__m256i ge = _mm256_cmpgt_epi32(res, pminus1);
return _mm256_sub_epi32(res, _mm256_and_si256(ge, p));
}
// Shoup (a in mont form, w normal, wp=floor(w*2^32/p)) -> (a*w) mont form
static inline __m256i shoup_mul8(__m256i a, __m256i w, __m256i wp) {
const __m256i p = _mm256_set1_epi32(MOD);
const __m256i pminus1 = _mm256_set1_epi32(MOD - 1);
__m256i lo = _mm256_mullo_epi32(a, w);
__m256i q = hi32_mul(a, wp);
__m256i r = _mm256_sub_epi32(lo, _mm256_mullo_epi32(q, p));
__m256i ge = _mm256_cmpgt_epi32(r, pminus1);
return _mm256_sub_epi32(r, _mm256_and_si256(ge, p));
}
static inline __m256i add_mod8(__m256i x, __m256i y) {
const __m256i p = _mm256_set1_epi32(MOD);
const __m256i pminus1 = _mm256_set1_epi32(MOD - 1);
__m256i s = _mm256_add_epi32(x, y);
__m256i ge = _mm256_cmpgt_epi32(s, pminus1);
return _mm256_sub_epi32(s, _mm256_and_si256(ge, p));
}
static inline __m256i sub_mod8(__m256i x, __m256i y) {
const __m256i p = _mm256_set1_epi32(MOD);
__m256i d = _mm256_sub_epi32(x, y);
__m256i mask = _mm256_cmpgt_epi32(y, x);
return _mm256_add_epi32(d, _mm256_and_si256(mask, p));
}
static inline u32 mont_mul(u32 x, u32 y) {
u64 t = (u64)x * y;
u32 m = (u32)t * NINV;
u64 u = (t + (u64)m * MOD) >> 32;
if (u >= MOD) u -= MOD;
return (u32)u;
}
static inline u32 shoup_mul(u32 a, u32 w, u32 wp) {
u32 q = (u32)(((u64)a * wp) >> 32);
u32 r = (u32)((u64)a * w - (u64)q * MOD);
if (r >= MOD) r -= MOD;
return r;
}
static inline u32 add_mod(u32 x, u32 y) { u32 s = x + y; return s >= MOD ? s - MOD : s; }
static inline u32 sub_mod(u32 x, u32 y) { return x >= y ? x - y : x + MOD - y; }
static inline u32 to_mont(u32 x) { return mont_mul(x, R2); }
static u32 mont_pow(u32 base, u64 e) {
u32 r = ONE, b = base;
while (e) { if (e & 1) r = mont_mul(r, b); b = mont_mul(b, b); e >>= 1; }
return r;
}
static void gen_roots(u32* rts, int n, u32 w_mont) {
u32 p[8];
p[0] = ONE;
for (int k = 1; k < 8; k++) p[k] = mont_mul(p[k-1], w_mont);
u32 w8 = mont_mul(p[7], w_mont);
__m256i pv = _mm256_loadu_si256((__m256i*)p);
u32 base = ONE;
for (int i = 0; i < n; i += 8) {
_mm256_storeu_si256((__m256i*)(rts + i), mont_mul8(pv, _mm256_set1_epi32(base)));
base = mont_mul(base, w8);
}
}
static inline __m256i to_mont8(__m256i x) { return mont_mul8(x, _mm256_set1_epi32(R2)); }
// radix-8 twiddle tables: [w1[0..m), w2[0..m), ..., w7[0..m)]
static void build_twiddles8_fwd(const u32* rts, u32* tw, int n) {
int off = 0;
for (int len = n; len >= 8; len >>= 3) {
int m = len >> 3;
int step = n / len;
for (int s = 1; s <= 7; s++) {
u32* w = tw + off + (s - 1) * m;
for (int j = 0; j < m; j++) w[j] = rts[s * j * step];
}
off += 7 * m;
}
}
static void build_twiddles8_inv(const u32* rts, u32* tw, int n) {
int off = 0;
for (int len = 8; len <= n; len <<= 3) {
int m = len >> 3;
int step = n / len;
for (int s = 1; s <= 7; s++) {
u32* w = tw + off + (s - 1) * m;
for (int j = 0; j < m; j++) w[j] = rts[s * j * step];
}
off += 7 * m;
}
}
// forward DIF radix-8: ii=+i, z1=zeta, z3=zeta^3
static void ntt_fwd8(u32 *x, int n, const u32 *tw, u32 ii, u32 z1, u32 z3) {
(void)ii; (void)z1; (void)z3;
int off = 0;
for (int len = n; len >= 8; len >>= 3) {
int m = len >> 3;
if (m >= 8) {
const u32* W[8];
for (int s = 0; s < 7; s++) W[s] = tw + off + s * m;
for (int i = 0; i < n; i += len) {
u32 *y = x + i;
for (int j = 0; j < m; j += 8) {
__m256i x0 = load8(y + j);
__m256i x1 = load8(y + j + m);
__m256i x2 = load8(y + j + 2*m);
__m256i x3 = load8(y + j + 3*m);
__m256i x4 = load8(y + j + 4*m);
__m256i x5 = load8(y + j + 5*m);
__m256i x6 = load8(y + j + 6*m);
__m256i x7 = load8(y + j + 7*m);
// even 4-DFT
__m256i e0 = add_mod8(x0, x4);
__m256i e1 = sub_mod8(x0, x4);
__m256i e2 = add_mod8(x2, x6);
__m256i e3 = sub_mod8(x2, x6);
__m256i ie3 = shoup_mul8(e3, _mm256_set1_epi32(IROOT), _mm256_set1_epi32(I_p));
__m256i E0 = add_mod8(e0, e2);
__m256i E1 = add_mod8(e1, ie3);
__m256i E2 = sub_mod8(e0, e2);
__m256i E3 = sub_mod8(e1, ie3);
// odd 4-DFT
__m256i o0 = add_mod8(x1, x5);
__m256i o1 = sub_mod8(x1, x5);
__m256i o2 = add_mod8(x3, x7);
__m256i o3 = sub_mod8(x3, x7);
__m256i io3 = shoup_mul8(o3, _mm256_set1_epi32(IROOT), _mm256_set1_epi32(I_p));
__m256i O0 = add_mod8(o0, o2);
__m256i O1 = add_mod8(o1, io3);
__m256i O2 = sub_mod8(o0, o2);
__m256i O3 = sub_mod8(o1, io3);
// combine
__m256i zO1 = shoup_mul8(O1, _mm256_set1_epi32(ZETA), _mm256_set1_epi32(Z1_p));
__m256i zO3 = shoup_mul8(O3, _mm256_set1_epi32(ZETA3), _mm256_set1_epi32(Z3_p));
__m256i iO2 = shoup_mul8(O2, _mm256_set1_epi32(IROOT), _mm256_set1_epi32(I_p));
__m256i Y0 = add_mod8(E0, O0);
__m256i Y1 = add_mod8(E1, zO1);
__m256i Y2 = add_mod8(E2, iO2);
__m256i Y3 = add_mod8(E3, zO3);
__m256i Y4 = sub_mod8(E0, O0);
__m256i Y5 = sub_mod8(E1, zO1);
__m256i Y6 = sub_mod8(E2, iO2);
__m256i Y7 = sub_mod8(E3, zO3);
// twiddle outputs 1..7
store8(y + j, Y0);
store8(y + j + m, mont_mul8(Y1, load8(W[0] + j)));
store8(y + j + 2*m, mont_mul8(Y2, load8(W[1] + j)));
store8(y + j + 3*m, mont_mul8(Y3, load8(W[2] + j)));
store8(y + j + 4*m, mont_mul8(Y4, load8(W[3] + j)));
store8(y + j + 5*m, mont_mul8(Y5, load8(W[4] + j)));
store8(y + j + 6*m, mont_mul8(Y6, load8(W[5] + j)));
store8(y + j + 7*m, mont_mul8(Y7, load8(W[6] + j)));
}
}
} else {
// m == 1 scalar
const u32* W[8];
for (int s = 0; s < 7; s++) W[s] = tw + off + s * m;
for (int i = 0; i < n; i += len) {
u32 *y = x + i;
int j = 0;
u32 x0 = y[0], x1 = y[m], x2 = y[2*m], x3 = y[3*m], x4 = y[4*m], x5 = y[5*m], x6 = y[6*m], x7 = y[7*m];
u32 e0 = add_mod(x0,x4), e1 = sub_mod(x0,x4), e2 = add_mod(x2,x6), e3 = sub_mod(x2,x6);
u32 ie3 = shoup_mul(e3, IROOT, I_p);
u32 E0 = add_mod(e0,e2), E1 = add_mod(e1,ie3), E2 = sub_mod(e0,e2), E3 = sub_mod(e1,ie3);
u32 o0 = add_mod(x1,x5), o1 = sub_mod(x1,x5), o2 = add_mod(x3,x7), o3 = sub_mod(x3,x7);
u32 io3 = shoup_mul(o3, IROOT, I_p);
u32 O0 = add_mod(o0,o2), O1 = add_mod(o1,io3), O2 = sub_mod(o0,o2), O3 = sub_mod(o1,io3);
u32 zO1 = shoup_mul(O1,ZETA,Z1_p), zO3 = shoup_mul(O3,ZETA3,Z3_p), iO2 = shoup_mul(O2,IROOT,I_p);
y[0] = add_mod(E0,O0);
y[m] = mont_mul(add_mod(E1,zO1), W[0][0]);
y[2*m] = mont_mul(add_mod(E2,iO2), W[1][0]);
y[3*m] = mont_mul(add_mod(E3,zO3), W[2][0]);
y[4*m] = mont_mul(sub_mod(E0,O0), W[3][0]);
y[5*m] = mont_mul(sub_mod(E1,zO1), W[4][0]);
y[6*m] = mont_mul(sub_mod(E2,iO2), W[5][0]);
y[7*m] = mont_mul(sub_mod(E3,zO3), W[6][0]);
}
}
off += 7 * m;
}
}
// inverse DIT radix-8: ii=-i, z1=zeta^7, z3=zeta^5
static void ntt_inv8(u32 *x, int n, const u32 *tw, u32 ii, u32 z1, u32 z3) {
(void)ii; (void)z1; (void)z3;
int off = 0;
for (int len = 8; len <= n; len <<= 3) {
int m = len >> 3;
if (m >= 8) {
const u32* W[8];
for (int s = 0; s < 7; s++) W[s] = tw + off + s * m;
for (int i = 0; i < n; i += len) {
u32 *y = x + i;
for (int j = 0; j < m; j += 8) {
__m256i y0 = load8(y + j);
__m256i y1 = mont_mul8(load8(y + j + m), load8(W[0] + j));
__m256i y2 = mont_mul8(load8(y + j + 2*m), load8(W[1] + j));
__m256i y3 = mont_mul8(load8(y + j + 3*m), load8(W[2] + j));
__m256i y4 = mont_mul8(load8(y + j + 4*m), load8(W[3] + j));
__m256i y5 = mont_mul8(load8(y + j + 5*m), load8(W[4] + j));
__m256i y6 = mont_mul8(load8(y + j + 6*m), load8(W[5] + j));
__m256i y7 = mont_mul8(load8(y + j + 7*m), load8(W[6] + j));
// inverse 8-DFT (same structure, ii/z1/z3 already inverse)
__m256i e0 = add_mod8(y0, y4);
__m256i e1 = sub_mod8(y0, y4);
__m256i e2 = add_mod8(y2, y6);
__m256i e3 = sub_mod8(y2, y6);
__m256i ie3 = shoup_mul8(e3, _mm256_set1_epi32(IMINUS), _mm256_set1_epi32(IM_p));
__m256i E0 = add_mod8(e0, e2);
__m256i E1 = add_mod8(e1, ie3);
__m256i E2 = sub_mod8(e0, e2);
__m256i E3 = sub_mod8(e1, ie3);
__m256i o0 = add_mod8(y1, y5);
__m256i o1 = sub_mod8(y1, y5);
__m256i o2 = add_mod8(y3, y7);
__m256i o3 = sub_mod8(y3, y7);
__m256i io3 = shoup_mul8(o3, _mm256_set1_epi32(IMINUS), _mm256_set1_epi32(IM_p));
__m256i O0 = add_mod8(o0, o2);
__m256i O1 = add_mod8(o1, io3);
__m256i O2 = sub_mod8(o0, o2);
__m256i O3 = sub_mod8(o1, io3);
__m256i zO1 = shoup_mul8(O1, _mm256_set1_epi32(ZETA7), _mm256_set1_epi32(Z7_p));
__m256i zO3 = shoup_mul8(O3, _mm256_set1_epi32(ZETA5), _mm256_set1_epi32(Z5_p));
__m256i iO2 = shoup_mul8(O2, _mm256_set1_epi32(IMINUS), _mm256_set1_epi32(IM_p));
store8(y + j, add_mod8(E0, O0));
store8(y + j + m, add_mod8(E1, zO1));
store8(y + j + 2*m, add_mod8(E2, iO2));
store8(y + j + 3*m, add_mod8(E3, zO3));
store8(y + j + 4*m, sub_mod8(E0, O0));
store8(y + j + 5*m, sub_mod8(E1, zO1));
store8(y + j + 6*m, sub_mod8(E2, iO2));
store8(y + j + 7*m, sub_mod8(E3, zO3));
}
}
} else {
const u32* W[8];
for (int s = 0; s < 7; s++) W[s] = tw + off + s * m;
for (int i = 0; i < n; i += len) {
u32 *y = x + i;
int j = 0;
u32 y0 = y[0];
u32 y1 = mont_mul(y[m], W[0][0]);
u32 y2 = mont_mul(y[2*m], W[1][0]);
u32 y3 = mont_mul(y[3*m], W[2][0]);
u32 y4 = mont_mul(y[4*m], W[3][0]);
u32 y5 = mont_mul(y[5*m], W[4][0]);
u32 y6 = mont_mul(y[6*m], W[5][0]);
u32 y7 = mont_mul(y[7*m], W[6][0]);
u32 e0 = add_mod(y0,y4), e1 = sub_mod(y0,y4), e2 = add_mod(y2,y6), e3 = sub_mod(y2,y6);
u32 ie3 = shoup_mul(e3, IMINUS, IM_p);
u32 E0 = add_mod(e0,e2), E1 = add_mod(e1,ie3), E2 = sub_mod(e0,e2), E3 = sub_mod(e1,ie3);
u32 o0 = add_mod(y1,y5), o1 = sub_mod(y1,y5), o2 = add_mod(y3,y7), o3 = sub_mod(y3,y7);
u32 io3 = shoup_mul(o3, IMINUS, IM_p);
u32 O0 = add_mod(o0,o2), O1 = add_mod(o1,io3), O2 = sub_mod(o0,o2), O3 = sub_mod(o1,io3);
u32 zO1 = shoup_mul(O1,ZETA7,Z7_p), zO3 = shoup_mul(O3,ZETA5,Z5_p), iO2 = shoup_mul(O2,IMINUS,IM_p);
y[0] = add_mod(E0,O0);
y[m] = add_mod(E1,zO1);
y[2*m] = add_mod(E2,iO2);
y[3*m] = add_mod(E3,zO3);
y[4*m] = sub_mod(E0,O0);
y[5*m] = sub_mod(E1,zO1);
y[6*m] = sub_mod(E2,iO2);
y[7*m] = sub_mod(E3,zO3);
}
}
off += 7 * m;
}
}
void poly_multiply(unsigned *a, int n, unsigned *b, int m, unsigned *c) {
int L = 1;
while (L < n + m + 2) L <<= 1;
// L must be a power of 8 for radix-8 (L = 2^21 = 8^7). If not, fall back handled by choosing L=2^21.
if (L < 8) L = 8;
int i = 0;
for (; i + 8 <= n + 1; i += 8) _mm256_storeu_si256((__m256i*)(A + i), to_mont8(_mm256_loadu_si256((__m256i*)(a + i))));
for (; i <= n; i++) A[i] = to_mont(a[i]);
for (; i < L; i++) A[i] = 0;
i = 0;
for (; i + 8 <= m + 1; i += 8) _mm256_storeu_si256((__m256i*)(B + i), to_mont8(_mm256_loadu_si256((__m256i*)(b + i))));
for (; i <= m; i++) B[i] = to_mont(b[i]);
for (; i < L; i++) B[i] = 0;
u32 w = mont_pow(to_mont(3), (MOD - 1) / L);
u32 wi = mont_pow(w, MOD - 2);
gen_roots(roots, L, w);
gen_roots(roots_inv, L, wi);
build_twiddles8_fwd(roots, TW_fwd, L);
build_twiddles8_inv(roots_inv, TW_inv, L);
u32 iplus = to_mont(IROOT);
u32 iminus = mont_mul(iplus, to_mont(MOD - 1));
ntt_fwd8(A, L, TW_fwd, iplus, to_mont(ZETA), to_mont(ZETA3));
ntt_fwd8(B, L, TW_fwd, iplus, to_mont(ZETA), to_mont(ZETA3));
for (int k = 0; k < L; k += 8)
_mm256_storeu_si256((__m256i*)(A + k), mont_mul8(_mm256_loadu_si256((__m256i*)(A + k)), _mm256_loadu_si256((__m256i*)(B + k))));
ntt_inv8(A, L, TW_inv, iminus, to_mont(ZETA7), to_mont(ZETA5));
u32 linv_mont = mont_pow(to_mont((u32)(L % MOD)), MOD - 2);
u32 linv_std = mont_mul(linv_mont, 1);
__m256i linvv = _mm256_set1_epi32(linv_std);
int outn = n + m + 1;
int k = 0;
for (; k + 8 <= outn; k += 8)
_mm256_storeu_si256((__m256i*)(c + k), mont_mul8(_mm256_loadu_si256((__m256i*)(A + k)), linvv));
for (; k < outn; k++) c[k] = mont_mul(A[k], linv_std);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 75.543 ms | 55 MB + 656 KB | Accepted | Score: 100 | 显示更多 |