提交记录 29030


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_bot 1002. 测测你的多项式乘法 Accepted 100 96.465 ms 32404 KB C++ 8.07 KB
提交时间 评测时间
2026-06-27 02:27:17 2026-06-27 02:27:20
#pragma GCC target("avx2")
#include <immintrin.h>
#include <stdint.h>
#include <stdlib.h>

static const uint32_t MOD = 998244353u;
static const uint32_t G = 3u;
static const uint32_t QINV = 998244351u;
static const uint32_t R2 = 932051910u;
static const uint32_t R = 301989884u;

static inline uint32_t add_mod(uint32_t a, uint32_t b) {
    uint32_t s = a + b;
    return s >= MOD ? s - MOD : s;
}

static inline uint32_t sub_mod(uint32_t a, uint32_t b) {
    return a >= b ? a - b : a + MOD - b;
}

static inline uint32_t mont_reduce(uint64_t x) {
    uint32_t y = (uint32_t)x * QINV;
    uint32_t r = (uint32_t)((x + (uint64_t)y * MOD) >> 32);
    return r >= MOD ? r - MOD : r;
}

static inline uint32_t mont_mul(uint32_t a, uint32_t b) {
    return mont_reduce((uint64_t)a * b);
}

static inline uint32_t to_mont(uint32_t x) {
    return mont_mul(x, R2);
}

static inline __m256i add_vec(__m256i a, __m256i b) {
    const __m256i mod = _mm256_set1_epi32((int)MOD);
    const __m256i modm1 = _mm256_set1_epi32((int)(MOD - 1));
    __m256i s = _mm256_add_epi32(a, b);
    __m256i mask = _mm256_cmpgt_epi32(s, modm1);
    return _mm256_sub_epi32(s, _mm256_and_si256(mask, mod));
}

static inline __m256i sub_vec(__m256i a, __m256i b) {
    const __m256i mod = _mm256_set1_epi32((int)MOD);
    __m256i d = _mm256_sub_epi32(a, b);
    __m256i mask = _mm256_cmpgt_epi32(b, a);
    return _mm256_add_epi32(d, _mm256_and_si256(mask, mod));
}

static inline __m256i mont_mul_vec(__m256i a, __m256i b) {
    const __m256i mod = _mm256_set1_epi32((int)MOD);
    const __m256i modm1 = _mm256_set1_epi32((int)(MOD - 1));
    const __m256i qinv = _mm256_set1_epi32((int)QINV);

    __m256i lo = _mm256_mullo_epi32(a, b);
    __m256i m = _mm256_mullo_epi32(lo, qinv);

    __m256i t0 = _mm256_mul_epu32(a, b);
    __m256i mp0 = _mm256_mul_epu32(m, mod);
    __m256i r0 = _mm256_srli_epi64(_mm256_add_epi64(t0, mp0), 32);

    __m256i a1 = _mm256_srli_epi64(a, 32);
    __m256i b1 = _mm256_srli_epi64(b, 32);
    __m256i m1 = _mm256_srli_epi64(m, 32);
    __m256i t1 = _mm256_mul_epu32(a1, b1);
    __m256i mp1 = _mm256_mul_epu32(m1, mod);
    __m256i r1 = _mm256_srli_epi64(_mm256_add_epi64(t1, mp1), 32);

    __m256i r = _mm256_or_si256(r0, _mm256_slli_epi64(r1, 32));
    __m256i mask = _mm256_cmpgt_epi32(r, modm1);
    return _mm256_sub_epi32(r, _mm256_and_si256(mask, mod));
}

static uint32_t pow_mod(uint32_t a, uint32_t e) {
    uint32_t r = 1;
    while (e) {
        if (e & 1) r = (uint32_t)((uint64_t)r * a % MOD);
        a = (uint32_t)((uint64_t)a * a % MOD);
        e >>= 1;
    }
    return r;
}

static uint32_t roots[(1 << 21) - 1];

static void build_roots_forward(int n) {
    int off = 0;
    for (int len = n; len > 1; len >>= 1) {
        int half = len >> 1;
        uint32_t wlen = to_mont(pow_mod(G, (MOD - 1) / (uint32_t)len));
        uint32_t w = R;
        for (int i = 0; i < half; ++i) {
            roots[off + i] = w;
            w = mont_mul(w, wlen);
        }
        off += half;
    }
}

static void build_roots_inverse(int n) {
    int off = 0;
    for (int len = 2; len <= n; len <<= 1) {
        int half = len >> 1;
        uint32_t base = pow_mod(G, (MOD - 1) / (uint32_t)len);
        uint32_t wlen = to_mont(pow_mod(base, MOD - 2));
        uint32_t w = R;
        for (int i = 0; i < half; ++i) {
            roots[off + i] = w;
            w = mont_mul(w, wlen);
        }
        off += half;
    }
}

static void ntt_forward_dif(uint32_t *a, int n) {
    int off = 0;
    for (int len = n; len > 1; len >>= 1) {
        int half = len >> 1;
        uint32_t *rw = roots + off;
        for (int i = 0; i < n; i += len) {
            uint32_t *p = a + i;
            int j = 0;
            for (; j + 8 <= half; j += 8) {
                __m256i x = _mm256_loadu_si256((const __m256i *)(p + j));
                __m256i y = _mm256_loadu_si256((const __m256i *)(p + j + half));
                __m256i w = _mm256_loadu_si256((const __m256i *)(rw + j));
                _mm256_storeu_si256((__m256i *)(p + j), add_vec(x, y));
                _mm256_storeu_si256((__m256i *)(p + j + half), mont_mul_vec(sub_vec(x, y), w));
            }
            for (; j < half; ++j) {
                uint32_t x = p[j];
                uint32_t y = p[j + half];
                p[j] = add_mod(x, y);
                p[j + half] = mont_mul(sub_mod(x, y), rw[j]);
            }
        }
        off += half;
    }
}

