提交记录 39923


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 1002. 测测你的多项式乘法 Accepted 100 73.156 ms 41620 KB C++ 13.57 KB
提交时间 评测时间
2026-08-17 02:08:53 2026-08-17 02:08:55
// NTT mod 998244353, RADIX-4 DIF/DIT (1 radix-2 top + 10 radix-4), AVX2.
// All-Montgomery data, Montgomery twiddles (mont8), Shoup internal i (shoup_mul8).
// Butterfly = 4 mults / 4 points, ~11 live YMM regs (no spill -> pipelines).
#pragma GCC target("avx2")
#pragma GCC optimize("O3,unroll-loops")
#include <immintrin.h>
#include <stdint.h>

typedef unsigned long long u64;
typedef unsigned int u32;

static const u32 MOD = 998244353u;
static const u32 NINV = 998244351u;
static const u32 R = 301989884u;
static const u32 R2 = 932051910u;
static const u32 G = 3u;
static const u32 IM = 911660635u;            // sqrt(-1) mod p
static const int MAXN = 1 << 21;

alignas(32) static u32 A[MAXN];
alignas(32) static u32 B[MAXN];
alignas(32) static u32 tw[MAXN + 64];
alignas(32) static u32 twi[MAXN + 64];
alignas(32) static u32 tmp[MAXN / 4 + 8];

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;
}
static inline u32 mont_pow(u32 base, u64 e) {
    u32 r = R;
    for (; e; e >>= 1) { if (e & 1) r = mont_s(r, base); base = mont_s(base, base); }
    return r;
}
static inline u32 add_s(u32 a, u32 b) { u32 s = a + b; return s >= MOD ? s - MOD : s; }
static inline u32 sub_s(u32 a, u32 b) { u32 d = a - b; return (int)d < 0 ? d + MOD : d; }

static inline __m256i mont8(__m256i x, __m256i y, const __m256i ninv, const __m256i modv) {
    __m256i ao = _mm256_srli_epi64(x, 32);
    __m256i bo = _mm256_srli_epi64(y, 32);
    __m256i te = _mm256_mul_epu32(x, y);
    __m256i to = _mm256_mul_epu32(ao, bo);
    __m256i me = _mm256_mul_epu32(te, ninv);
    __m256i mo = _mm256_mul_epu32(to, ninv);
    __m256i ue = _mm256_srli_epi64(_mm256_add_epi64(te, _mm256_mul_epu32(me, modv)), 32);
    __m256i uo = _mm256_srli_epi64(_mm256_add_epi64(to, _mm256_mul_epu32(mo, modv)), 32);
    __m256i u = _mm256_or_si256(ue, _mm256_slli_epi64(uo, 32));
    u = _mm256_min_epu32(u, _mm256_sub_epi32(u, modv));
    return u;
}
static inline __m256i addm(__m256i a, __m256i b, const __m256i modv) {
    __m256i s = _mm256_add_epi32(a, b);
    return _mm256_min_epu32(s, _mm256_sub_epi32(s, modv));
}
static inline __m256i subm(__m256i a, __m256i b, const __m256i modv) {
    __m256i d = _mm256_sub_epi32(a, b);
    return _mm256_min_epu32(d, _mm256_add_epi32(d, modv));
}
const u64 MAGIC = 18479187002ULL;
static inline u32 shp_quot(u32 w) { return (u32)(((u64)w * MAGIC) >> 32); }
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 __m256i shoup_mul8(__m256i a, __m256i w, __m256i wp, const __m256i modv) {
    __m256i lo = _mm256_mullo_epi32(a, w);
    __m256i qe = _mm256_mul_epu32(a, wp);
    __m256i qo = _mm256_mul_epu32(_mm256_srli_epi64(a, 32), _mm256_srli_epi64(wp, 32));
    __m256i q = _mm256_or_si256(_mm256_srli_epi64(qe, 32), _mm256_slli_epi64(_mm256_srli_epi64(qo, 32), 32));
    __m256i r = _mm256_sub_epi32(lo, _mm256_mullo_epi32(q, modv));
    return _mm256_min_epu32(r, _mm256_sub_epi32(r, modv));
}

