提交记录 88850


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 1002e7. 测测你的多项式乘法×10 Accepted 100 612.213 ms 471624 KB C++17 40.44 KB
提交时间 评测时间
2026-09-25 05:17:38 2026-09-25 05:17:43
// Blocked radix-4 DIF/DIT NTT mod 2013265921 using Shoup modular multiplication.
// Per stage: 6 tables (t1,t2,t3 = twiddle values; s1,s2,s3 = Shoup constants), contiguous.
// DRAM levels: block sizes n, n/4, ...; then in-cache base blocks of size ~2^15..2^19.
// DIF forward: natural in -> digit-reversed out.  DIT inverse: digit-reversed in -> natural out.
#pragma once
#include <immintrin.h>
#include <cstdint>
#include <cstring>
#include <cstdlib>

typedef uint32_t u32;
typedef uint64_t u64;

#define NTP 2013265921u
#define TGT __attribute__((target("avx2")))

#ifndef NTT_MAXN
#define NTT_MAXN (1u << 25)
#endif

static u32 g_n, g_Sb;
static u32 g_pinv;
static int g_log2n, g_L, g_odd, g_nb, g_BT = 15;
static u32 g_ciF, g_cisF, g_ciI, g_cisI;
static u32 g_tabF[5 * NTT_MAXN / 2 + 4096];
static u32 g_tabI[5 * NTT_MAXN / 2 + 4096];

struct Tabs {
  const u32 *l1[40], *l2[40], *l3[40], *m1[40], *m2[40], *m3[40];
  const u32 *b2, *b2s;
  const u32 *h1[40], *h2[40], *h3[40], *q1[40], *q2[40], *q3[40];
};
static Tabs g_TF, g_TI;

static inline u32 mulmod(u32 a, u32 b) { return (u32)((u64)a * b % NTP); }
static u32 powmod(u32 a, u32 e) { u32 r = 1; while (e) { if (e & 1) r = mulmod(r, a); a = mulmod(a, a); e >>= 1; } return r; }
static inline u32 shoup1s(u32 a, u32 w, u32 ws) {
  u32 t = (u32)((u64)a * w);
  u32 q = (u32)(((u64)a * ws) >> 32);
  u32 r = t - (u32)((u64)q * NTP);
  return r >= NTP ? r - NTP : r;
}
static inline u32 shoup_c(u32 w) { return (u32)(((u64)w << 32) / NTP); }

TGT static inline __m256i vmulhi(__m256i a, __m256i b) {
  __m256i e = _mm256_mul_epu32(a, b);
  __m256i o = _mm256_mul_epu32(_mm256_srli_epi64(a, 32), _mm256_srli_epi64(b, 32));
  e = _mm256_srli_epi64(e, 32);
  return _mm256_blend_epi32(e, o, 0xAA);
}
TGT static inline __m256i vmont(__m256i a, __m256i b, __m256i pv) {
  const __m256i pinv = _mm256_set1_epi32((int)g_pinv);
  __m256i tlo = _mm256_mullo_epi32(a, b);
  __m256i thi = vmulhi(a, b);
  __m256i m   = _mm256_mullo_epi32(tlo, pinv);
  __m256i mhi = vmulhi(m, pv);
  __m256i nz  = _mm256_andnot_si256(_mm256_cmpeq_epi32(tlo, _mm256_setzero_si256()), _mm256_set1_epi32(1));
  __m256i u   = _mm256_add_epi32(_mm256_add_epi32(thi, mhi), nz);
  return _mm256_min_epu32(u, _mm256_sub_epi32(u, pv));
}
// M = floor(2^64/p) = 2^33 + MLO ;  ws ~= 2*w + mulhi(w, MLO)  (error <= 1, valid for Shoup)
#define MLO_V 572662301u
TGT static inline __m256i der_ws(__m256i w) {
  __m256i mlo = _mm256_set1_epi32((int)MLO_V);
  return _mm256_add_epi32(_mm256_add_epi32(w, w), vmulhi(w, mlo));
}

TGT static inline __m256i vshoup(__m256i a, __m256i w, __m256i ws, __m256i pv) {
  __m256i t = _mm256_mullo_epi32(a, w);
  __m256i q = vmulhi(a, ws);
  __m256i r = _mm256_sub_epi32(t, _mm256_mullo_epi32(q, pv));
  return _mm256_min_epu32(r, _mm256_sub_epi32(r, pv));
}
TGT static inline __m256i vaddm(__m256i a, __m256i b, __m256i pv) {
  __m256i s = _mm256_add_epi32(a, b);
  return _mm256_min_epu32(s, _mm256_sub_epi32(s, pv));
}
TGT static inline __m256i vsubm(__m256i a, __m256i b, __m256i pv) {
  __m256i d = _mm256_sub_epi32(a, b);
  return _mm256_min_epu32(d, _mm256_add_epi32(d, pv));
}
TGT static inline __m128i vmulhi4(__m128i a, __m128i b) {
  __m128i e = _mm_mul_epu32(a, b);
  __m128i o = _mm_mul_epu32(_mm_srli_epi64(a, 32), _mm_srli_epi64(b, 32));
  e = _mm_srli_epi64(e, 32);
  return _mm_blend_epi32(e, o, 0xAA);
}
TGT static inline __m128i sh4(__m128i a, __m128i w, __m128i ws, __m128i pv) {
  __m128i t = _mm_mullo_epi32(a, w);
  __m128i q = vmulhi4(a, ws);
  __m128i r = _mm_sub_epi32(t, _mm_mullo_epi32(q, pv));
  return _mm_min_epu32(r, _mm_sub_epi32(r, pv));
}
TGT static inline __m128i add4(__m128i a, __m128i b, __m128i pv) {
  __m128i s = _mm_add_epi32(a, b);
  return _mm_min_epu32(s, _mm_sub_epi32(s, pv));
}
TGT static inline __m128i sub4(__m128i a, __m128i b, __m128i pv) {
  __m128i d = _mm_sub_epi32(a, b);
  return _mm_min_epu32(d, _mm_add_epi32(d, pv));
}


TGT static inline __m256i ld2(const __m128i *lo, const __m128i *hi) {
  return _mm256_inserti128_si256(_mm256_castsi128_si256(_mm_loadu_si128(lo)), _mm_loadu_si128(hi), 1);
}
TGT static inline void st2(__m128i *lo, __m128i *hi, __m256i v) {
  _mm_storeu_si128(lo, _mm256_castsi256_si128(v));
  _mm_storeu_si128(hi, _mm256_extracti128_si256(v, 1));
}
TGT static inline __m256i bc128(__m128i v) {
  return _mm256_inserti128_si256(_mm256_castsi128_si256(v), v, 1);
}

