// 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);
}
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 612.213 ms | 460 MB + 584 KB | Accepted | Score: 100 | 显示更多 |