提交记录 33741


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 1002i. 【模板题】多项式乘法 Accepted 100 14.097 ms 8596 KB C++ 10.53 KB
提交时间 评测时间
2026-08-14 22:31:47 2026-08-14 22:31:54
// 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;
}

CompilationN/AN/ACompile OKScore: N/A

Subtask #1 Testcase #112.37 us40 KBAcceptedScore: 100

Subtask #1 Testcase #213.87 ms8 MB + 240 KBAcceptedScore: 0

Subtask #1 Testcase #36.088 ms3 MB + 284 KBAcceptedScore: 0

Subtask #1 Testcase #46.15 ms3 MB + 264 KBAcceptedScore: 0

Subtask #1 Testcase #510.03 us40 KBAcceptedScore: 0

Subtask #1 Testcase #69.09 us40 KBAcceptedScore: 0

Subtask #1 Testcase #78.85 us40 KBAcceptedScore: 0

Subtask #1 Testcase #813.341 ms7 MB + 664 KBAcceptedScore: 0

Subtask #1 Testcase #913.313 ms7 MB + 664 KBAcceptedScore: 0

Subtask #1 Testcase #1012.743 ms7 MB + 60 KBAcceptedScore: 0

Subtask #1 Testcase #1114.097 ms8 MB + 404 KBAcceptedScore: 0

Subtask #1 Testcase #1211.661 ms6 MB + 160 KBAcceptedScore: 0

Subtask #1 Testcase #138.89 us32 KBAcceptedScore: 0


Judge Duck Online | 评测鸭在线
Server Time: 2026-09-09 13:52:21 | Loaded in 1 ms | Server Status
个人娱乐项目,仅供学习交流使用 | 捐赠