// generate tmp[0..s-1] = mont(wm^p) via serial 8-wide carry chain
static void gen_powers(u32 *t, int s, u32 wm, const __m256i ninv, const __m256i modv) {
    t[0] = R;
    u32 P[8];
    P[0] = wm;
    for (int i = 1; i < 8; i++) P[i] = mont_s(P[i - 1], wm);
    __m256i Pv = _mm256_loadu_si256((__m256i *)P);
    __m256i carry = _mm256_set1_epi32((int)R);
    int p = 1;
    for (; p + 8 <= s; p += 8) {
        __m256i chunk = mont8(carry, Pv, ninv, modv);
        _mm256_storeu_si256((__m256i *)(t + p), chunk);
        carry = _mm256_set1_epi32((int)t[p + 7]);
    }
    for (; p < s; p++) t[p] = mont_s(t[p - 1], wm);
}

static void gen_twiddles(int n, int top) {
    const __m256i ninv = _mm256_set1_epi32(NINV), modv = _mm256_set1_epi32(MOD);
    u32 w = modpow(G, (MOD - 1) / n);
    u32 wm = mont_s(w, R2);
    // forward: radix-2 top twiddles w^j for j=0..n/2-1 (only if top)
    if (top) gen_powers(tw, n / 2, wm, ninv, modv);
    int off = top ? n / 2 : 0;
    for (int len = top ? (n >> 1) : n; len >= 4; len >>= 2) {
        int m = len >> 2;
        int step = n / len;
        u32 b = mont_pow(wm, step);
        gen_powers(tmp, m, b, ninv, modv);
        for (int p = 0; p < m; p += 8) {
            __m256i t1 = _mm256_loadu_si256((const __m256i *)(tmp + p));
            __m256i t2 = mont8(t1, t1, ninv, modv);
            __m256i t3 = mont8(t2, t1, ninv, modv);
            _mm256_storeu_si256((__m256i *)(tw + off + 3 * p + 0), t1);
            _mm256_storeu_si256((__m256i *)(tw + off + 3 * p + 8), t2);
            _mm256_storeu_si256((__m256i *)(tw + off + 3 * p + 16), t3);
        }
        off += 3 * ((m + 7) & ~7);
    }
    // inverse twiddles
    u32 iw = modpow(w, MOD - 2);
    u32 iwm = mont_s(iw, R2);
    if (top) gen_powers(twi, n / 2, iwm, ninv, modv);
    off = top ? n / 2 : 0;
    for (int len = 4; len <= (top ? (n >> 1) : n); len <<= 2) {
        int m = len >> 2;
        int step = n / len;
        u32 b = mont_pow(iwm, step);
        gen_powers(tmp, m, b, ninv, modv);
        for (int p = 0; p < m; p += 8) {
            __m256i t1 = _mm256_loadu_si256((const __m256i *)(tmp + p));
            __m256i t2 = mont8(t1, t1, ninv, modv);
            __m256i t3 = mont8(t2, t1, ninv, modv);
            _mm256_storeu_si256((__m256i *)(twi + off + 3 * p + 0), t1);
            _mm256_storeu_si256((__m256i *)(twi + off + 3 * p + 8), t2);
            _mm256_storeu_si256((__m256i *)(twi + off + 3 * p + 16), t3);
        }
        off += 3 * ((m + 7) & ~7);
    }
}