// fused level-0 DIF stage: reads coefficients directly from src ([0,len] valid, zero beyond).
// Requires len >= h and len < 2h+... (x2 = x3 = 0 assumed).
TGT static void stage0_dif_src(u32 *a, const u32 *src, u32 len, u32 h,
                               const u32 *t1, const u32 *t2, const u32 *t3, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  const __m256i zero = _mm256_setzero_si256();
  u32 j = 0;
  u32 lim1 = len - h;
  for (; j + 8 <= lim1 + 1; j += 8) {
    __builtin_prefetch(src + j + 256);
    __builtin_prefetch(src + j + h + 256);
    __builtin_prefetch(a + j + 256);
    __builtin_prefetch(a + j + h + 256);
    __builtin_prefetch(a + j + 2 * h + 256);
    __builtin_prefetch(a + j + 3 * h + 256);
    __builtin_prefetch(t1 + j + 256);
    __m256i x0 = _mm256_loadu_si256((const __m256i *)(src + j));
    __m256i x1 = _mm256_loadu_si256((const __m256i *)(src + j + h));
    __m256i u0 = vaddm(x0, x1, pv), u2 = vsubm(x0, x1, pv);
    __m256i E = vshoup(x1, cI, cIs, pv);
    __m256i u1 = vaddm(x0, E, pv), u3 = vsubm(x0, E, pv);
    { __m256i w = _mm256_loadu_si256((const __m256i *)(t1 + j)); u1 = vshoup(u1, w, der_ws(w), pv); }
    { __m256i w = _mm256_loadu_si256((const __m256i *)(t2 + j)); u2 = vshoup(u2, w, der_ws(w), pv); }
    { __m256i w = _mm256_loadu_si256((const __m256i *)(t3 + j)); u3 = vshoup(u3, w, der_ws(w), pv); }
    _mm256_storeu_si256((__m256i *)(a + j), u0);
    _mm256_storeu_si256((__m256i *)(a + j + h), u1);
    _mm256_storeu_si256((__m256i *)(a + j + 2 * h), u2);
    _mm256_storeu_si256((__m256i *)(a + j + 3 * h), u3);
  }
  for (; j <= lim1 && j < h; j++) {
    u32 x0 = (j <= len) ? src[j] : 0, x1 = (j + h <= len) ? src[j + h] : 0;
    u32 u0 = (u32)(((u64)x0 + x1) % NTP);
    u32 u2 = (u32)(((u64)x0 + NTP - x1) % NTP);
    u32 E = mulmod(x1, ci);
    u32 u1 = (u32)(((u64)x0 + E) % NTP), u3 = (u32)(((u64)x0 + NTP - E) % NTP);
    u1 = shoup1s(u1, t1[j], shoup_c(t1[j]));
    u2 = shoup1s(u2, t2[j], shoup_c(t2[j]));
    u3 = shoup1s(u3, t3[j], shoup_c(t3[j]));
    a[j] = u0; a[j + h] = u1; a[j + 2 * h] = u2; a[j + 3 * h] = u3;
  }
  for (; j + 8 <= h; j += 8) {
    __m256i x0 = (j + 8 <= len + 1) ? _mm256_loadu_si256((const __m256i *)(src + j)) : zero;
    __m256i u0 = x0, u1 = x0, u2 = x0, u3 = x0;
    { __m256i w = _mm256_loadu_si256((const __m256i *)(t1 + j)); u1 = vshoup(u1, w, der_ws(w), pv); }
    { __m256i w = _mm256_loadu_si256((const __m256i *)(t2 + j)); u2 = vshoup(u2, w, der_ws(w), pv); }
    { __m256i w = _mm256_loadu_si256((const __m256i *)(t3 + j)); u3 = vshoup(u3, w, der_ws(w), pv); }
    _mm256_storeu_si256((__m256i *)(a + j), u0);
    _mm256_storeu_si256((__m256i *)(a + j + h), u1);
    _mm256_storeu_si256((__m256i *)(a + j + 2 * h), u2);
    _mm256_storeu_si256((__m256i *)(a + j + 3 * h), u3);
  }
  for (; j < h; j++) {
    u32 x0 = (j <= len) ? src[j] : 0;
    u32 u1 = shoup1s(x0, t1[j], shoup_c(t1[j]));
    u32 u2 = shoup1s(x0, t2[j], shoup_c(t2[j]));
    u32 u3 = shoup1s(x0, t3[j], shoup_c(t3[j]));
    a[j] = x0; a[j + h] = u1; a[j + 2 * h] = u2; a[j + 3 * h] = u3;
  }
}

