// v10: AVX2 Montgomery NTT, dense u32 (8 lanes), radix-2 DIF/DIT, contiguous twiddles.
#pragma GCC target("avx2")
#include <cstdio>
#include <cstring>
#include <cstdlib>
#include <immintrin.h>
typedef unsigned long long u64;
typedef unsigned int u32;
const u32 MOD = 998244353u;
const u32 NINV = 998244351u;
const u32 R2 = 932051910u;
const u32 G = 3u;
const int MAXN = 1 << 18;
alignas(32) static u32 a[MAXN], b[MAXN];
alignas(32) static u32 roots[MAXN];
alignas(32) static u32 tw_fwd[MAXN];
alignas(32) static u32 tw_inv[MAXN];
static inline u32 modpow(u32 base, u64 e) {
u64 r = 1, bb = base % MOD;
for (; e; e >>= 1) {
if (e & 1) r = r * bb % MOD;
bb = bb * bb % MOD;
}
return (u32)r;
}
static inline u32 mont_s(u32 a, u32 b) {
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;
}
// 8-lane Montgomery multiply: x*y*2^-32 mod MOD
static inline __m256i mont8(__m256i x, __m256i y, const __m256i &ninv, const __m256i &modv, const __m256i &modm1, const __m256i &idx0246) {
__m256i t_lo = _mm256_mullo_epi32(x, y);
__m256i m = _mm256_mullo_epi32(t_lo, ninv);
__m256i pe = _mm256_mul_epu32(x, y);
__m256i po = _mm256_mul_epu32(_mm256_srli_epi64(x, 32), _mm256_srli_epi64(y, 32));
__m256i me = _mm256_mul_epu32(m, modv);
__m256i mo = _mm256_mul_epu32(_mm256_srli_epi64(m, 32), modv);
__m256i ue = _mm256_srli_epi64(_mm256_add_epi64(pe, me), 32);
__m256i uo = _mm256_srli_epi64(_mm256_add_epi64(po, mo), 32);
__m256i e = _mm256_permutevar8x32_epi32(ue, idx0246);
__m256i o = _mm256_permutevar8x32_epi32(uo, idx0246);
__m256i lo = _mm256_unpacklo_epi32(e, o);
__m256i hi = _mm256_unpackhi_epi32(e, o);
__m256i u = _mm256_blend_epi32(lo, hi, 0xF0);
__m256i mask = _mm256_cmpgt_epi32(u, modm1);
u = _mm256_sub_epi32(u, _mm256_and_si256(mask, modv));
return u;
}
static void ntt_fwd(u32 *x, int n, const u32 *tw) {
const __m256i ninv = _mm256_set1_epi32(NINV);
const __m256i modv = _mm256_set1_epi32(MOD);
const __m256i modm1 = _mm256_set1_epi32(MOD - 1);
const __m256i idx0246 = _mm256_setr_epi32(0,2,4,6,0,2,4,6);
for (int len = n; len > 1; len >>= 1) {
int half = len >> 1;
for (int i = 0; i < n; i += len) {
u32 *y = x + i;
int j = 0;
for (; j + 16 <= half; j += 16) {
__m256i u0 = _mm256_loadu_si256((__m256i*)(y + j));
__m256i v0 = _mm256_loadu_si256((__m256i*)(y + j + half));
__m256i u1 = _mm256_loadu_si256((__m256i*)(y + j + 8));
__m256i v1 = _mm256_loadu_si256((__m256i*)(y + j + 8 + half));
__m256i s0 = _mm256_add_epi32(u0, v0); s0 = _mm256_min_epu32(s0, _mm256_sub_epi32(s0, modv));
__m256i d0 = _mm256_sub_epi32(u0, v0); d0 = _mm256_min_epu32(d0, _mm256_add_epi32(d0, modv));
__m256i s1 = _mm256_add_epi32(u1, v1); s1 = _mm256_min_epu32(s1, _mm256_sub_epi32(s1, modv));
__m256i d1 = _mm256_sub_epi32(u1, v1); d1 = _mm256_min_epu32(d1, _mm256_add_epi32(d1, modv));
__m256i tw0 = _mm256_loadu_si256((__m256i*)(tw + j));
__m256i tw1 = _mm256_loadu_si256((__m256i*)(tw + j + 8));
__m256i dv0 = mont8(d0, tw0, ninv, modv, modm1, idx0246);
__m256i dv1 = mont8(d1, tw1, ninv, modv, modm1, idx0246);
_mm256_storeu_si256((__m256i*)(y + j), s0);
_mm256_storeu_si256((__m256i*)(y + j + half), dv0);
_mm256_storeu_si256((__m256i*)(y + j + 8), s1);
_mm256_storeu_si256((__m256i*)(y + j + 8 + half), dv1);
}
for (; j + 8 <= half; j += 8) {
__m256i u = _mm256_loadu_si256((__m256i*)(y + j));
__m256i v = _mm256_loadu_si256((__m256i*)(y + j + half));
__m256i s = _mm256_add_epi32(u, v);
s = _mm256_min_epu32(s, _mm256_sub_epi32(s, modv));
__m256i d = _mm256_sub_epi32(u, v);
d = _mm256_min_epu32(d, _mm256_add_epi32(d, modv));
__m256i twv = _mm256_loadu_si256((__m256i*)(tw + j));
__m256i dv = mont8(d, twv, ninv, modv, modm1, idx0246);
_mm256_storeu_si256((__m256i*)(y + j), s);
_mm256_storeu_si256((__m256i*)(y + j + half), dv);
}
for (; 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] = (half == 1) ? d : mont_s(d, tw[j]);
}
}
tw += half;
}
}
static void ntt_inv(u32 *x, int n, const u32 *tw) {
const __m256i ninv = _mm256_set1_epi32(NINV);
const __m256i modv = _mm256_set1_epi32(MOD);
const __m256i modm1 = _mm256_set1_epi32(MOD - 1);
const __m256i idx0246 = _mm256_setr_epi32(0,2,4,6,0,2,4,6);
for (int len = 2; len <= n; len <<= 1) {
int half = len >> 1;
for (int i = 0; i < n; i += len) {
u32 *y = x + i;
int j = 0;
for (; j + 16 <= half; j += 16) {
__m256i u0 = _mm256_loadu_si256((__m256i*)(y + j));
__m256i v0 = _mm256_loadu_si256((__m256i*)(y + j + half));
__m256i u1 = _mm256_loadu_si256((__m256i*)(y + j + 8));
__m256i v1 = _mm256_loadu_si256((__m256i*)(y + j + 8 + half));
__m256i tw0 = _mm256_loadu_si256((__m256i*)(tw + j));
__m256i tw1 = _mm256_loadu_si256((__m256i*)(tw + j + 8));
v0 = mont8(v0, tw0, ninv, modv, modm1, idx0246);
v1 = mont8(v1, tw1, ninv, modv, modm1, idx0246);
__m256i s0 = _mm256_add_epi32(u0, v0); s0 = _mm256_min_epu32(s0, _mm256_sub_epi32(s0, modv));
__m256i d0 = _mm256_sub_epi32(u0, v0); d0 = _mm256_min_epu32(d0, _mm256_add_epi32(d0, modv));
__m256i s1 = _mm256_add_epi32(u1, v1); s1 = _mm256_min_epu32(s1, _mm256_sub_epi32(s1, modv));
__m256i d1 = _mm256_sub_epi32(u1, v1); d1 = _mm256_min_epu32(d1, _mm256_add_epi32(d1, modv));
_mm256_storeu_si256((__m256i*)(y + j), s0);
_mm256_storeu_si256((__m256i*)(y + j + half), d0);
_mm256_storeu_si256((__m256i*)(y + j + 8), s1);
_mm256_storeu_si256((__m256i*)(y + j + 8 + half), d1);
}
for (; j + 8 <= half; j += 8) {
__m256i u = _mm256_loadu_si256((__m256i*)(y + j));
__m256i v = _mm256_loadu_si256((__m256i*)(y + j + half));
__m256i twv = _mm256_loadu_si256((__m256i*)(tw + j));
v = mont8(v, twv, ninv, modv, modm1, idx0246);
__m256i s = _mm256_add_epi32(u, v);
s = _mm256_min_epu32(s, _mm256_sub_epi32(s, modv));
__m256i d = _mm256_sub_epi32(u, v);
d = _mm256_min_epu32(d, _mm256_add_epi32(d, modv));
_mm256_storeu_si256((__m256i*)(y + j), s);
_mm256_storeu_si256((__m256i*)(y + j + half), d);
}
for (; j < half; j++) {
u32 u = y[j];
u32 v = (half == 1) ? y[j + half] : mont_s(y[j + half], tw[j]);
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;
}
}
tw += half;
}
}
// fast input buffer
static const int BUFSZ = 1 << 20;
static char inbuf[BUFSZ];
static size_t inpos = 0, inlen = 0;
static inline int readbyte() {
if (inpos >= inlen) {
inlen = fread(inbuf, 1, BUFSZ, stdin);
inpos = 0;
if (inlen == 0) return -1;
}
return (unsigned char)inbuf[inpos++];
}
static inline int readint() {
int c = readbyte();
while (c == ' ' || c == '\n' || c == '\r' || c == '\t') c = readbyte();
int x = 0;
while (c >= '0' && c <= '9') { x = x * 10 + (c - '0'); c = readbyte(); }
return x;
}
static char outbuf[1 << 21];
static size_t outpos = 0;
static inline void putc(char c) { outbuf[outpos++] = c; }
static inline void putint(int x) {
if (x == 0) { putc('0'); return; }
char tmp[12]; int t = 0;
while (x) { tmp[t++] = '0' + (x % 10); x /= 10; }
while (t) putc(tmp[--t]);
}
int main() {
int n = readint();
int m = readint();
int na = n + 1;
int nb = m + 1;
for (int i = 0; i < na; i++) a[i] = (u32)readint();
for (int i = 0; i < nb; i++) b[i] = (u32)readint();
int size = 1;
while (size < na + nb - 1) size <<= 1;
u32 wstep = (MOD - 1) / (u32)size;
u32 w = modpow(G, wstep);
u32 invw = modpow(w, MOD - 2);
u64 cur = 1;
for (int i = 0; i < size; i++) { roots[i] = mont_s((u32)cur, R2); cur = cur * w % MOD; }
// build contiguous forward twiddles
{
u32 *p = tw_fwd;
for (int len = size; len > 1; len >>= 1) {
int half = len >> 1;
int step = size / len;
for (int j = 0; j < half; j++) p[j] = roots[j * step];
p += half;
}
}
// inverse roots
cur = 1;
for (int i = 0; i < size; i++) { roots[i] = mont_s((u32)cur, R2); cur = cur * invw % MOD; }
{
u32 *p = tw_inv;
for (int len = 2; len <= size; len <<= 1) {
int half = len >> 1;
int step = size / len;
for (int j = 0; j < half; j++) p[j] = roots[j * step];
p += half;
}
}
for (int i = 0; i < size; i++) a[i] = mont_s(a[i], R2);
for (int i = 0; i < size; i++) b[i] = mont_s(b[i], R2);
ntt_fwd(a, size, tw_fwd);
ntt_fwd(b, size, tw_fwd);
{
const __m256i ninv = _mm256_set1_epi32(NINV);
const __m256i modv = _mm256_set1_epi32(MOD);
const __m256i modm1 = _mm256_set1_epi32(MOD - 1);
const __m256i idx0246 = _mm256_setr_epi32(0,2,4,6,0,2,4,6);
for (int i = 0; i < size; i += 8) {
__m256i x = _mm256_loadu_si256((__m256i*)(a + i));
__m256i y = _mm256_loadu_si256((__m256i*)(b + i));
_mm256_storeu_si256((__m256i*)(a + i), mont8(x, y, ninv, modv, modm1, idx0246));
}
}
ntt_inv(a, size, tw_inv);
for (int i = 0; i < size; i++) a[i] = mont_s(a[i], 1);
u32 ninv_scale = modpow((u32)size, MOD - 2);
for (int i = 0; i < size; i++) a[i] = (u32)((u64)a[i] * ninv_scale % MOD);
int outn = n + m + 1;
for (int i = 0; i < outn; i++) {
if (i) putc(' ');
putint((int)a[i]);
}
putc('\n');
fwrite(outbuf, 1, outpos, stdout);
return 0;
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Subtask #1 Testcase #1 | 12.37 us | 40 KB | Accepted | Score: 100 | 显示更多 |
| Subtask #1 Testcase #2 | 13.87 ms | 8 MB + 240 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #3 | 6.088 ms | 3 MB + 284 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #4 | 6.15 ms | 3 MB + 264 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #5 | 10.03 us | 40 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #6 | 9.09 us | 40 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #7 | 8.85 us | 40 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #8 | 13.341 ms | 7 MB + 664 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #9 | 13.313 ms | 7 MB + 664 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #10 | 12.743 ms | 7 MB + 60 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #11 | 14.097 ms | 8 MB + 404 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #12 | 11.661 ms | 6 MB + 160 KB | Accepted | Score: 0 | 显示更多 |
| Subtask #1 Testcase #13 | 8.89 us | 32 KB | Accepted | Score: 0 | 显示更多 |