提交记录 39920


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 1002. 测测你的多项式乘法 Accepted 100 91.96 ms 32400 KB C++ 7.20 KB
提交时间 评测时间
2026-08-17 01:11:25 2026-08-17 01:11:27
// NTT mod 998244353, RADIX-2 DIF forward + DIT inverse (no bitrev), all-Montgomery AVX2.
// Single per-stage-block twiddle table (contiguous access for every stage).
// vpmuludq-only mont8 kernel. Butterfly = 1 twiddle multiply.
#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;   // -MOD^-1 mod 2^32
static const u32 R = 301989884u;      // 2^32 mod MOD = mont(1)
static const u32 R2 = 932051910u;     // R^2 mod MOD
static const u32 G = 3u;
static const int MAXN = 1 << 21;
static const int MAXL = 21;

alignas(32) static u32 A[MAXN];
alignas(32) static u32 B[MAXN];
alignas(32) static u32 TB[MAXN + 64];
static int off[MAXL + 1];

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 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));
}

// per-stage blocks: block k (step=2^k) = [w^0, w^step, ..., w^{(half-1)*step}, -1]
//   half = N/(2*step).  all Montgomery form.
static void gen_twiddles(int n, int lg) {
    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);
    u32 ws[MAXL + 1];
    ws[0] = wm;
    for (int k = 1; k <= lg; k++) ws[k] = mont_s(ws[k - 1], ws[k - 1]);
    int o = 0;
    for (int k = 0; k <= lg; k++) {
        int half = (n >> 1) >> k;
        off[k] = o;
        u32 *blk = TB + o;
        blk[0] = R;
        for (int m = 1; m < half; m <<= 1) {
            int t = __builtin_ctz((unsigned)m);
            u32 f = ws[k + t];
            if (m < 8) {
                for (int j = 0; j < m; j++) blk[m + j] = mont_s(blk[j], f);
            } else {
                __m256i factor = _mm256_set1_epi32(f);
                for (int j = 0; j < m; j += 8) {
                    __m256i x = _mm256_loadu_si256((const __m256i *)(blk + j));
                    _mm256_storeu_si256((__m256i *)(blk + m + j), mont8(x, factor, ninv, modv));
                }
            }
        }
        blk[half] = MOD - R;   // w^{n/2} = -1
        o += half + 1;
    }
}

static void dif_r2(u32 *x, int n, int lg) {
    const __m256i ninv = _mm256_set1_epi32(NINV), modv = _mm256_set1_epi32(MOD);
    for (int k = 0; k < lg; k++) {
        int len = n >> k;
        int half = len >> 1;
        const u32 *blk = TB + off[k];
        for (int i = 0; i < n; i += len) {
            u32 *y = x + i;
            int j = 0;
            for (; j + 8 <= half; j += 8) {
                __m256i x0 = _mm256_loadu_si256((const __m256i *)(y + j));
                __m256i x1 = _mm256_loadu_si256((const __m256i *)(y + 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 *)(y + j), d0);
                _mm256_storeu_si256((__m256i *)(y + j + half), e1);
            }
            for (; j < half; j++) {
                u32 u = y[j], v = y[j + half];
                y[j] = add_s(u, v);
                y[j + half] = mont_s(sub_s(u, v), blk[j]);
            }
        }
    }
}

static void dit_r2(u32 *x, int n, int lg) {
    const __m256i ninv = _mm256_set1_epi32(NINV), modv = _mm256_set1_epi32(MOD);
    const __m256i rev = _mm256_setr_epi32(7, 6, 5, 4, 3, 2, 1, 0);
    for (int k = lg - 1; k >= 0; k--) {
        int len = n >> k;
        int half = len >> 1;
        const u32 *blk = TB + off[k];
        for (int i = 0; i < n; i += len) {
            u32 *y = x + i;
            int j = 0;
            for (; j + 8 <= half; j += 8) {
                __m256i x0 = _mm256_loadu_si256((const __m256i *)(y + j));
                __m256i x1 = _mm256_loadu_si256((const __m256i *)(y + j + half));
                __m256i wv = _mm256_loadu_si256((const __m256i *)(blk + half - j - 7));
                wv = _mm256_permutevar8x32_epi32(wv, rev);
                __m256i v = mont8(x1, wv, ninv, modv);
                _mm256_storeu_si256((__m256i *)(y + j), subm(x0, v, modv));
                _mm256_storeu_si256((__m256i *)(y + j + half), addm(x0, v, modv));
            }
            for (; j < half; j++) {
                u32 u = y[j];
                u32 v = mont_s(y[j + half], blk[half - j]);
                y[j] = sub_s(u, v);
                y[j + half] = add_s(u, v);
            }
        }
    }
}

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; while ((1 << lg) < size) lg++;
    if (size > MAXN) size = MAXN;

    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, lg);
    dif_r2(A, size, lg);
    dif_r2(B, size, lg);

    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));
    }
    dit_r2(A, size, lg);

    // Montgomery -> normal (multiply by 1)
    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 #191.96 ms31 MB + 656 KBAcceptedScore: 100


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