// ================= forward (DIF) =================
TGT static void stage4_dif(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
                           const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  for (u32 i = 0; i < blk; i += 4 * h) {
    u32 *p = a + i;
    u32 j = 0;
    for (; j + 16 <= h; j += 16) {
      __builtin_prefetch(p + j + 256);
      __builtin_prefetch(p + j + h + 256);
      __builtin_prefetch(p + j + 2 * h + 256);
      __builtin_prefetch(p + j + 3 * h + 256);
      __builtin_prefetch(t1 + j + 256);
      __builtin_prefetch(t2 + j + 256);
      __builtin_prefetch(t3 + j + 256);
      __m256i x0 = _mm256_loadu_si256((const __m256i *)(p + j));
      __m256i x1 = _mm256_loadu_si256((const __m256i *)(p + j + h));
      __m256i x2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h));
      __m256i x3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h));
      __m256i y0 = _mm256_loadu_si256((const __m256i *)(p + j + 8));
      __m256i y1 = _mm256_loadu_si256((const __m256i *)(p + j + h + 8));
      __m256i y2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h + 8));
      __m256i y3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h + 8));
      __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
      __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
      __m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
      __m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
      __m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
      __m256i v0 = vaddm(A2, C2, pv), v2 = vsubm(A2, C2, pv);
      __m256i E = vshoup(D, cI, cIs, pv);
      __m256i F = vshoup(D2, cI, cIs, pv);
      __m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
      __m256i v1 = vaddm(B2, F, pv), v3 = vsubm(B2, F, pv);
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t1 + j)); u1 = vshoup(u1, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t2 + j)); u2 = vshoup(u2, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t3 + j)); u3 = vshoup(u3, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t1 + j + 8)); v1 = vshoup(v1, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t2 + j + 8)); v2 = vshoup(v2, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t3 + j + 8)); v3 = vshoup(v3, w, der_ws(w), pv); }
      _mm256_storeu_si256((__m256i *)(p + j), u0);
      _mm256_storeu_si256((__m256i *)(p + j + h), u1);
      _mm256_storeu_si256((__m256i *)(p + j + 2 * h), u2);
      _mm256_storeu_si256((__m256i *)(p + j + 3 * h), u3);
      _mm256_storeu_si256((__m256i *)(p + j + 8), v0);
      _mm256_storeu_si256((__m256i *)(p + j + h + 8), v1);
      _mm256_storeu_si256((__m256i *)(p + j + 2 * h + 8), v2);
      _mm256_storeu_si256((__m256i *)(p + j + 3 * h + 8), v3);
    }
    for (; j + 8 <= h; j += 8) {
      __builtin_prefetch(p + j + 256);
      __builtin_prefetch(p + j + h + 256);
      __builtin_prefetch(p + j + 2 * h + 256);
      __builtin_prefetch(p + j + 3 * h + 256);
      __m256i x0 = _mm256_loadu_si256((const __m256i *)(p + j));
      __m256i x1 = _mm256_loadu_si256((const __m256i *)(p + j + h));
      __m256i x2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h));
      __m256i x3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h));
      __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
      __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
      __m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
      __m256i E = vshoup(D, cI, cIs, pv);
      __m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t1 + j)); u1 = vshoup(u1, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t2 + j)); u2 = vshoup(u2, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t3 + j)); u3 = vshoup(u3, w, der_ws(w), pv); }
      _mm256_storeu_si256((__m256i *)(p + j), u0);
      _mm256_storeu_si256((__m256i *)(p + j + h), u1);
      _mm256_storeu_si256((__m256i *)(p + j + 2 * h), u2);
      _mm256_storeu_si256((__m256i *)(p + j + 3 * h), u3);
    }
  }
}
TGT static void stage4_dif_nt(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
                           const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  for (u32 i = 0; i < blk; i += 4 * h) {
    u32 *p = a + i;
    u32 j = 0;
    for (; j + 16 <= h; j += 16) {
      __m256i x0 = _mm256_loadu_si256((const __m256i *)(p + j));
      __m256i x1 = _mm256_loadu_si256((const __m256i *)(p + j + h));
      __m256i x2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h));
      __m256i x3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h));
      __m256i y0 = _mm256_loadu_si256((const __m256i *)(p + j + 8));
      __m256i y1 = _mm256_loadu_si256((const __m256i *)(p + j + h + 8));
      __m256i y2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h + 8));
      __m256i y3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h + 8));
      __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
      __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
      __m256i A2 = vaddm(y0, y2, pv), B2 = vsubm(y0, y2, pv);
      __m256i C2 = vaddm(y1, y3, pv), D2 = vsubm(y1, y3, pv);
      __m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
      __m256i v0 = vaddm(A2, C2, pv), v2 = vsubm(A2, C2, pv);
      __m256i E = vshoup(D, cI, cIs, pv);
      __m256i F = vshoup(D2, cI, cIs, pv);
      __m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
      __m256i v1 = vaddm(B2, F, pv), v3 = vsubm(B2, F, pv);
      u1 = vshoup(u1, _mm256_loadu_si256((const __m256i *)(t1 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t1 + j))), pv);
      u2 = vshoup(u2, _mm256_loadu_si256((const __m256i *)(t2 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t2 + j))), pv);
      u3 = vshoup(u3, _mm256_loadu_si256((const __m256i *)(t3 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t3 + j))), pv);
      v1 = vshoup(v1, _mm256_loadu_si256((const __m256i *)(t1 + j + 8)), der_ws(_mm256_loadu_si256((const __m256i *)(t1 + j + 8))), pv);
      v2 = vshoup(v2, _mm256_loadu_si256((const __m256i *)(t2 + j + 8)), der_ws(_mm256_loadu_si256((const __m256i *)(t2 + j + 8))), pv);
      v3 = vshoup(v3, _mm256_loadu_si256((const __m256i *)(t3 + j + 8)), der_ws(_mm256_loadu_si256((const __m256i *)(t3 + j + 8))), pv);
      _mm256_stream_si256((__m256i *)(p + j), u0);
      _mm256_stream_si256((__m256i *)(p + j + h), u1);
      _mm256_stream_si256((__m256i *)(p + j + 2 * h), u2);
      _mm256_stream_si256((__m256i *)(p + j + 3 * h), u3);
      _mm256_stream_si256((__m256i *)(p + j + 8), v0);
      _mm256_stream_si256((__m256i *)(p + j + h + 8), v1);
      _mm256_stream_si256((__m256i *)(p + j + 2 * h + 8), v2);
      _mm256_stream_si256((__m256i *)(p + j + 3 * h + 8), v3);
    }
    for (; j + 8 <= h; j += 8) {
      __m256i x0 = _mm256_loadu_si256((const __m256i *)(p + j));
      __m256i x1 = _mm256_loadu_si256((const __m256i *)(p + j + h));
      __m256i x2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h));
      __m256i x3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h));
      __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
      __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
      __m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
      __m256i E = vshoup(D, cI, cIs, pv);
      __m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
      u1 = vshoup(u1, _mm256_loadu_si256((const __m256i *)(t1 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t1 + j))), pv);
      u2 = vshoup(u2, _mm256_loadu_si256((const __m256i *)(t2 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t2 + j))), pv);
      u3 = vshoup(u3, _mm256_loadu_si256((const __m256i *)(t3 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t3 + j))), pv);
      _mm256_stream_si256((__m256i *)(p + j), u0);
      _mm256_stream_si256((__m256i *)(p + j + h), u1);
      _mm256_stream_si256((__m256i *)(p + j + 2 * h), u2);
      _mm256_stream_si256((__m256i *)(p + j + 3 * h), u3);
    }
  }
}
TGT static void stage4_dif_h4(u32 *a, u32 nsub, const u32 *t1, const u32 *t2, const u32 *t3,
                              const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  if (nsub < 32) {
    const __m128i pv4 = _mm256_castsi256_si128(pv), cI4 = _mm256_castsi256_si128(cI), cIs4 = _mm256_castsi256_si128(cIs);
    __m128i W1 = _mm_loadu_si128((const __m128i *)t1), X1 = _mm_loadu_si128((const __m128i *)s1);
    __m128i W2 = _mm_loadu_si128((const __m128i *)t2), X2 = _mm_loadu_si128((const __m128i *)s2);
    __m128i W3 = _mm_loadu_si128((const __m128i *)t3), X3 = _mm_loadu_si128((const __m128i *)s3);
    for (u32 i = 0; i < nsub; i += 16) {
      u32 *p = a + i;
      __m128i x0 = _mm_loadu_si128((const __m128i *)(p));
      __m128i x1 = _mm_loadu_si128((const __m128i *)(p + 4));
      __m128i x2 = _mm_loadu_si128((const __m128i *)(p + 8));
      __m128i x3 = _mm_loadu_si128((const __m128i *)(p + 12));
      __m128i A = add4(x0, x2, pv4), B = sub4(x0, x2, pv4);
      __m128i C = add4(x1, x3, pv4), D = sub4(x1, x3, pv4);
      __m128i u0 = add4(A, C, pv4), u2 = sub4(A, C, pv4);
      __m128i E = sh4(D, cI4, cIs4, pv4);
      __m128i u1 = add4(B, E, pv4), u3 = sub4(B, E, pv4);
      u1 = sh4(u1, W1, X1, pv4); u2 = sh4(u2, W2, X2, pv4); u3 = sh4(u3, W3, X3, pv4);
      _mm_storeu_si128((__m128i *)(p), u0);
      _mm_storeu_si128((__m128i *)(p + 4), u1);
      _mm_storeu_si128((__m128i *)(p + 8), u2);
      _mm_storeu_si128((__m128i *)(p + 12), u3);
    }
    return;
  }
  __m256i W1 = bc128(_mm_loadu_si128((const __m128i *)t1));
  __m256i W2 = bc128(_mm_loadu_si128((const __m128i *)t2));
  __m256i W3 = bc128(_mm_loadu_si128((const __m128i *)t3));
  __m256i X1 = bc128(_mm_loadu_si128((const __m128i *)s1));
  __m256i X2 = bc128(_mm_loadu_si128((const __m128i *)s2));
  __m256i X3 = bc128(_mm_loadu_si128((const __m128i *)s3));
  for (u32 i = 0; i < nsub; i += 32) {
    u32 *p = a + i;
    __m256i x0 = ld2((const __m128i *)(p), (const __m128i *)(p + 16));
    __m256i x1 = ld2((const __m128i *)(p + 4), (const __m128i *)(p + 20));
    __m256i x2 = ld2((const __m128i *)(p + 8), (const __m128i *)(p + 24));
    __m256i x3 = ld2((const __m128i *)(p + 12), (const __m128i *)(p + 28));
    __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
    __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
    __m256i u0 = vaddm(A, C, pv), u2 = vsubm(A, C, pv);
    __m256i E = vshoup(D, cI, cIs, pv);
    __m256i u1 = vaddm(B, E, pv), u3 = vsubm(B, E, pv);
    u1 = vshoup(u1, W1, X1, pv); u2 = vshoup(u2, W2, X2, pv); u3 = vshoup(u3, W3, X3, pv);
    st2((__m128i *)(p), (__m128i *)(p + 16), u0);
    st2((__m128i *)(p + 4), (__m128i *)(p + 20), u1);
    st2((__m128i *)(p + 8), (__m128i *)(p + 24), u2);
    st2((__m128i *)(p + 12), (__m128i *)(p + 28), u3);
  }
}
TGT static void stage4_dif_h1(u32 *a, u32 blk, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  for (u32 i = 0; i < blk; i += 8) {
    u32 *p = a + i;
    __m256i x = _mm256_loadu_si256((const __m256i *)p);
    __m256i t = _mm256_shuffle_epi32(x, _MM_SHUFFLE(1, 0, 3, 2));
    __m256i s = vaddm(x, t, pv);
    __m256i d = vsubm(x, t, pv);
    __m256i Z = _mm256_shuffle_epi32(d, 0x00);
    __m256i P = _mm256_blend_epi32(s, Z, 0xAA);
    __m256i E = vshoup(d, cI, cIs, pv);
    __m256i E2 = _mm256_shuffle_epi32(E, 0x55);
    __m256i W = _mm256_shuffle_epi32(s, 0x55);
    __m256i Q = _mm256_blend_epi32(W, E2, 0xAA);
    __m256i sum = vaddm(P, Q, pv);
    __m256i dif = vsubm(P, Q, pv);
    _mm256_storeu_si256((__m256i *)p, _mm256_blend_epi32(sum, dif, 0xCC));
  }
}
TGT static void stage2_dif(u32 *a, u32 blk, const u32 *t, const u32 *s) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  u32 h = blk >> 1;
  for (u32 j = 0; j < h; j += 8) {
    __m256i x0 = _mm256_loadu_si256((const __m256i *)(a + j));
    __m256i x1 = _mm256_loadu_si256((const __m256i *)(a + j + h));
    __m256i u0 = vaddm(x0, x1, pv);
    __m256i u1 = vshoup(vsubm(x0, x1, pv), _mm256_loadu_si256((const __m256i *)(t + j)),
                        _mm256_loadu_si256((const __m256i *)(s + j)), pv);
    _mm256_storeu_si256((__m256i *)(a + j), u0);
    _mm256_storeu_si256((__m256i *)(a + j + h), u1);
  }
}