static void dif4(u32 *x, int n, int top) {
    const __m256i ninv = _mm256_set1_epi32(NINV), modv = _mm256_set1_epi32(MOD);
    u32 iv = IM;                 // sqrt(-1)
    u32 ivp = shp_quot(iv);
    __m256i ivv = _mm256_set1_epi32((int)iv);
    __m256i ivpv = _mm256_set1_epi32((int)ivp);
    // radix-2 top
    if (top) {
        int half = n >> 1;
        const u32 *blk = tw;
        int j = 0;
        for (; j + 8 <= half; j += 8) {
            __m256i x0 = _mm256_loadu_si256((const __m256i *)(x + j));
            __m256i x1 = _mm256_loadu_si256((const __m256i *)(x + j + half));
            __m256i wv = _mm256_loadu_si256((const __m256i *)(blk + j));
            __m256i d0 = addm(x0, x1, modv);
            __m256i e1 = mont8(subm(x0, x1, modv), wv, ninv, modv);
            _mm256_storeu_si256((__m256i *)(x + j), d0);
            _mm256_storeu_si256((__m256i *)(x + j + half), e1);
        }
        for (; j < half; j++) {
            u32 u = x[j], v = x[j + half];
            x[j] = add_s(u, v);
            x[j + half] = mont_s(sub_s(u, v), blk[j]);
        }
    }
    // radix-4 stages
    int off = top ? n / 2 : 0;
    for (int len = top ? (n >> 1) : n; len >= 4; len >>= 2) {
        int m = len >> 2;
        for (int i = 0; i < n; i += len) {
            u32 *y = x + i;
            const u32 *blk = tw + off;
            int p = 0;
            for (; p + 8 <= m; p += 8) {
                __m256i x0 = _mm256_loadu_si256((const __m256i *)(y + p));
                __m256i x1 = _mm256_loadu_si256((const __m256i *)(y + p + m));
                __m256i x2 = _mm256_loadu_si256((const __m256i *)(y + p + 2 * m));
                __m256i x3 = _mm256_loadu_si256((const __m256i *)(y + p + 3 * m));
                __m256i w1 = _mm256_loadu_si256((const __m256i *)(blk + 3 * p + 0));
                __m256i w2 = _mm256_loadu_si256((const __m256i *)(blk + 3 * p + 8));
                __m256i w3 = _mm256_loadu_si256((const __m256i *)(blk + 3 * p + 16));
                __m256i a0 = addm(x0, x2, modv), a1 = subm(x0, x2, modv);
                __m256i a2 = addm(x1, x3, modv), a3 = subm(x1, x3, modv);
                __m256i ia3 = shoup_mul8(a3, ivv, ivpv, modv);
                _mm256_storeu_si256((__m256i *)(y + p), addm(a0, a2, modv));
                _mm256_storeu_si256((__m256i *)(y + p + m), mont8(addm(a1, ia3, modv), w1, ninv, modv));
                _mm256_storeu_si256((__m256i *)(y + p + 2 * m), mont8(subm(a0, a2, modv), w2, ninv, modv));
                _mm256_storeu_si256((__m256i *)(y + p + 3 * m), mont8(subm(a1, ia3, modv), w3, ninv, modv));
            }
            for (; p < m; p++) {
                u32 x0 = y[p], x1 = y[p + m], x2 = y[p + 2 * m], x3 = y[p + 3 * m];
                u32 w1 = blk[p], w2 = blk[8 + p], w3 = blk[16 + p];
                u32 a0 = add_s(x0, x2), a1 = sub_s(x0, x2), a2 = add_s(x1, x3), a3 = sub_s(x1, x3);
                u32 ia3 = shoup_mul(a3, iv, ivp);
                y[p] = add_s(a0, a2);
                y[p + m] = mont_s(add_s(a1, ia3), w1);
                y[p + 2 * m] = mont_s(sub_s(a0, a2), w2);
                y[p + 3 * m] = mont_s(sub_s(a1, ia3), w3);
            }
        }
        off += 3 * ((m + 7) & ~7);
    }
}

