// 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 uint16_t u16;
typedef uint64_t u64;
#define NTP 2013265921u
#define TGT __attribute__((target("avx2")))
#ifndef NTT_MAXN
#define NTT_MAXN (1u << 22)
#endif
static u32 g_n, g_Sb;
static u32 g_P = 2013265921u, g_Ml = 572662301u, g_Msh = 1;
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 % g_P); }
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 * g_P);
return r >= g_P ? r - g_P : r;
}
static inline u32 shoup_c(u32 w) { return (u32)(((u64)w << 32) / g_P); }
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)g_Ml);
__m256i t = _mm256_add_epi32(vmulhi(w, mlo), _mm256_sll_epi32(w, _mm_cvtsi32_si128(g_Msh)));
return t;
}
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)g_P);
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) % g_P);
u32 u2 = (u32)(((u64)x0 + g_P - x1) % g_P);
u32 E = mulmod(x1, ci);
u32 u1 = (u32)(((u64)x0 + E) % g_P), u3 = (u32)(((u64)x0 + g_P - E) % g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
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)g_P);
u32 seq[8], cur = v0 % g_P;
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);
_mm256_storeu_si256((__m256i *)(sc + i), der_ws(v));
v = vshoup(v, r8v, rs8v, pv);
}
if (i < cnt) {
u32 x = v0 % g_P;
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)g_P);
u32 seq[8], cur = v0 % g_P;
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 % g_P; 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 u32 g_initP = 2013265921u;
static void ntt_init(u32 n) {
g_P = g_initP; g_pinv = 0;
{ u32 iv = 1; for (int q = 0; q < 5; q++) iv *= 2u - g_P * iv; g_pinv = 0u - iv;
u64 M = (u64)(~0ull) / g_P; u32 mh = (u32)(M >> 32); g_Msh = 0; while ((1u << g_Msh) < mh) g_Msh++; g_Ml = (u32)(M & 0xffffffffu); }
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, (g_P - 1) / n);
u32 wi = powmod(w, g_P - 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 && g_L >= 1 && 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)g_P);
const u32 R2v = (u32)(((u64)1 << 32) % g_P * (((u64)1 << 32) % g_P) % g_P);
__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();
}
}
// ==================== 1004e7: decimal multiply, n = m = 1e7 digits ====================
struct DI {
unsigned long abi;
const char *s; unsigned long sn;
char *o; unsigned long ol; unsigned long os;
char *e; unsigned long el; unsigned long es;
const char *IB; unsigned long IBl;
char *OB; unsigned long OBl;
unsigned long tsc;
} __attribute__((packed));
static const u32 P1 = 998244353u, P2 = 2013265921u;
#define MAXLIMB 1700000
static u32 LB1[MAXLIMB], LB2[MAXLIMB];
static u32 RA[MAXLIMB * 2 + 8], RB[MAXLIMB * 2 + 8];
static u64 COEF[MAXLIMB * 2 + 8];
static u32 PA_[NTT_MAXN], PB_[NTT_MAXN] __attribute__((aligned(64)));
static const char DIG2[201] =
"00010203040506070809101112131415161718192021222324252627282930313233343536373839"
"40414243444546474849505152535455565758596061626364656667686970717273747576777879"
"8081828384858687888990919293949596979899";
// little-endian limbs: out[0] = least significant group of up to 6 digits
// 8 ASCII digits -> numeric value (SWAR, 4 multiplies)
static inline u64 swar8(const char *s) {
u64 x;
memcpy(&x, s, 8);
x -= 0x3030303030303030ull;
x = (x * 10 + (x >> 8)) & 0x00FF00FF00FF00FFull;
x = (x * 100 + (x >> 16)) & 0x0000FFFF0000FFFFull;
return (x * 10000 + (x >> 32)) & 0xFFFFFFFFull;
}
static u32 parse_limbs(const char *s, long len, u32 *out) {
u32 k = 0;
long i = len;
// 24 digits at a time: 3 SWAR groups -> 4 six-digit limbs
while (i >= 24) {
i -= 24;
const char *q = s + i;
u64 g0 = swar8(q), g1 = swar8(q + 8), g2 = swar8(q + 16);
out[k] = (u32)(g2 % 1000000ull);
out[k + 1] = (u32)((g1 % 10000ull) * 100ull + (g2 / 1000000ull));
out[k + 2] = (u32)((g0 % 100ull) * 10000ull + (g1 / 10000ull));
out[k + 3] = (u32)(g0 / 100ull);
k += 4;
}
while (i >= 6) {
i -= 6;
u32 v = (u32)(s[i] - '0');
v = v * 10 + (u32)(s[i+1] - '0');
v = v * 10 + (u32)(s[i+2] - '0');
v = v * 10 + (u32)(s[i+3] - '0');
v = v * 10 + (u32)(s[i+4] - '0');
v = v * 10 + (u32)(s[i+5] - '0');
out[k++] = v;
}
if (i > 0) {
u32 v = 0;
for (long j = 0; j < i; j++) v = v * 10 + (u32)(s[j] - '0');
out[k++] = v;
}
return k;
}
__attribute__((target("avx2"))) static void mul_run_job(DI *d) {
const char *s = d->s;
long sn = (long)d->sn;
long l1 = 0;
while (l1 < sn && s[l1] >= '0' && s[l1] <= '9') l1++;
long b0 = l1;
while (b0 < sn && (s[b0] < '0' || s[b0] > '9')) b0++;
long l2 = b0;
while (l2 < sn && s[l2] >= '0' && s[l2] <= '9') l2++;
l2 -= b0;
if (l1 <= 0 || l2 <= 0) return;
u32 na = parse_limbs(s, l1, LB1);
u32 nb = parse_limbs(s + b0, l2, LB2);
u32 need = na + nb;
u32 n = 1; while (n < need) n <<= 1;
for (int pr = 0; pr < 2; pr++) {
g_initP = pr ? P2 : P1;
ntt_init(n);
u32 *A = PA_, *B = PB_;
ntt_fwd(A, LB1, na);
ntt_fwd(B, LB2, nb);
ntt_inv(A, B);
u32 ninv = powmod(n % g_P, g_P - 2);
u32 sh = shoup_c(ninv);
u32 *dst = pr ? RB : RA;
u32 i = 0;
{
const __m256i pv = _mm256_set1_epi32((int)g_P);
__m256i nv = _mm256_set1_epi32((int)ninv), ns = _mm256_set1_epi32((int)sh);
for (; i + 8 <= need; i += 8) {
__m256i x = vshoup(_mm256_loadu_si256((const __m256i *)(A + i)), nv, ns, pv);
_mm256_storeu_si256((__m256i *)(dst + i), x);
}
for (; i < need; i++) dst[i] = shoup1s(A[i], ninv, sh);
}
}
// CRT (Montgomery) fused with the base-1e6 carry (r1 < P1 < P2, so r1 mod P2 = r1)
u32 invp1 = powmod(P1 % P2, P2 - 2);
const u32 Kinv = (u32)((u64)invp1 % P2 * (((u64)1 << 32) % P2) % P2); // invp1 * 2^32 mod P2
const u32 p2inv = g_pinv, p2v = g_P; // P2's Montgomery constants
u64 carry = 0;
u32 *res = RB;
u32 top = 0;
{
u64 tmp[512];
for (u32 base = 0; base < need; base += 512) {
u32 m = (need - base < 512) ? (need - base) : 512;
for (u32 j = 0; j < m; j++) {
u32 i = base + j;
u32 r1 = RA[i], r2 = RB[i];
u32 dd = r2 - r1; if (dd >= P2) dd += P2; // (r2 - r1) mod P2
u32 mm = (u32)((u64)dd * Kinv) * p2inv;
u64 u = (((u64)dd * Kinv) + (u64)mm * p2v) >> 32; // dd * invp1 mod P2
if (u >= p2v) u -= p2v;
tmp[j] = r1 + (u64)P1 * (u32)u;
}
for (u32 j = 0; j < m; j++) {
u64 v = tmp[j] + carry;
res[base + j] = (u32)(v % 1000000u);
carry = v / 1000000u;
if (res[base + j]) top = base + j + 1;
}
}
}
while (carry) { res[top++] = (u32)(carry % 1000000u); carry /= 1000000u; }
if (top == 0) top = 1;
// output
char *o = d->o;
char *olim = d->o + d->ol - 64;
{
u32 v = res[top - 1];
char tmp[12]; int k = 0;
if (v == 0) tmp[k++] = '0';
while (v) { tmp[k++] = (char)('0' + v % 10); v /= 10; }
while (k) { if (o > olim) { d->os = (unsigned long)(o - d->o); return; } *o++ = tmp[--k]; }
}
u32 d3[1000];
for (u32 j = 0; j < 1000; j++) {
u32 hi = j / 100, mid = (j / 10) % 10, lo = j % 10;
d3[j] = (u32)('0' + hi) | ((u32)('0' + mid) << 8) | ((u32)('0' + lo) << 16);
}
for (int i = (int)top - 2; i >= 0; i--) {
u32 v = res[i];
if (o > olim) { d->os = (unsigned long)(o - d->o); return; }
u64 w = (u64)d3[v / 1000] | ((u64)d3[v % 1000] << 24);
memcpy(o, &w, 8); o += 6;
}
if (o <= olim) *o++ = '\n';
d->os = (unsigned long)(o - d->o);
}
#ifndef NO_LOCAL_TEST
extern "C" void __libc_start_main(void *m, int argc, char **argv) {
(void)m;
unsigned long *p = (unsigned long *)(argv + argc + 1);
while (*p) p++;
p++;
DI *d = 0;
for (int i = 0; i < 32 && p[0]; i++, p += 2)
if (p[0] == 0x6b637564UL) { d = (DI *)p[1]; break; }
if (d) mul_run_job(d);
__asm__ volatile("syscall" ::"a"(60), "D"(0) : "rcx", "r11", "memory");
for (;;);
}
int main() { return 0; }
#endif
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 213.829 ms | 121 MB + 780 KB | Accepted | Score: 100 | 显示更多 |