// ================= inverse (DIT) =================
TGT static void merge4_dit(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
                           const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  for (u32 i = 0; i < blk; i += 4 * h) {
    u32 *p = a + i;
    for (u32 j = 0; j < h; j += 8) {
      __builtin_prefetch(p + j + 256);
      __builtin_prefetch(p + j + h + 256);
      __builtin_prefetch(p + j + 2 * h + 256);
      __builtin_prefetch(p + j + 3 * h + 256);
      __builtin_prefetch(t1 + j + 256);
      __builtin_prefetch(t2 + j + 256);
      __builtin_prefetch(t3 + j + 256);
      __m256i x0 = _mm256_loadu_si256((const __m256i *)(p + j));
      __m256i x1 = _mm256_loadu_si256((const __m256i *)(p + j + h));
      __m256i x2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h));
      __m256i x3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h));
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t1 + j)); x1 = vshoup(x1, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t2 + j)); x2 = vshoup(x2, w, der_ws(w), pv); }
      { __m256i w = _mm256_loadu_si256((const __m256i *)(t3 + j)); x3 = vshoup(x3, w, der_ws(w), pv); }
      __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
      __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
      __m256i E = vshoup(D, cI, cIs, pv);
      _mm256_storeu_si256((__m256i *)(p + j), vaddm(A, C, pv));
      _mm256_storeu_si256((__m256i *)(p + j + h), vaddm(B, E, pv));
      _mm256_storeu_si256((__m256i *)(p + j + 2 * h), vsubm(A, C, pv));
      _mm256_storeu_si256((__m256i *)(p + j + 3 * h), vsubm(B, E, pv));
    }
  }
}
TGT static void merge4_dit_nt(u32 *a, u32 blk, u32 h, const u32 *t1, const u32 *t2, const u32 *t3,
                           const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  for (u32 i = 0; i < blk; i += 4 * h) {
    u32 *p = a + i;
    for (u32 j = 0; j < h; j += 8) {
      __m256i x0 = _mm256_loadu_si256((const __m256i *)(p + j));
      __m256i x1 = _mm256_loadu_si256((const __m256i *)(p + j + h));
      __m256i x2 = _mm256_loadu_si256((const __m256i *)(p + j + 2 * h));
      __m256i x3 = _mm256_loadu_si256((const __m256i *)(p + j + 3 * h));
      x1 = vshoup(x1, _mm256_loadu_si256((const __m256i *)(t1 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t1 + j))), pv);
      x2 = vshoup(x2, _mm256_loadu_si256((const __m256i *)(t2 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t2 + j))), pv);
      x3 = vshoup(x3, _mm256_loadu_si256((const __m256i *)(t3 + j)), der_ws(_mm256_loadu_si256((const __m256i *)(t3 + j))), pv);
      __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
      __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
      __m256i E = vshoup(D, cI, cIs, pv);
      _mm256_stream_si256((__m256i *)(p + j), vaddm(A, C, pv));
      _mm256_stream_si256((__m256i *)(p + j + h), vaddm(B, E, pv));
      _mm256_stream_si256((__m256i *)(p + j + 2 * h), vsubm(A, C, pv));
      _mm256_stream_si256((__m256i *)(p + j + 3 * h), vsubm(B, E, pv));
    }
  }
}
TGT static void merge4_dit_h4(u32 *a, u32 nsub, const u32 *t1, const u32 *t2, const u32 *t3,
                              const u32 *s1, const u32 *s2, const u32 *s3, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  if (nsub < 32) {
    const __m128i pv4 = _mm256_castsi256_si128(pv), cI4 = _mm256_castsi256_si128(cI), cIs4 = _mm256_castsi256_si128(cIs);
    __m128i W1 = _mm_loadu_si128((const __m128i *)t1), X1 = _mm_loadu_si128((const __m128i *)s1);
    __m128i W2 = _mm_loadu_si128((const __m128i *)t2), X2 = _mm_loadu_si128((const __m128i *)s2);
    __m128i W3 = _mm_loadu_si128((const __m128i *)t3), X3 = _mm_loadu_si128((const __m128i *)s3);
    for (u32 i = 0; i < nsub; i += 16) {
      u32 *p = a + i;
      __m128i x0 = _mm_loadu_si128((const __m128i *)(p));
      __m128i x1 = sh4(_mm_loadu_si128((const __m128i *)(p + 4)), W1, X1, pv4);
      __m128i x2 = sh4(_mm_loadu_si128((const __m128i *)(p + 8)), W2, X2, pv4);
      __m128i x3 = sh4(_mm_loadu_si128((const __m128i *)(p + 12)), W3, X3, pv4);
      __m128i A = add4(x0, x2, pv4), B = sub4(x0, x2, pv4);
      __m128i C = add4(x1, x3, pv4), D = sub4(x1, x3, pv4);
      __m128i E = sh4(D, cI4, cIs4, pv4);
      _mm_storeu_si128((__m128i *)(p), add4(A, C, pv4));
      _mm_storeu_si128((__m128i *)(p + 4), add4(B, E, pv4));
      _mm_storeu_si128((__m128i *)(p + 8), sub4(A, C, pv4));
      _mm_storeu_si128((__m128i *)(p + 12), sub4(B, E, pv4));
    }
    return;
  }
  __m256i W1 = bc128(_mm_loadu_si128((const __m128i *)t1));
  __m256i W2 = bc128(_mm_loadu_si128((const __m128i *)t2));
  __m256i W3 = bc128(_mm_loadu_si128((const __m128i *)t3));
  __m256i X1 = bc128(_mm_loadu_si128((const __m128i *)s1));
  __m256i X2 = bc128(_mm_loadu_si128((const __m128i *)s2));
  __m256i X3 = bc128(_mm_loadu_si128((const __m128i *)s3));
  for (u32 i = 0; i < nsub; i += 32) {
    u32 *p = a + i;
    __m256i x0 = ld2((const __m128i *)(p), (const __m128i *)(p + 16));
    __m256i x1 = ld2((const __m128i *)(p + 4), (const __m128i *)(p + 20));
    __m256i x2 = ld2((const __m128i *)(p + 8), (const __m128i *)(p + 24));
    __m256i x3 = ld2((const __m128i *)(p + 12), (const __m128i *)(p + 28));
    x1 = vshoup(x1, W1, X1, pv); x2 = vshoup(x2, W2, X2, pv); x3 = vshoup(x3, W3, X3, pv);
    __m256i A = vaddm(x0, x2, pv), B = vsubm(x0, x2, pv);
    __m256i C = vaddm(x1, x3, pv), D = vsubm(x1, x3, pv);
    __m256i E = vshoup(D, cI, cIs, pv);
    st2((__m128i *)(p), (__m128i *)(p + 16), vaddm(A, C, pv));
    st2((__m128i *)(p + 4), (__m128i *)(p + 20), vaddm(B, E, pv));
    st2((__m128i *)(p + 8), (__m128i *)(p + 24), vsubm(A, C, pv));
    st2((__m128i *)(p + 12), (__m128i *)(p + 28), vsubm(B, E, pv));
  }
}
TGT static void merge4_dit_h1(u32 *a, u32 blk, u32 ci, u32 cis) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i cI = _mm256_set1_epi32((int)ci), cIs = _mm256_set1_epi32((int)cis);
  for (u32 i = 0; i < blk; i += 8) {
    u32 *p = a + i;
    __m256i x = _mm256_loadu_si256((const __m256i *)p);
    __m256i t = _mm256_shuffle_epi32(x, _MM_SHUFFLE(1, 0, 3, 2));
    __m256i s = vaddm(x, t, pv);
    __m256i d = vsubm(x, t, pv);
    __m256i Z = _mm256_shuffle_epi32(d, 0x00);
    __m256i P = _mm256_blend_epi32(s, Z, 0xAA);
    __m256i E = vshoup(d, cI, cIs, pv);
    __m256i E2 = _mm256_shuffle_epi32(E, 0x55);
    __m256i W = _mm256_shuffle_epi32(s, 0x55);
    __m256i Q = _mm256_blend_epi32(W, E2, 0xAA);
    __m256i sum = vaddm(P, Q, pv);
    __m256i dif = vsubm(P, Q, pv);
    _mm256_storeu_si256((__m256i *)p, _mm256_blend_epi32(sum, dif, 0xCC));
  }
}
TGT static void merge2_dit(u32 *a, u32 blk, const u32 *t, const u32 *s) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  u32 h = blk >> 1;
  for (u32 j = 0; j < h; j += 8) {
    __m256i x0 = _mm256_loadu_si256((const __m256i *)(a + j));
    __m256i x1 = vshoup(_mm256_loadu_si256((const __m256i *)(a + j + h)),
                        _mm256_loadu_si256((const __m256i *)(t + j)),
                        _mm256_loadu_si256((const __m256i *)(s + j)), pv);
    _mm256_storeu_si256((__m256i *)(a + j), vaddm(x0, x1, pv));
    _mm256_storeu_si256((__m256i *)(a + j + h), vsubm(x0, x1, pv));
  }
}


