// For p=81788929 and L=3: sum <= 128*p*p; REDC < 4*p.
// Cache-fused adaptation by Codex. Forward transforms, leaf products and
// inverse transforms share each cache-sized pair of blocks.
// Reference: Qwerty1232, https://duck.ac/submission/28087
// Fix B copy length for unequal degrees.
#include <immintrin.h>
#include <math.h>
#include <algorithm>
#include <array>
#include <cassert>
#include <cstdint>
#include <cstring>
#include <vector>
#pragma GCC target("avx2,bmi,tune=skylake")
using u32 = uint32_t;
using u64 = uint64_t;
struct Montgomery {
u32 mod; // mod
u32 mod2; // 2 * mod
u32 n_inv; // n_inv * mod == -1 (mod 2^32)
u32 r; // 2^32 % mod
u32 r2; // (2^32)^2 % mod
Montgomery() = default;
Montgomery(u32 mod) : mod(mod) {
assert(mod % 2 == 1);
assert(mod < (1 << 30));
mod2 = 2 * mod;
n_inv = 1;
for (int i = 0; i < 5; i++) {
n_inv *= 2 + n_inv * mod;
}
r = (u64(1) << 32) % mod;
r2 = u64(r) * r % mod;
}
u32 shrink(u32 val) const {
return std::min(val, val - mod);
}
u32 shrink2(u32 val) const {
return std::min(val, val - mod2);
}
template <bool strict = true>
u32 reduce(u64 val) const {
u32 res = val + u32(val) * n_inv * u64(mod) >> 32;
if /* constexpr */ (strict)
res = shrink(res);
return res;
}
template <bool strict = true>
u32 mul(u32 a, u32 b) const {
return reduce<strict>(u64(a) * b);
}
template <bool input_in_space = false, bool output_in_space = false>
u32 power(u32 b, u32 e) const {
if (!input_in_space)
b = mul<false>(b, r2);
u32 r = output_in_space ? this->r : 1;
for (; e > 0; e >>= 1) {
if (e & 1)
r = mul<false>(r, b);
b = mul<false>(b, b);
}
return shrink(r);
}
};
using i256 = __m256i;
using u32x8 = u32 __attribute__((vector_size(32)));
using u64x4 = u64 __attribute__((vector_size(32)));
u32x8 load_u32x8(const u32* ptr) {
return (u32x8)_mm256_load_si256((const i256*)ptr);
}
void store_u32x8(u32* ptr, u32x8 vec) {
_mm256_store_si256((i256*)ptr, (i256)vec);
}
struct Montgomery_simd {
u32x8 mod; // mod
u32x8 mod2; // 2 * mod
u32x8 n_inv; // n_inv * mod == -1 (mod 2^32)
u32x8 r; // 2^32 % mod
u32x8 r2; // (2^32)^2 % mod
Montgomery_simd() = default;
Montgomery_simd(u32 mod) {
Montgomery mt(mod);
this->mod = (u32x8)_mm256_set1_epi32(mt.mod);
this->mod2 = (u32x8)_mm256_set1_epi32(mt.mod2);
this->n_inv = (u32x8)_mm256_set1_epi32(mt.n_inv);
this->r = (u32x8)_mm256_set1_epi32(mt.r);
this->r2 = (u32x8)_mm256_set1_epi32(mt.r2);
}
u32x8 shrink(u32x8 vec) const {
return (u32x8)_mm256_min_epu32((i256)vec, _mm256_sub_epi32((i256)vec, (i256)mod));
}
template<int Low=0> u32x8 canonical_wide(u32x8 v) const {
for(int shift=5;shift>=Low;shift--){
u32x8 p=mod<<shift;
v=(u32x8)_mm256_min_epu32((i256)v,(i256)(v-p));
}
return v;
}
u32x8 shrink2(u32x8 vec) const {
return (u32x8)_mm256_min_epu32((i256)vec, _mm256_sub_epi32((i256)vec, (i256)mod2));
}
u32x8 shrink_n(u32x8 vec) const {
return (u32x8)_mm256_min_epu32((i256)vec, _mm256_add_epi32((i256)vec, (i256)mod));
}
u32x8 shrink2_n(u32x8 vec) const {
return (u32x8)_mm256_min_epu32((i256)vec, _mm256_add_epi32((i256)vec, (i256)mod2));
}
template <bool strict = true>
u32x8 reduce(u64x4 x0246, u64x4 x1357) const {
u64x4 x0246_ninv = (u64x4)_mm256_mul_epu32((i256)x0246, (i256)n_inv);
u64x4 x1357_ninv = (u64x4)_mm256_mul_epu32((i256)x1357, (i256)n_inv);
u64x4 x0246_res = (u64x4)_mm256_add_epi64((i256)x0246, _mm256_mul_epu32((i256)x0246_ninv, (i256)mod));
u64x4 x1357_res = (u64x4)_mm256_add_epi64((i256)x1357, _mm256_mul_epu32((i256)x1357_ninv, (i256)mod));
u32x8 res = (u32x8)_mm256_or_si256(_mm256_bsrli_epi128((i256)x0246_res, 4), (i256)x1357_res);
if (strict)
res = shrink(res);
return res;
}
template <bool strict = true, bool b_use_only_even = false>
u32x8 mul_u32x8(u32x8 a, u32x8 b) const {
u32x8 a_sh = (u32x8)_mm256_bsrli_epi128((i256)a, 4);
u32x8 b_sh = b_use_only_even ? b : (u32x8)_mm256_bsrli_epi128((i256)b, 4);
u64x4 x0246 = (u64x4)_mm256_mul_epu32((i256)a, (i256)b);
u64x4 x1357 = (u64x4)_mm256_mul_epu32((i256)a_sh, (i256)b_sh);
return reduce<strict>(x0246, x1357);
}
template <bool strict = true>
u64x4 mul_u64x4(u64x4 a, u64x4 b) const {
u64x4 pr = (u64x4)_mm256_mul_epu32((i256)a, (i256)b);
u64x4 pr2 = (u64x4)_mm256_mul_epu32(_mm256_mul_epu32((i256)pr, (i256)n_inv), (i256)mod);
u64x4 res = (u64x4)_mm256_bsrli_epi128(_mm256_add_epi64((i256)pr, (i256)pr2), 4);
if (strict)
res = (u64x4)shrink((u32x8)res);
return res;
}
};
class NTT {
public:
u32 mod, pr_root;
private:
static const /* constexpr */ int LG = 32; // more than enough for u32
Montgomery mt;
Montgomery_simd mts;
u32 w[4], wr[4];
u32 wd[LG], wrd[LG];
u64x4 wt_init, wrt_init;
u64x4 wd_x4[LG], wrd_x4[LG];
u64x4 wl_init;
u64x4 wld_x4[LG];
static u32 find_pr_root(u32 mod, const Montgomery& mt) {
std::vector<u32> factors;
u32 n = mod - 1;
for (u32 i = 2; u64(i) * i <= n; i++) {
if (n % i == 0) {
factors.push_back(i);
do {
n /= i;
} while (n % i == 0);
}
}
if (n > 1) {
factors.push_back(n);
}
for (u32 i = 2; i < mod; i++) {
if (std::all_of(factors.begin(), factors.end(), [&](u32 f) { return mt.power<false, false>(i, (mod - 1) / f) != 1; })) {
return i;
}
}
assert(false && "primitive root not found");
}
public:
NTT() = default;
NTT(u32 mod) : mod(mod), mt(mod), mts(mod) {
const Montgomery mt = this->mt;
const Montgomery_simd mts = this->mts;
pr_root = find_pr_root(mod, mt);
int lg = __builtin_ctz(mod - 1);
assert(lg <= LG);
memset(w, 0, sizeof(w));
memset(wr, 0, sizeof(wr));
memset(wd_x4, 0, sizeof(wd_x4));
memset(wrd_x4, 0, sizeof(wrd_x4));
memset(wld_x4, 0, sizeof(wld_x4));
std::vector<u32> vec(lg + 1), vecr(lg + 1);
vec[lg] = mt.power<false, true>(pr_root, mod - 1 >> lg);
vecr[lg] = mt.power<true, true>(vec[lg], mod - 2);
for (int i = lg - 1; i >= 0; i--) {
vec[i] = mt.mul<true>(vec[i + 1], vec[i + 1]);
vecr[i] = mt.mul<true>(vecr[i + 1], vecr[i + 1]);
}
w[0] = wr[0] = mt.r;
if (lg >= 2) {
w[1] = vec[2], wr[1] = vecr[2];
if (lg >= 3) {
w[2] = vec[3], wr[2] = vecr[3];
w[3] = mt.mul<true>(w[1], w[2]);
wr[3] = mt.mul<true>(wr[1], wr[2]);
}
}
wt_init = (u64x4)_mm256_setr_epi64x(w[0], w[0], w[0], w[1]);
wrt_init = (u64x4)_mm256_setr_epi64x(wr[0], wr[0], wr[0], wr[1]);
wl_init = (u64x4)_mm256_setr_epi64x(w[0], w[1], w[2], w[3]);
u32 prf = mt.r, prf_r = mt.r;
for (int i = 0; i < lg - 2; i++) {
u32 f = mt.mul<true>(prf, vec[i + 3]), fr = mt.mul<true>(prf_r, vecr[i + 3]);
prf = mt.mul<true>(prf, vecr[i + 3]), prf_r = mt.mul<true>(prf_r, vec[i + 3]);
u32 f2 = mt.mul<true>(f, f), f2r = mt.mul<true>(fr, fr);
wd_x4[i] = (u64x4)_mm256_setr_epi64x(f2, f, f2, f);
wrd_x4[i] = (u64x4)_mm256_setr_epi64x(f2r, fr, f2r, fr);
}
prf = mt.r;
for (int i = 0; i < lg - 3; i++) {
u32 f = mt.mul<true>(prf, vec[i + 4]);
prf = mt.mul<true>(prf, vecr[i + 4]);
wld_x4[i] = (u64x4)_mm256_set1_epi64x(f);
}
}
private:
static const /* constexpr */ int L0 = 3;
int get_low_lg(int lg) const {
return lg % 2 == L0 % 2 ? L0 : L0 + 1;
}
// public:
// bool lg_available(int lg) {
// return L0 <= lg && lg <= __builtin_ctz(mod - 1) + get_low_lg(lg);
// }
private:
template <bool transposed, bool trivial = false>
static void butterfly_x2(u32* ptr_a, u32* ptr_b, u32x8 w, const Montgomery_simd& mts) {
u32x8 a = load_u32x8(ptr_a), b = load_u32x8(ptr_b);
u32x8 a2, b2;
if (!transposed) {
a = mts.shrink2(a), b = trivial ? mts.shrink2(b) : mts.mul_u32x8<false, true>(b, w);
a2 = a + b, b2 = a + mts.mod2 - b;
} else {
a2 = mts.shrink2(a + b), b2 = trivial ? mts.shrink2_n(a - b) : mts.mul_u32x8<false, true>(a + mts.mod2 - b, w);
}
store_u32x8(ptr_a, a2), store_u32x8(ptr_b, b2);
}
template <bool transposed, bool trivial = false>
static void butterfly_x4(u32* ptr_a, u32* ptr_b, u32* ptr_c, u32* ptr_d, u32x8 w1, u32x8 w2, u32x8 w3, const Montgomery_simd& mts) {
u32x8 a = load_u32x8(ptr_a), b = load_u32x8(ptr_b), c = load_u32x8(ptr_c), d = load_u32x8(ptr_d);
if (!transposed) {
butterfly_x2<false, trivial>((u32*)&a, (u32*)&c, w1, mts);
butterfly_x2<false, trivial>((u32*)&b, (u32*)&d, w1, mts);
butterfly_x2<false, trivial>((u32*)&a, (u32*)&b, w2, mts);
butterfly_x2<false, false>((u32*)&c, (u32*)&d, w3, mts);
} else {
butterfly_x2<true, trivial>((u32*)&a, (u32*)&b, w2, mts);
butterfly_x2<true, false>((u32*)&c, (u32*)&d, w3, mts);
butterfly_x2<true, trivial>((u32*)&a, (u32*)&c, w1, mts);
butterfly_x2<true, trivial>((u32*)&b, (u32*)&d, w1, mts);
}
store_u32x8(ptr_a, a), store_u32x8(ptr_b, b), store_u32x8(ptr_c, c), store_u32x8(ptr_d, d);
}
static u32x8 mul_pre(u32x8 a,u32x8 w,u32x8 wn,const Montgomery_simd& mts) {
i256 a1=_mm256_srli_epi64((i256)a,32);
i256 m0=_mm256_mul_epu32((i256)a,(i256)wn),m1=_mm256_mul_epu32(a1,(i256)wn);
i256 p0=_mm256_mul_epu32((i256)a,(i256)w),p1=_mm256_mul_epu32(a1,(i256)w);
p0=_mm256_add_epi64(p0,_mm256_mul_epu32(m0,(i256)mts.mod));
p1=_mm256_add_epi64(p1,_mm256_mul_epu32(m1,(i256)mts.mod));
return (u32x8)_mm256_blend_epi32(_mm256_srli_epi64(p0,32),p1,0xaa);
}
template<bool inverse,bool trivial>
static void prebut(u32x8& a,u32x8& b,u32x8 w,u32x8 wn,const Montgomery_simd& mts) {
if constexpr(!inverse){
b=trivial?b:mul_pre(b,w,wn,mts);
auto x=a+b;b=a+mts.mod2-b;a=x;
}else{
auto x=mts.shrink2(a+b);
b=trivial?mts.shrink2_n(a-b):mul_pre(a+mts.mod2-b,w,wn,mts);a=x;
}
}
template <bool inverse, bool trivial = false>
void transform_aux(int k, int i, u32* data, u64x4& wi, const Montgomery_simd& mts) const {
u32x8 w1 = (u32x8)_mm256_shuffle_epi32((i256)wi, 0b00'00'00'00);
u32x8 w2 = (u32x8)_mm256_permute4x64_epi64((i256)wi, 0b01'01'01'01); // only even indices will be used
u32x8 w3 = (u32x8)_mm256_permute4x64_epi64((i256)wi, 0b11'11'11'11); // only even indices will be used
u32x8 n1=(u32x8)_mm256_mul_epu32((i256)w1,(i256)mts.n_inv);
u32x8 n2=(u32x8)_mm256_mul_epu32((i256)w2,(i256)mts.n_inv);
u32x8 n3=(u32x8)_mm256_mul_epu32((i256)w3,(i256)mts.n_inv);
for(int j=0;j<(1<<k);j+=8){
u32* p=data+i+j;int step=1<<k;
auto a=load_u32x8(p),b=load_u32x8(p+step),c=load_u32x8(p+step*2),d=load_u32x8(p+step*3);
if constexpr(!inverse){
prebut<false,trivial>(a,c,w1,n1,mts);prebut<false,trivial>(b,d,w1,n1,mts);
prebut<false,trivial>(a,b,w2,n2,mts);prebut<false,false>(c,d,w3,n3,mts);
}else{
if constexpr(trivial){
prebut<true,true>(a,b,w2,n2,mts);prebut<true,false>(c,d,w3,n3,mts);
prebut<true,true>(a,c,w1,n1,mts);prebut<true,true>(b,d,w1,n1,mts);
}else{
auto u=a+b, v=mul_pre(a+mts.mod2-b,w2,n2,mts);
auto s=c+d, t=mul_pre(c+mts.mod2-d,w3,n3,mts);
auto sum=u+s;
sum=(u32x8)_mm256_min_epu32((i256)sum,(i256)(sum-mts.mod2-mts.mod2));
a=mts.shrink2(sum);b=mts.shrink2(v+t);
c=mul_pre(u+mts.mod2+mts.mod2-s,w1,n1,mts);
d=mul_pre(v+mts.mod2-t,w1,n1,mts);
}
}
store_u32x8(p,a);store_u32x8(p+step,b);store_u32x8(p+2*step,c);store_u32x8(p+3*step,d);
}
wi = mts.mul_u64x4<true>(wi, (inverse ? wrd_x4 : wd_x4)[__builtin_ctz(~i >> k + 2)]);
}
public:
// input in [0, 4 * mod)
// output in [0, 4 * mod)
// data must be 32-byte aligned
void transform_forward(int lg, u32* data) const {
const Montgomery_simd mts = this->mts;
const int L = get_low_lg(lg);
// for (int k = lg - 2; k >= L; k -= 2) {
// u64x4 wi = wt_init;
// transform_aux<false, true>(k, 0, data, wi, mts);
// for (int i = (1 << k + 2); i < (1 << lg); i += (1 << k + 2)) {
// transform_aux<false>(k, i, data, wi, mts);
// }
// }
if (L < lg) {
const int lc = (lg - L) / 2;
u64x4 wi_data[LG / 2];
std::fill(wi_data, wi_data + lc, wt_init);
for (int k = lg - 2; k >= L; k -= 2) {
transform_aux<false, true>(k, 0, data, wi_data[k - L >> 1], mts);
}
for (int i = 1; i < (1 << lc * 2 - 2); i++) {
int s = __builtin_ctz(i) >> 1;
for (int k = s; k >= 0; k--) {
transform_aux<false>(2 * k + L, i * (1 << L + 2), data, wi_data[k], mts);
}
}
}
}
// input in [0, 2 * mod)
// output in [0, mod)
// data must be 32-byte aligned
template <bool mul_by_sc = false>
void transform_inverse(int lg, u32* data, /* as normal number */ u32 sc = u32()) const {
const Montgomery_simd mts = this->mts;
const int L = get_low_lg(lg);
// for (int k = L; k + 2 <= lg; k += 2) {
// u64x4 wi = wrt_init;
// transform_aux<true, true>(k, 0, data, wi, mts);
// for (int i = (1 << k + 2); i < (1 << lg); i += (1 << k + 2)) {
// transform_aux<true>(k, i, data, wi, mts);
// }
// }
if (L < lg) {
const int lc = (lg - L) / 2;
u64x4 wi_data[LG / 2];
std::fill(wi_data, wi_data + lc, wrt_init);
for (int i = 0; i < (1 << lc * 2 - 2); i++) {
int s = __builtin_ctz(~i) >> 1;
if (i + 1 == (1 << 2 * s)) {
s--;
}
for (int k = 0; k <= s; k++) {
transform_aux<true>(2 * k + L, (i + 1 - (1 << 2 * k)) * (1 << L + 2), data, wi_data[k], mts);
}
if (i + 1 == (1 << 2 * (s + 1))) {
s++;
transform_aux<true, true>(2 * s + L, (i + 1 - (1 << 2 * s)) * (1 << L + 2), data, wi_data[s], mts);
}
}
}
const Montgomery mt = this->mt;
u32 f = mt.power<false, true>(mod + 1 >> 1, lg - L);
if /* constexpr */ (mul_by_sc)
f = mt.mul<true>(f, mt.mul<false>(mt.r2, sc));
u32x8 f_x8 = (u32x8)_mm256_set1_epi32(f);
for (int i = 0; i < (1 << lg); i += 8) {
store_u32x8(data + i, mts.mul_u32x8<true, true>(load_u32x8(data + i), f_x8));
}
}
private:
// input in [0, 4 * mod)
// output in [0, 2 * mod)
// multiplies mod (x^2^L - w)
template <int L, int K, bool remove_montgomery_reduction_factor = true>
/* !!! O3 is crucial here !!! */ __attribute__((optimize("O3"))) static void aux_mul_mod_x2L(const u32* a, const u32* b, u32* c, const std::array<u32x8, K>& ar_w, const Montgomery_simd& mts) {
static_assert(L >= 3);
// static_assert(L == L0 || L == L0 + 1);
/* constexpr */ int n = 1 << L;
alignas(64) u32 aux_a[K][n];
alignas(64) u64 aux_b[K][n * 2];
for (int k = 0; k < K; k++) {
for (int i = 0; i < n; i += 8) {
u32x8 ai = load_u32x8(a + n * k + i);
if /* constexpr */ (remove_montgomery_reduction_factor) {
ai = mts.mul_u32x8<true, true>(ai, mts.r2);
} else {
ai = mts.canonical_wide<L==3?2:0>(ai);
}
store_u32x8(aux_a[k] + i, ai);
u32x8 bi = load_u32x8(b + n * k + i);
u32x8 bi_0 = mts.canonical_wide<L==3?2:0>(bi);
u32x8 bi_w = mts.mul_u32x8<true, true>(bi, ar_w[k]);
store_u32x8((u32*)(aux_b[k] + i + 0), (u32x8)_mm256_permutevar8x32_epi32((i256)bi_w, _mm256_setr_epi64x(0, 1, 2, 3)));
store_u32x8((u32*)(aux_b[k] + i + 4), (u32x8)_mm256_permutevar8x32_epi32((i256)bi_w, _mm256_setr_epi64x(4, 5, 6, 7)));
store_u32x8((u32*)(aux_b[k] + n + i + 0), (u32x8)_mm256_permutevar8x32_epi32((i256)bi_0, _mm256_setr_epi64x(0, 1, 2, 3)));
store_u32x8((u32*)(aux_b[k] + n + i + 4), (u32x8)_mm256_permutevar8x32_epi32((i256)bi_0, _mm256_setr_epi64x(4, 5, 6, 7)));
}
}
u64x4 aux_ans[K][n / 4];
memset(aux_ans, 0, sizeof(aux_ans));
for (int i = 0; i < n; i++) {
for (int k = 0; k < K; k++) {
u64x4 ai = (u64x4)_mm256_set1_epi32(aux_a[k][i]);
for (int j = 0; j < n; j += 4) {
u64x4 bi = (u64x4)_mm256_loadu_si256((i256*)(aux_b[k] + n - i + j));
aux_ans[k][j / 4] += /* 64-bit addition */ (u64x4)_mm256_mul_epu32((i256)ai, (i256)bi);
}
}
if (i >= 8 && (i & 7) == 7) {
for (int k = 0; k < K; k++) {
for (int j = 0; j < n; j += 4) {
aux_ans[k][j / 4] = (u64x4)mts.shrink2((u32x8)aux_ans[k][j / 4]);
}
}
}
}
for (int k = 0; k < K; k++) {
for (int i = 0; i < n; i += 8) {
u64x4 c0 = aux_ans[k][i / 4], c1 = aux_ans[k][i / 4 + 1];
u32x8 res = (u32x8)_mm256_permutevar8x32_epi32((i256)mts.reduce<false>(c0, c1), _mm256_setr_epi32(0, 2, 4, 6, 1, 3, 5, 7));
store_u32x8(c + k * n + i, mts.shrink2(res));
}
}
}
template <int L, bool remove_montgomery_reduction_factor = true>
void aux_mul_mod_full(int lg, const u32* a, const u32* b, u32* c) const {
/* constexpr */ int sz = 1 << L;
const Montgomery_simd mts = this->mts;
int cnt = 1 << lg - L;
if (cnt == 1) {
aux_mul_mod_x2L<L, 1, remove_montgomery_reduction_factor>(a, b, c, {mts.r}, mts);
return;
}
if (cnt <= 8) {
for (int i = 0; i < cnt; i += 2) {
u32x8 wi = (u32x8)_mm256_set1_epi32(w[i / 2]);
aux_mul_mod_x2L<L, 2, remove_montgomery_reduction_factor>(a + i * sz, b + i * sz, c + i * sz, {wi, (mts.mod - wi)}, mts);
}
return;
}
u64x4 wi = wl_init;
for (int i = 0; i < cnt; i += 8) {
u32x8 w_ar[4] = {
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0b00'00'00'00),
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0b01'01'01'01),
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0b10'10'10'10),
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0b11'11'11'11),
};
if /* constexpr */ (L == L0) {
for (int j = 0; j < 8; j += 4) {
aux_mul_mod_x2L<L, 4, remove_montgomery_reduction_factor>(a + (i + j) * sz, b + (i + j) * sz, c + (i + j) * sz,
{w_ar[j / 2], mts.mod - w_ar[j / 2], w_ar[j / 2 + 1], mts.mod - w_ar[j / 2 + 1]}, mts);
}
} else {
for (int j = 0; j < 8; j += 2) {
aux_mul_mod_x2L<L, 2, remove_montgomery_reduction_factor>(a + (i + j) * sz, b + (i + j) * sz, c + (i + j) * sz,
{w_ar[j / 2], mts.mod - w_ar[j / 2]}, mts);
}
}
wi = mts.mul_u64x4<true>(wi, wld_x4[__builtin_ctz(~i >> 3)]);
}
}
public:
// input in [0, 4 * mod)
// output in [0, 2 * mod)
template <bool remove_montgomery_reduction_factor = true>
void aux_dot_mod(int lg, const u32* a, const u32* b, u32* c) const {
int L = get_low_lg(lg);
if (L == L0) {
aux_mul_mod_full<L0, remove_montgomery_reduction_factor>(lg, a, b, c);
} else {
aux_mul_mod_full<L0 + 1, remove_montgomery_reduction_factor>(lg, a, b, c);
}
}
template<int L>
void dot_range(int begin, int end, u32* a, u32* b, u64x4& wi) const {
const Montgomery_simd mts = this->mts;
constexpr int sz = 1 << L;
for (int i = begin >> L; i < (end >> L); i += 8) {
u32x8 w_ar[4] = {
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0x00),
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0x55),
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0xaa),
(u32x8)_mm256_permute4x64_epi64((i256)wi, 0xff),
};
if constexpr (L == L0) {
for(int j=0;j<8;j+=4)
aux_mul_mod_x2L<L,4,false>(a+(i+j)*sz,b+(i+j)*sz,a+(i+j)*sz,
{w_ar[j/2],mts.mod-w_ar[j/2],w_ar[j/2+1],mts.mod-w_ar[j/2+1]},mts);
} else {
for(int j=0;j<8;j+=2)
aux_mul_mod_x2L<L,2,false>(a+(i+j)*sz,b+(i+j)*sz,a+(i+j)*sz,
{w_ar[j/2],mts.mod-w_ar[j/2]},mts);
}
wi=mts.mul_u64x4<true>(wi,wld_x4[__builtin_ctz(~i>>3)]);
}
}
template<int L>
void fused(int lg,int offset,u32* a,u32* b,u64x4* fw,u64x4* iw,u64x4& dot) const {
const Montgomery_simd mts = this->mts;
if(lg > 11) {
int k=lg-2;
u64x4 w=fw[k];
if(offset==0) {
transform_aux<false,true>(k,offset,a,w,mts);
transform_aux<false,true>(k,offset,b,fw[k],mts);
} else {
transform_aux<false>(k,offset,a,w,mts);
transform_aux<false>(k,offset,b,fw[k],mts);
}
for(int j=0;j<4;j++)fused<L>(k,offset+(j<<k),a,b,fw,iw,dot);
if(offset==0)transform_aux<true,true>(k,offset,a,iw[k],mts);
else transform_aux<true>(k,offset,a,iw[k],mts);
return;
}
int end=offset+(1<<lg);
for(int k=lg-2;k>=L;k-=2) {
for(int i=offset;i<end;i+=(1<<(k+2))) {
u64x4 w=fw[k];
if(i==0) {
transform_aux<false,true>(k,i,a,w,mts);
transform_aux<false,true>(k,i,b,fw[k],mts);
}else {
transform_aux<false>(k,i,a,w,mts);
transform_aux<false>(k,i,b,fw[k],mts);
}
}
}
dot_range<L>(offset,end,a,b,dot);
for(int k=L;k<=lg-2;k+=2)
for(int i=offset;i<end;i+=(1<<(k+2))) {
if(i==0)transform_aux<true,true>(k,i,a,iw[k],mts);
else transform_aux<true>(k,i,a,iw[k],mts);
}
}
void init_top(const u32* src,int n,u32* dst,int lg) const {
auto stream=[](u32* p,u32x8 x){_mm256_stream_si256((i256*)p,(i256)x);};
const int q=1<<(lg-2);
alignas(32) u32 table[16];
u32 root=mt.mul(w[1],1);
for(int i=0;i<16;i++)table[i]=u64(root)*i%mod;
u32x8 t0=load_u32x8(table),t1=load_u32x8(table+8);
auto run=[&](int i,u32x8 a,u32x8 b){
u32x8 v=(u32x8)_mm256_blendv_epi8(
_mm256_permutevar8x32_epi32((i256)t0,(i256)b),
_mm256_permutevar8x32_epi32((i256)t1,(i256)b),
_mm256_cmpgt_epi32((i256)b,_mm256_set1_epi32(7)));
stream(dst+i,a+b);stream(dst+q+i,a+mts.mod2-b);
stream(dst+2*q+i,a+v);stream(dst+3*q+i,a+mts.mod2-v);
};
int i=0;
for(;i+8<=n-q;i+=8)run(i,(u32x8)_mm256_loadu_si256((const i256*)(src+i)),(u32x8)_mm256_loadu_si256((const i256*)(src+q+i)));
if(i<n-q){alignas(32)u32 tail[8]={};memcpy(tail,src+q+i,(n-q-i)*4);run(i,(u32x8)_mm256_loadu_si256((const i256*)(src+i)),load_u32x8(tail));i+=8;}
for(;i<q;i+=8)run(i,(u32x8)_mm256_loadu_si256((const i256*)(src+i)),u32x8{});
}
void convolve_inputs(const u32* A,int n,const u32* B,int m,u32* c,int lg,u32* a,u32* b) const {
init_top(A,n,a,lg);init_top(B,m,b,lg);_mm_sfence();
u64x4 fw[LG],iw[LG],dot=wl_init;
std::fill(fw,fw+LG,wt_init);std::fill(iw,iw+LG,wrt_init);
int k=lg-2,L=get_low_lg(lg);
for(int j=0;j<4;j++) {
if(L==L0)fused<L0>(k,j<<k,a,b,fw,iw,dot);
else fused<L0+1>(k,j<<k,a,b,fw,iw,dot);
}
u32 f=mt.power<false,true>((mod+1)>>1,lg-L);
f=mt.mul<true>(f,mt.mul<false>(mt.r2,mt.r));
u32x8 fx=(u32x8)_mm256_set1_epi32(f);
u32x8 w3=(u32x8)_mm256_set1_epi32(wr[1]);
u32x8 froot=(u32x8)_mm256_set1_epi32(mt.mul(f,wr[1]));
int q=1<<k,sz=n+m-1;
for(int i=0;i<q;i+=8){
u32x8 x=load_u32x8(a+i),y=load_u32x8(a+q+i),z=load_u32x8(a+2*q+i),t=load_u32x8(a+3*q+i);
auto u=mts.mul_u32x8<true,true>(x+y,fx);
auto v=mts.mul_u32x8<true,true>(x+mts.mod2-y,fx);
auto s=mts.mul_u32x8<true,true>(z+t,fx);
auto r=mts.mul_u32x8<true,true>(z+mts.mod2-t,froot);
x=mts.shrink(u+s);y=mts.shrink(v+r);
z=mts.shrink(u+mts.mod-s);t=mts.shrink(v+mts.mod-r);
_mm256_storeu_si256((i256*)(c+i),(i256)x);_mm256_storeu_si256((i256*)(c+q+i),(i256)y);
_mm256_storeu_si256((i256*)(c+2*q+i),(i256)z);
if(i+3*q+8<=sz)_mm256_storeu_si256((i256*)(c+3*q+i),(i256)t);
else if(i+3*q<sz)memcpy(c+3*q+i,&t,4*(sz-3*q-i));
}
}
void convolve_cyclic(int lg,u32* a,u32* b) const {
if(lg<7){
transform_forward(lg,a);transform_forward(lg,b);
aux_dot_mod<false>(lg,a,b,a);transform_inverse<true>(lg,a,mt.r);return;
}
u64x4 fw[LG],iw[LG],dot=wl_init;
std::fill(fw,fw+LG,wt_init);std::fill(iw,iw+LG,wrt_init);
int L=get_low_lg(lg);
if(L==L0)fused<L0>(lg,0,a,b,fw,iw,dot);
else fused<L0+1>(lg,0,a,b,fw,iw,dot);
u32 f=mt.power<false,true>((mod+1)>>1,lg-L);
f=mt.mul<true>(f,mt.mul<false>(mt.r2,mt.r));
u32x8 fx=(u32x8)_mm256_set1_epi32(f);
for(int i=0;i<(1<<lg);i+=8)store_u32x8(a+i,mts.mul_u32x8<true,true>(load_u32x8(a+i),fx));
}
};
void poly_multiply(unsigned* A, int n, unsigned* B, int m, unsigned* c) {
n++, m++;
u32 mod = 81'788'929;
NTT ntt(mod);
// int lg = std::max<int>(3, std::bit_width<size_t>((n - 1) + (m - 1)));
int lg = 3;
while ((1 << lg) < (n + m - 1)) {
lg++;
}
u32* a = (u32*)_mm_malloc(4 << lg, 32);
u32* b = (u32*)_mm_malloc(4 << lg, 32);
if(lg>=9 && n>=(1<<(lg-2)) && m>=(1<<(lg-2)) && n<=(1<<(lg-1)) && m<=(1<<(lg-1)) && n+m-1>=(3<<(lg-2))) {
ntt.convolve_inputs(A,n,B,m,c,lg,a,b);
_mm_free(a);_mm_free(b);return;
}
std::copy(A, A + n, a);
std::copy(B, B + m, b);
std::fill(a + n, a + (1 << lg), 0);
std::fill(b + m, b + (1 << lg), 0);
ntt.convolve_cyclic(lg, a, b);
std::copy(a, a + n + m - 1, c);
_mm_free(a), _mm_free(b);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 22.69 ms | 23 MB + 676 KB | Accepted | Score: 100 | 显示更多 |