static void dit4(u32 *x, int n, int top) {
    const __m256i ninv = _mm256_set1_epi32(NINV), modv = _mm256_set1_epi32(MOD);
    u32 iiv = MOD - IM;           // -i = i^{-1}
    u32 iivp = shp_quot(iiv);
    __m256i iivv = _mm256_set1_epi32((int)iiv);
    __m256i iivpv = _mm256_set1_epi32((int)iivp);
    // radix-4 stages (increasing len); inverse blocks stored for len=4,16,... in order
    int off = top ? n / 2 : 0;
    for (int len = 4; len <= (top ? (n >> 1) : n); len <<= 2) {
        int m = len >> 2;
        for (int i = 0; i < n; i += len) {
            u32 *y = x + i;
            const u32 *blk = twi + off;
            int p = 0;
            for (; p + 8 <= m; p += 8) {
                __m256i y0 = _mm256_loadu_si256((const __m256i *)(y + p));
                __m256i y1 = _mm256_loadu_si256((const __m256i *)(y + p + m));
                __m256i y2 = _mm256_loadu_si256((const __m256i *)(y + p + 2 * m));
                __m256i y3 = _mm256_loadu_si256((const __m256i *)(y + p + 3 * m));
                __m256i w1 = _mm256_loadu_si256((const __m256i *)(blk + 3 * p + 0));
                __m256i w2 = _mm256_loadu_si256((const __m256i *)(blk + 3 * p + 8));
                __m256i w3 = _mm256_loadu_si256((const __m256i *)(blk + 3 * p + 16));
                __m256i t1 = mont8(y1, w1, ninv, modv);
                __m256i t2 = mont8(y2, w2, ninv, modv);
                __m256i t3 = mont8(y3, w3, ninv, modv);
                __m256i b0 = addm(y0, t2, modv), b1 = subm(y0, t2, modv);
                __m256i b2 = addm(t1, t3, modv), b3 = subm(t1, t3, modv);
                __m256i ib3 = shoup_mul8(b3, iivv, iivpv, modv);
                _mm256_storeu_si256((__m256i *)(y + p), addm(b0, b2, modv));
                _mm256_storeu_si256((__m256i *)(y + p + m), addm(b1, ib3, modv));
                _mm256_storeu_si256((__m256i *)(y + p + 2 * m), subm(b0, b2, modv));
                _mm256_storeu_si256((__m256i *)(y + p + 3 * m), subm(b1, ib3, modv));
            }
            for (; p < m; p++) {
                u32 y0 = y[p], y1 = y[p + m], y2 = y[p + 2 * m], y3 = y[p + 3 * m];
                u32 w1 = blk[p], w2 = blk[8 + p], w3 = blk[16 + p];
                u32 t1 = mont_s(y1, w1), t2 = mont_s(y2, w2), t3 = mont_s(y3, w3);
                u32 b0 = add_s(y0, t2), b1 = sub_s(y0, t2), b2 = add_s(t1, t3), b3 = sub_s(t1, t3);
                u32 ib3 = shoup_mul(b3, iiv, iivp);
                y[p] = add_s(b0, b2);
                y[p + m] = add_s(b1, ib3);
                y[p + 2 * m] = sub_s(b0, b2);
                y[p + 3 * m] = sub_s(b1, ib3);
            }
        }
        off += 3 * ((m + 7) & ~7);
    }
    // radix-2 bottom (len=2)
    if (top) {
        int half = n >> 1;
        const u32 *blk = twi;  // inverse w^{-j}
        int j = 0;
        for (; j + 8 <= half; j += 8) {
            __m256i y0 = _mm256_loadu_si256((const __m256i *)(x + j));
            __m256i y1 = _mm256_loadu_si256((const __m256i *)(x + j + half));
            __m256i wv = _mm256_loadu_si256((const __m256i *)(blk + j));
            __m256i t = mont8(y1, wv, ninv, modv);
            _mm256_storeu_si256((__m256i *)(x + j), addm(y0, t, modv));
            _mm256_storeu_si256((__m256i *)(x + j + half), subm(y0, t, modv));
        }
        for (; j < half; j++) {
            u32 u = x[j], t = mont_s(x[j + half], blk[j]);
            x[j] = add_s(u, t);
            x[j + half] = sub_s(u, t);
        }
    }
}

void poly_multiply(unsigned *a, int n, unsigned *b, int m, unsigned *c) {
    int size = 1; while (size < n + m + 1) size <<= 1;
    int lg = 0; { int t = size; while (t > 1) { t >>= 1; lg++; } }
    int top = lg & 1;
    u32 ninv_s = modpow((u32)size, MOD - 2);
    u32 ninv_mont = mont_s(ninv_s, R2);
    u32 a_val[10], b_val[10];
    for (int v = 0; v < 10; v++) {
        a_val[v] = mont_s(mont_s((u32)v, R2), ninv_mont);
        b_val[v] = mont_s((u32)v, R2);
    }
    int i = 0;
    for (; i <= n; i++) A[i] = a_val[a[i]];
    for (; i < size; i++) A[i] = 0;
    for (i = 0; i <= m; i++) B[i] = b_val[b[i]];
    for (; i < size; i++) B[i] = 0;

    gen_twiddles(size, top);
    dif4(A, size, top);
    dif4(B, size, top);

    const __m256i ninv = _mm256_set1_epi32(NINV), modv = _mm256_set1_epi32(MOD);
    for (i = 0; i < size; i += 8) {
        __m256i x = _mm256_loadu_si256((const __m256i *)(A + i));
        __m256i y = _mm256_loadu_si256((const __m256i *)(B + i));
        _mm256_storeu_si256((__m256i *)(A + i), mont8(x, y, ninv, modv));
    }
    dit4(A, size, top);

    const __m256i ones = _mm256_set1_epi32(1);
    int outn = n + m + 1;
    for (i = 0; i + 8 <= outn; i += 8) {
        __m256i x = _mm256_loadu_si256((const __m256i *)(A + i));
        _mm256_storeu_si256((__m256i *)(c + i), mont8(x, ones, ninv, modv));
    }
    for (; i < outn; i++) c[i] = mont_s(A[i], 1);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #173.156 ms40 MB + 660 KBAcceptedScore: 100


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