// ================= table generation =================
// dst[i] = v0*r^i ; sc[i] = floor(dst[i]*2^32/p) (within 1, sufficient for Shoup)
TGT static void gen_geo(u32 *dst, u32 *sc, u32 cnt, u32 v0, u32 r) {
  if (!cnt) return;
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const u32 Mlo = 572662301u;
  __m256i mlo = _mm256_set1_epi32((int)Mlo);
  u32 seq[8], cur = v0 % NTP;
  for (int i = 0; i < 8; i++) { seq[i] = cur; cur = mulmod(cur, r); }
  u32 r8 = 1; for (int t = 0; t < 8; t++) r8 = mulmod(r8, r);
  u32 rs8 = shoup_c(r8);
  __m256i v = _mm256_loadu_si256((const __m256i *)seq);
  __m256i r8v = _mm256_set1_epi32((int)r8), rs8v = _mm256_set1_epi32((int)rs8);
  u32 i = 0;
  for (; i + 8 <= cnt; i += 8) {
    _mm256_storeu_si256((__m256i *)(dst + i), v);
    __m256i two = _mm256_add_epi32(v, v);
    _mm256_storeu_si256((__m256i *)(sc + i), _mm256_add_epi32(two, vmulhi(v, mlo)));
    v = vshoup(v, r8v, rs8v, pv);
  }
  if (i < cnt) {
    u32 x = v0 % NTP;
    for (u32 t = 0; t < i; t++) x = mulmod(x, r);
    for (; i < cnt; i++) { dst[i] = x; sc[i] = shoup_c(x); x = mulmod(x, r); }
  }
}