static void ntt_inverse_dit(uint32_t *a, int n) {
    int off = 0;
    for (int len = 2; len <= n; len <<= 1) {
        int half = len >> 1;
        uint32_t *rw = roots + off;
        for (int i = 0; i < n; i += len) {
            uint32_t *p = a + i;
            int j = 0;
            for (; j + 8 <= half; j += 8) {
                __m256i x = _mm256_loadu_si256((const __m256i *)(p + j));
                __m256i y0 = _mm256_loadu_si256((const __m256i *)(p + j + half));
                __m256i w = _mm256_loadu_si256((const __m256i *)(rw + j));
                __m256i y = mont_mul_vec(y0, w);
                _mm256_storeu_si256((__m256i *)(p + j), add_vec(x, y));
                _mm256_storeu_si256((__m256i *)(p + j + half), sub_vec(x, y));
            }
            for (; j < half; ++j) {
                uint32_t x = p[j];
                uint32_t y = mont_mul(p[j + half], rw[j]);
                p[j] = add_mod(x, y);
                p[j + half] = sub_mod(x, y);
            }
        }
        off += half;
    }
}

void poly_multiply(unsigned *a, int n, unsigned *b, int m, unsigned *c) {
    int need = n + m + 1;
    int len = 1;
    while (len < need) len <<= 1;

    uint32_t *fa = (uint32_t *)calloc((size_t)len, sizeof(uint32_t));
    uint32_t *fb = (uint32_t *)calloc((size_t)len, sizeof(uint32_t));

    uint32_t digit[10];
    for (int i = 0; i < 10; ++i) digit[i] = to_mont((uint32_t)i);
    for (int i = 0; i <= n; ++i) fa[i] = digit[a[i]];
    for (int i = 0; i <= m; ++i) fb[i] = digit[b[i]];

    build_roots_forward(len);
    ntt_forward_dif(fa, len);
    ntt_forward_dif(fb, len);

    int i = 0;
    for (; i + 8 <= len; i += 8) {
        __m256i x = _mm256_loadu_si256((const __m256i *)(fa + i));
        __m256i y = _mm256_loadu_si256((const __m256i *)(fb + i));
        _mm256_storeu_si256((__m256i *)(fa + i), mont_mul_vec(x, y));
    }
    for (; i < len; ++i) fa[i] = mont_mul(fa[i], fb[i]);

    build_roots_inverse(len);
    ntt_inverse_dit(fa, len);

    uint32_t inv_n = pow_mod((uint32_t)len, MOD - 2);
    __m256i vinv = _mm256_set1_epi32((int)inv_n);
    i = 0;
    for (; i + 8 <= need; i += 8) {
        __m256i x = _mm256_loadu_si256((const __m256i *)(fa + i));
        _mm256_storeu_si256((__m256i *)(c + i), mont_mul_vec(x, vinv));
    }
    for (; i < need; ++i) c[i] = mont_mul(fa[i], inv_n);

    free(fa);
    free(fb);
}

#ifdef LOCAL_TEST
#include <stdio.h>
static unsigned aa[1 << 20], bb[1 << 20], cc[1 << 21], dd[4096];
int main() {
    for (int n = 0; n < 64; ++n) {
        for (int m = 0; m < 64; ++m) {
            for (int i = 0; i <= n; ++i) aa[i] = (unsigned)((i * 7 + n) % 10);
            for (int i = 0; i <= m; ++i) bb[i] = (unsigned)((i * 5 + m) % 10);
            poly_multiply(aa, n, bb, m, cc);
            for (int i = 0; i <= n + m; ++i) dd[i] = 0;
            for (int i = 0; i <= n; ++i)
                for (int j = 0; j <= m; ++j)
                    dd[i + j] += aa[i] * bb[j];
            for (int i = 0; i <= n + m; ++i) {
                if (cc[i] != dd[i]) {
                    printf("bad n=%d m=%d i=%d got=%u want=%u\n", n, m, i, cc[i], dd[i]);
                    return 1;
                }
            }
        }
    }
    puts("ok");
    return 0;
}
#endif

#ifdef LOCAL_BENCH
#include <stdio.h>
#include <time.h>
static unsigned aa[1000001], bb[1000001], cc[2000001];
int main() {
    for (int i = 0; i <= 1000000; ++i) {
        aa[i] = (unsigned)((i * 7 + 3) % 10);
        bb[i] = (unsigned)((i * 5 + 1) % 10);
    }
    clock_t st = clock();
    poly_multiply(aa, 1000000, bb, 1000000, cc);
    clock_t ed = clock();
    unsigned long long sample = 0;
    for (int i = 0; i <= 2000000; i += 137) sample += cc[i];
    printf("%.3f ms sample=%llu\n", 1000.0 * (double)(ed - st) / CLOCKS_PER_SEC, sample);
    return 0;
}
#endif

CompilationN/AN/ACompile OKScore: N/A

Testcase #196.465 ms31 MB + 660 KBAcceptedScore: 100


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