#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];
static void build_roots_forward(int n) {
int off = 0;
int lg = __builtin_ctz((unsigned)n);
int start = n;
if (lg & 1) {
int half = n >> 1;
uint32_t wlen = to_mont(pow_mod(G, (MOD - 1) / (uint32_t)n));
uint32_t w = R;
for (int i = 0; i < half; ++i) {
roots[off + i] = w;
w = mont_mul(w, wlen);
}
off += half;
start = n >> 1;
}
for (int len = start; len >= 4; len >>= 2) {
int q = len >> 2;
uint32_t step = to_mont(pow_mod(G, (MOD - 1) / (uint32_t)len));
uint32_t w1 = R;
for (int i = 0; i < q; ++i) {
uint32_t w2 = mont_mul(w1, w1);
roots[off + i] = w1;
roots[off + q + i] = w2;
roots[off + q + q + i] = mont_mul(w2, w1);
w1 = mont_mul(w1, step);
}
off += q * 3;
}
}
static void build_roots_inverse(int n) {
int off = 0;
int lg = __builtin_ctz((unsigned)n);
int start = (lg & 1) ? (n >> 1) : n;
for (int len = 4; len <= start; len <<= 2) {
int q = len >> 2;
uint32_t base = pow_mod(G, (MOD - 1) / (uint32_t)len);
uint32_t step = to_mont(pow_mod(base, MOD - 2));
uint32_t w1 = R;
for (int i = 0; i < q; ++i) {
uint32_t w2 = mont_mul(w1, w1);
roots[off + i] = w1;
roots[off + q + i] = w2;
roots[off + q + q + i] = mont_mul(w2, w1);
w1 = mont_mul(w1, step);
}
off += q * 3;
}
if (lg & 1) {
int half = n >> 1;
uint32_t base = pow_mod(G, (MOD - 1) / (uint32_t)n);
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);
}
}
}
static void ntt_forward(uint32_t *a, int n) {
int off = 0;
int lg = __builtin_ctz((unsigned)n);
int start = n;
if (lg & 1) {
int half = n >> 1;
uint32_t *rw = roots;
uint32_t *p = a;
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;
start = n >> 1;
}
const uint32_t imag = to_mont(pow_mod(G, (MOD - 1) / 4));
const __m256i vi = _mm256_set1_epi32((int)imag);
for (int len = start; len >= 4; len >>= 2) {
int q = len >> 2;
uint32_t *r1 = roots + off;
uint32_t *r2 = r1 + q;
uint32_t *r3 = r2 + q;
for (int i = 0; i < n; i += len) {
uint32_t *p = a + i;
int j = 0;
for (; j + 8 <= q; j += 8) {
__m256i a0 = _mm256_loadu_si256((const __m256i *)(p + j));
__m256i a1 = _mm256_loadu_si256((const __m256i *)(p + j + q));
__m256i a2 = _mm256_loadu_si256((const __m256i *)(p + j + q + q));
__m256i a3 = _mm256_loadu_si256((const __m256i *)(p + j + q + q + q));
__m256i s02 = add_vec(a0, a2);
__m256i d02 = sub_vec(a0, a2);
__m256i s13 = add_vec(a1, a3);
__m256i d13 = mont_mul_vec(sub_vec(a1, a3), vi);
__m256i w1 = _mm256_loadu_si256((const __m256i *)(r1 + j));
__m256i w2 = _mm256_loadu_si256((const __m256i *)(r2 + j));
__m256i w3 = _mm256_loadu_si256((const __m256i *)(r3 + j));
_mm256_storeu_si256((__m256i *)(p + j), add_vec(s02, s13));
_mm256_storeu_si256((__m256i *)(p + j + q), mont_mul_vec(add_vec(d02, d13), w1));
_mm256_storeu_si256((__m256i *)(p + j + q + q), mont_mul_vec(sub_vec(s02, s13), w2));
_mm256_storeu_si256((__m256i *)(p + j + q + q + q), mont_mul_vec(sub_vec(d02, d13), w3));
}
for (; j < q; ++j) {
uint32_t a0 = p[j], a1 = p[j + q], a2 = p[j + q + q], a3 = p[j + q + q + q];
uint32_t s02 = add_mod(a0, a2);
uint32_t d02 = sub_mod(a0, a2);
uint32_t s13 = add_mod(a1, a3);
uint32_t d13 = mont_mul(sub_mod(a1, a3), imag);
p[j] = add_mod(s02, s13);
p[j + q] = mont_mul(add_mod(d02, d13), r1[j]);
p[j + q + q] = mont_mul(sub_mod(s02, s13), r2[j]);
p[j + q + q + q] = mont_mul(sub_mod(d02, d13), r3[j]);
}
}
off += q * 3;
}
}
static void ntt_inverse(uint32_t *a, int n) {
int off = 0;
const uint32_t imag = to_mont(pow_mod(G, (MOD - 1) / 4));
const __m256i vi = _mm256_set1_epi32((int)imag);
int lg = __builtin_ctz((unsigned)n);
int start = (lg & 1) ? (n >> 1) : n;
for (int len = 4; len <= start; len <<= 2) {
int q = len >> 2;
uint32_t *r1 = roots + off;
uint32_t *r2 = r1 + q;
uint32_t *r3 = r2 + q;
for (int i = 0; i < n; i += len) {
uint32_t *p = a + i;
int j = 0;
for (; j + 8 <= q; j += 8) {
__m256i a0 = _mm256_loadu_si256((const __m256i *)(p + j));
__m256i a1 = mont_mul_vec(_mm256_loadu_si256((const __m256i *)(p + j + q)),
_mm256_loadu_si256((const __m256i *)(r1 + j)));
__m256i a2 = mont_mul_vec(_mm256_loadu_si256((const __m256i *)(p + j + q + q)),
_mm256_loadu_si256((const __m256i *)(r2 + j)));
__m256i a3 = mont_mul_vec(_mm256_loadu_si256((const __m256i *)(p + j + q + q + q)),
_mm256_loadu_si256((const __m256i *)(r3 + j)));
__m256i s02 = add_vec(a0, a2);
__m256i d02 = sub_vec(a0, a2);
__m256i s13 = add_vec(a1, a3);
__m256i d13 = mont_mul_vec(sub_vec(a1, a3), vi);
_mm256_storeu_si256((__m256i *)(p + j), add_vec(s02, s13));
_mm256_storeu_si256((__m256i *)(p + j + q), sub_vec(d02, d13));
_mm256_storeu_si256((__m256i *)(p + j + q + q), sub_vec(s02, s13));
_mm256_storeu_si256((__m256i *)(p + j + q + q + q), add_vec(d02, d13));
}
for (; j < q; ++j) {
uint32_t a0 = p[j];
uint32_t a1 = mont_mul(p[j + q], r1[j]);
uint32_t a2 = mont_mul(p[j + q + q], r2[j]);
uint32_t a3 = mont_mul(p[j + q + q + q], r3[j]);
uint32_t s02 = add_mod(a0, a2);
uint32_t d02 = sub_mod(a0, a2);
uint32_t s13 = add_mod(a1, a3);
uint32_t d13 = mont_mul(sub_mod(a1, a3), imag);
p[j] = add_mod(s02, s13);
p[j + q] = sub_mod(d02, d13);
p[j + q + q] = sub_mod(s02, s13);
p[j + q + q + q] = add_mod(d02, d13);
}
}
off += q * 3;
}
if (lg & 1) {
int half = n >> 1;
uint32_t *rw = roots + off;
int j = 0;
for (; j + 8 <= half; j += 8) {
__m256i x = _mm256_loadu_si256((const __m256i *)(a + j));
__m256i y = mont_mul_vec(_mm256_loadu_si256((const __m256i *)(a + j + half)),
_mm256_loadu_si256((const __m256i *)(rw + j)));
_mm256_storeu_si256((__m256i *)(a + j), add_vec(x, y));
_mm256_storeu_si256((__m256i *)(a + j + half), sub_vec(x, y));
}
for (; j < half; ++j) {
uint32_t x = a[j];
uint32_t y = mont_mul(a[j + half], rw[j]);
a[j] = add_mod(x, y);
a[j + half] = sub_mod(x, y);
}
}
}
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(fa, len);
ntt_forward(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(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
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 87.983 ms | 31 MB + 660 KB | Accepted | Score: 100 | 显示更多 |