TGT static void gen_geo_val(u32 *dst, u32 cnt, u32 v0, u32 r) {
  if (!cnt) return;
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  u32 seq[8], cur = v0 % NTP;
  for (int i = 0; i < 8; i++) { seq[i] = cur; cur = mulmod(cur, r); }
  u32 r8 = 1; for (int t = 0; t < 8; t++) r8 = mulmod(r8, r);
  u32 rs8 = shoup_c(r8);
  __m256i v = _mm256_loadu_si256((const __m256i *)seq);
  __m256i r8v = _mm256_set1_epi32((int)r8), rs8v = _mm256_set1_epi32((int)rs8);
  u32 i = 0;
  for (; i + 8 <= cnt; i += 8) { _mm256_storeu_si256((__m256i *)(dst + i), v); v = vshoup(v, r8v, rs8v, pv); }
  if (i < cnt) { u32 x = v0 % NTP; for (u32 t = 0; t < i; t++) x = mulmod(x, r); for (; i < cnt; i++) { dst[i] = x; x = mulmod(x, r); } }
}

static void build_tabs(Tabs &T, u32 *tab, u32 n, u32 Sb, int L, int odd, u32 root) {
  u32 *p = tab;
  u32 S = n;
  for (int l = 0; l < L; l++) {
    u32 M = S >> 2;
    T.l1[l] = p; T.l2[l] = p + M; T.l3[l] = p + 2 * M;
    T.m1[l] = p; T.m2[l] = p + M; T.m3[l] = p + 2 * M;
    u32 base = powmod(root, n / S);
    u32 b1 = base, b2 = mulmod(base, base), b3 = mulmod(b2, base);
    gen_geo_val(p, M, 1, b1);
    gen_geo_val(p + M, M, 1, b2);
    gen_geo_val(p + 2 * M, M, 1, b3);
    p += 3 * M;
    S >>= 2;
  }
  T.b2 = 0; T.b2s = 0;
  if (odd) {
    u32 M = Sb >> 1;
    T.b2 = p; T.b2s = p + M;
    gen_geo(p, p + M, M, 1, powmod(root, n / Sb));
    p += 2 * M;
  }
  int s = 0;
  u32 S2v = 4;
  while (s < g_nb) {
    u32 M = S2v >> 2;
    T.h1[s] = p; T.h2[s] = p + M; T.h3[s] = p + 2 * M;
    T.q1[s] = p + 3 * M; T.q2[s] = p + 4 * M; T.q3[s] = p + 5 * M;
    u32 base = powmod(root, n / S2v);
    u32 b1 = base, b2 = mulmod(base, base), b3 = mulmod(b2, base);
    gen_geo(p, p + 3 * M, M, 1, b1);
    gen_geo(p + M, p + 4 * M, M, 1, b2);
    gen_geo(p + 2 * M, p + 5 * M, M, 1, b3);
    p += 6 * M;
    S2v <<= 2;
    s++;
  }
}

static u32 g_tabwords;
static void ntt_init(u32 n) {
  { u32 iv = 1; for (int q = 0; q < 5; q++) iv *= 2u - NTP * iv; g_pinv = 0u - iv; }
  g_n = n;
  int lg = 0; while ((1u << lg) < n) lg++;
  g_log2n = lg;
  const int BT = g_BT;
  int L = 0;
  while (lg - 2 * (L + 1) >= BT) L++;
  g_L = L;
  g_Sb = n >> (2 * L);
  g_odd = ((lg - 2 * L) & 1) ? 1 : 0;
  u32 S2 = g_Sb; if (g_odd) S2 >>= 1;
  g_nb = 0; while (S2 >= 4) { S2 >>= 2; g_nb++; }
  u32 w = powmod(31u, (NTP - 1) / n);
  u32 wi = powmod(w, NTP - 2);
  u32 iF = powmod(w, n >> 2), iI = powmod(wi, n >> 2);
  g_ciF = iF; g_cisF = shoup_c(iF);
  g_ciI = iI; g_cisI = shoup_c(iI);
  { u32 S=n; u32 tot=0; for (int l=0;l<L;l++){tot+=3*(S>>2);S>>=2;} u32 S2v=4; for(int s=0;s<g_nb;s++){tot+=6*(S2v>>2);S2v<<=2;} if(g_odd) tot+=2*(g_Sb>>1); g_tabwords=tot; }
  build_tabs(g_TF, g_tabF, n, g_Sb, L, g_odd, w);
  build_tabs(g_TI, g_tabI, n, g_Sb, L, g_odd, wi);
}

TGT static void ntt_fwd(u32 *a, const u32 *src = 0, u32 len = 0) {
  u32 S = g_n;
  int l0 = 0;
  if (src && len >= (g_n >> 2) && len <= (g_n >> 1)) {
    stage0_dif_src(a, src, len, g_n >> 2, g_TF.l1[0], g_TF.l2[0], g_TF.l3[0], g_ciF, g_cisF);
    S = g_n >> 2;
    l0 = 1;
  }
  for (int l = l0; l < g_L; l++) {
    u32 h = S >> 2;
    for (u32 b = 0; b < g_n; b += S)
      stage4_dif(a + b, S, h, g_TF.l1[l], g_TF.l2[l], g_TF.l3[l], g_TF.m1[l], g_TF.m2[l], g_TF.m3[l], g_ciF, g_cisF);
    _mm_sfence();
    S = h;
  }
  for (u32 b = 0; b < g_n; b += g_Sb) {
    u32 *p = a + b;
    if (g_odd) stage2_dif(p, g_Sb, g_TF.b2, g_TF.b2s);
    for (int s = g_nb - 1; s >= 0; s--) {
      u32 S2 = 4u << (2 * s);
      if (S2 >= 32) {
        for (u32 i = 0; i < g_Sb; i += S2)
          stage4_dif(p + i, S2, S2 >> 2, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF);
      } else if (S2 == 16) {
        stage4_dif_h4(p, g_Sb, g_TF.h1[s], g_TF.h2[s], g_TF.h3[s], g_TF.q1[s], g_TF.q2[s], g_TF.q3[s], g_ciF, g_cisF);
      } else {
        stage4_dif_h1(p, g_Sb, g_ciF, g_cisF);
      }
    }
  }
}

TGT static void ntt_inv(u32 *a, const u32 *Bmul = 0) {
  for (u32 b = 0; b < g_n; b += g_Sb) {
    u32 *p = a + b;
    if (Bmul) {
      const __m256i pv = _mm256_set1_epi32((int)NTP);
      const u32 R2v = (u32)(((u64)1 << 32) % NTP * (((u64)1 << 32) % NTP) % NTP);
      __m256i R2 = _mm256_set1_epi32((int)R2v);
      const u32 *q = Bmul + b;
      for (u32 i = 0; i < g_Sb; i += 8) {
        __m256i x = _mm256_loadu_si256((const __m256i *)(p + i));
        __m256i y = _mm256_loadu_si256((const __m256i *)(q + i));
        _mm256_storeu_si256((__m256i *)(p + i), vmont(vmont(x, y, pv), R2, pv));
      }
    }
    for (int s = 0; s < g_nb; s++) {
      u32 S = 4u << (2 * s);
      if (S >= 32) {
        for (u32 i = 0; i < g_Sb; i += S)
          merge4_dit(p + i, S, S >> 2, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI);
      } else if (S == 16) {
        merge4_dit_h4(p, g_Sb, g_TI.h1[s], g_TI.h2[s], g_TI.h3[s], g_TI.q1[s], g_TI.q2[s], g_TI.q3[s], g_ciI, g_cisI);
      } else {
        merge4_dit_h1(p, g_Sb, g_ciI, g_cisI);
      }
    }
    if (g_odd) merge2_dit(p, g_Sb, g_TI.b2, g_TI.b2s);
  }
  for (int l = g_L - 1; l >= 0; l--) {
    u32 SB = g_n >> (2 * l);
    u32 h = SB >> 2;
    for (u32 b = 0; b < g_n; b += SB)
      merge4_dit(a + b, SB, h, g_TI.l1[l], g_TI.l2[l], g_TI.l3[l], g_TI.m1[l], g_TI.m2[l], g_TI.m3[l], g_ciI, g_cisI);
    _mm_sfence();
  }
}

// ============================================================================
// Mixed-radix radix-3 top level for N = 3 * 2^23 (same prime).
// ============================================================================
static u32 g_M3;                    // 2^23
static u32 *g_t31, *g_t32, *g_u31, *g_u32;   // forward / inverse radix-3 twiddles
static u32 g_rho, g_rhos, g_rho2, g_rho2s;

static void init_r3(u32 N) {
  u32 M = N / 3;
  g_M3 = M;
  u32 w = powmod(31u, (NTP - 1) / N);        // primitive N-th root
  u32 wi = powmod(w, NTP - 2);
  g_rho = powmod(w, M);                      // primitive cube root
  g_rhos = shoup_c(g_rho);
  g_rho2 = mulmod(g_rho, g_rho);
  g_rho2s = shoup_c(g_rho2);
  // place the four tables at the tail of the pools
  u32 *pf = g_tabF + (5 * NTT_MAXN / 2 + 4096) - 2 * M;
  u32 *pi = g_tabI + (5 * NTT_MAXN / 2 + 4096) - 2 * M;
  g_t31 = pf; g_t32 = pf + M;
  g_u31 = pi; g_u32 = pi + M;
  gen_geo_val(pf, M, 1, powmod(w, 2));       // omega_N^{2n}: step 2 per entry? no: geometric in n
  // NOTE: gen_geo_val fills v0 * r^i, so for T1[n]=w^n use r = w ; for T2[n]=w^{2n} use r = w^2
  gen_geo_val(g_t31, M, 1, w);
  gen_geo_val(g_t32, M, 1, mulmod(w, w));
  gen_geo_val(g_u31, M, 1, wi);
  gen_geo_val(g_u32, M, 1, mulmod(wi, wi));
}

// forward radix-3 DIF over the whole array, fused with the coefficient build
TGT static void stage3_dif_src(u32 *a, const u32 *src, u32 len, u32 M) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i rho = _mm256_set1_epi32((int)g_rho), rhos = _mm256_set1_epi32((int)g_rhos);
  u32 j = 0;
  u32 lim = (len > M) ? (len - M) : 0;       // x1 = src[j+M] valid for j <= lim
  for (; j + 8 <= lim + 1; j += 8) {
    __m256i x0 = _mm256_loadu_si256((const __m256i *)(src + j));
    __m256i x1 = _mm256_loadu_si256((const __m256i *)(src + j + M));
    __m256i u0 = vaddm(x0, x1, pv);
    __m256i u1 = vsubm(x0, x1, pv);
    __m256i D = x1;                                   // x2 = 0 -> D = x1 - x2 = x1
    __m256i E1 = vshoup(D, rho, rhos, pv);
    __m256i E2 = vsubm(_mm256_setzero_si256(), vaddm(D, E1, pv), pv);
    __m256i v1 = vaddm(x0, E1, pv);                   // u1 = x0 + rho*x1
    __m256i v2 = vaddm(x0, E2, pv);                   // u2 = x0 + rho^2*x1
    __m256i w1 = _mm256_loadu_si256((const __m256i *)(g_t31 + j));
    __m256i w2 = _mm256_loadu_si256((const __m256i *)(g_t32 + j));
    v1 = vshoup(v1, w1, der_ws(w1), pv);
    v2 = vshoup(v2, w2, der_ws(w2), pv);
    _mm256_storeu_si256((__m256i *)(a + j), vaddm(u0, _mm256_setzero_si256(), pv));
    _mm256_storeu_si256((__m256i *)(a + j + M), v1);
    _mm256_storeu_si256((__m256i *)(a + j + 2 * M), v2);
  }
  for (; j <= lim && j < M; j++) {
    u32 x0 = (j <= len) ? src[j] : 0, x1 = (j + M <= len) ? src[j + M] : 0;
    u32 u0 = (u32)(((u64)x0 + x1) % NTP);
    u32 D = x1;
    u32 E1 = shoup1s(D, g_rho, g_rhos);
    u32 E2 = (u32)((NTP - (u64)((D + E1) % NTP)) % NTP);
    u32 v1 = (u32)(((u64)x0 + E1) % NTP), v2 = (u32)(((u64)x0 + E2) % NTP);
    v1 = shoup1s(v1, g_t31[j], shoup_c(g_t31[j]));
    v2 = shoup1s(v2, g_t32[j], shoup_c(g_t32[j]));
    a[j] = u0; a[j + M] = v1; a[j + 2 * M] = v2;
  }
  for (; j + 8 <= M; j += 8) {
    __m256i x0 = (j + 8 <= len + 1) ? _mm256_loadu_si256((const __m256i *)(src + j)) : _mm256_setzero_si256();
    __m256i u0 = x0;
    __m256i w1 = _mm256_loadu_si256((const __m256i *)(g_t31 + j));
    __m256i w2 = _mm256_loadu_si256((const __m256i *)(g_t32 + j));
    _mm256_storeu_si256((__m256i *)(a + j), u0);
    _mm256_storeu_si256((__m256i *)(a + j + M), vshoup(x0, w1, der_ws(w1), pv));
    _mm256_storeu_si256((__m256i *)(a + j + 2 * M), vshoup(x0, w2, der_ws(w2), pv));
  }
  for (; j < M; j++) {
    u32 x0 = (j <= len) ? src[j] : 0;
    a[j] = x0;
    a[j + M] = shoup1s(x0, g_t31[j], shoup_c(g_t31[j]));
    a[j + 2 * M] = shoup1s(x0, g_t32[j], shoup_c(g_t32[j]));
  }
}

// inverse radix-3 DIT merge over the whole array (3 sub-blocks of size M)
TGT static void merge3_dit(u32 *a, u32 M) {
  const __m256i pv = _mm256_set1_epi32((int)NTP);
  const __m256i rho2 = _mm256_set1_epi32((int)g_rho2), rho2s = _mm256_set1_epi32((int)g_rho2s);
  for (u32 s = 0; s + 8 <= M; s += 8) {
    __m256i t0 = _mm256_loadu_si256((const __m256i *)(a + s));
    __m256i w1 = _mm256_loadu_si256((const __m256i *)(g_u31 + s));
    __m256i w2 = _mm256_loadu_si256((const __m256i *)(g_u32 + s));
    __m256i t1 = vshoup(_mm256_loadu_si256((const __m256i *)(a + s + M)), w1, der_ws(w1), pv);
    __m256i t2 = vshoup(_mm256_loadu_si256((const __m256i *)(a + s + 2 * M)), w2, der_ws(w2), pv);
    __m256i out0 = vaddm(vaddm(t0, t1, pv), t2, pv);
    __m256i D = vsubm(t1, t2, pv);
    __m256i E1 = vshoup(D, rho2, rho2s, pv);
    __m256i out1 = vaddm(vsubm(t0, t2, pv), E1, pv);
    __m256i out2 = vsubm(vsubm(t0, t1, pv), E1, pv);
    _mm256_storeu_si256((__m256i *)(a + s), out0);
    _mm256_storeu_si256((__m256i *)(a + s + M), out1);
    _mm256_storeu_si256((__m256i *)(a + s + 2 * M), out2);
  }
  for (u32 s = (M & ~7u); s < M; s++) {
    u32 t0 = a[s];
    u32 t1 = shoup1s(a[s + M], g_u31[s], shoup_c(g_u31[s]));
    u32 t2 = shoup1s(a[s + 2 * M], g_u32[s], shoup_c(g_u32[s]));
    u32 out0 = (u32)(((u64)t0 + t1 + t2) % NTP);
    u32 D = (u32)(((u64)t1 + NTP - t2) % NTP);
    u32 E1 = shoup1s(D, g_rho2, g_rho2s);
    u32 out1 = (u32)(((u64)t0 + NTP - t2 + E1) % NTP);
    u32 out2 = (u32)(((u64)t0 + NTP - t1 + NTP - E1) % NTP);
    a[s] = out0; a[s + M] = out1; a[s + 2 * M] = out2;
  }
}

// ===================== mixed-radix driver =====================
static void ntt3_fwd(u32 *a, const u32 *src, u32 len) {
  stage3_dif_src(a, src, len, g_M3);
  for (int q = 0; q < 3; q++) ntt_fwd(a + (size_t)q * g_M3, 0, 0);
}
static void ntt3_inv(u32 *a, const u32 *Bmul) {
  for (int q = 0; q < 3; q++) ntt_inv(a + (size_t)q * g_M3, Bmul ? (Bmul + (size_t)q * g_M3) : 0);
  merge3_dit(a, g_M3);
}

// ===================== 1002e7: poly_multiply, n=m=1e7 (mixed radix 3*2^23) ====
static u32 g_dataA[3 * (1u << 23)] __attribute__((aligned(64)));
static u32 g_dataB[3 * (1u << 23)] __attribute__((aligned(64)));
static int g_ready = 0;

__attribute__((target("avx2"))) void poly_multiply(unsigned *a, int n, unsigned *b, int m, unsigned *c) {
  const u32 M = 1u << 23;
  const u32 N = 3u * M;
  if (!g_ready) { ntt_init(M); init_r3(N); g_ready = 1; }
  u32 *A = g_dataA, *B = g_dataB;
  u32 na = (u32)n, nb = (u32)m;
  ntt3_fwd(A, a, na);
  ntt3_fwd(B, b, nb);
  ntt3_inv(A, B);
  {
    const __m256i pv = _mm256_set1_epi32((int)NTP);
    u32 ninv = powmod(N % NTP, NTP - 2);
    __m256i nv = _mm256_set1_epi32((int)ninv), ns = _mm256_set1_epi32((int)shoup_c(ninv));
    u32 lim = na + nb + 1;
    u32 i = 0;
    for (; i + 8 <= lim; i += 8) {
      __m256i x = vshoup(_mm256_loadu_si256((const __m256i *)(A + i)), nv, ns, pv);
      _mm256_storeu_si256((__m256i *)(c + i), x);
    }
    for (; i < lim; i++) c[i] = (u32)((u64)A[i] * ninv % NTP);
  }
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #1612.213 ms460 MB + 584 KBAcceptedScore: 100


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