// 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")
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 {
static constexpr u32x8 mod = {81788929,81788929,81788929,81788929,81788929,81788929,81788929,81788929};
static constexpr u32x8 mod2 = {163577858,163577858,163577858,163577858,163577858,163577858,163577858,163577858};
static constexpr u32x8 n_inv = {81788927,81788927,81788927,81788927,81788927,81788927,81788927,81788927};
static constexpr u32x8 r = {41942988,41942988,41942988,41942988,41942988,41942988,41942988,41942988};
static constexpr u32x8 r2 = {56088131,56088131,56088131,56088131,56088131,56088131,56088131,56088131};
Montgomery_simd() = default;
explicit Montgomery_simd(u32 p) { assert(p == 81788929); }
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;
}
};
namespace fixed_roots {
struct Twiddle{u32 w1,w2,w3,n1,n2,n3;};
struct Table{Twiddle data[65536];};
extern const Table forward asm("poly_roots_forward");
extern const Table inverse asm("poly_roots_inverse");
struct DotTable{u32 data[32768][4];};
extern const DotTable dot asm("poly_roots_dot");
}
asm(R"asm(
.pushsection .rodata
.p2align 6
.globl poly_roots_forward
.type poly_roots_forward,@object
poly_roots_forward:
.set _i,0
.set _w1,41942988
.set _w2,41942988
.set _w3,42958308
.rept 65536
.set _w12,(((_w1*_w2)%81788929)*1557504)%81788929
.long _w1,_w2,_w12,((_w1*81788927)&0xffffffff),((_w2*81788927)&0xffffffff),((_w12*81788927)&0xffffffff)
.if ((_i & 1) == 0)
.set _fac,1977387
.set _fac2,57807995
.elseif ((_i & 2) == 0)
.set _fac,49739338
.set _fac2,1883838
.elseif ((_i & 4) == 0)
.set _fac,76551861
.set _fac2,27152551
.elseif ((_i & 8) == 0)
.set _fac,57685215
.set _fac2,62819432
.elseif ((_i & 16) == 0)
.set _fac,25318722
.set _fac2,22367481
.elseif ((_i & 32) == 0)
.set _fac,22305379
.set _fac2,25489457
.elseif ((_i & 64) == 0)
.set _fac,75160758
.set _fac2,54748607
.elseif ((_i & 128) == 0)
.set _fac,77449485
.set _fac2,18371892
.elseif ((_i & 256) == 0)
.set _fac,49050524
.set _fac2,60074596
.elseif ((_i & 512) == 0)
.set _fac,58847824
.set _fac2,43336831
.elseif ((_i & 1024) == 0)
.set _fac,69356575
.set _fac2,16579980
.elseif ((_i & 2048) == 0)
.set _fac,69052175
.set _fac2,78708963
.elseif ((_i & 4096) == 0)
.set _fac,45043381
.set _fac2,26101542
.elseif ((_i & 8192) == 0)
.set _fac,68811137
.set _fac2,51041304
.elseif ((_i & 16384) == 0)
.set _fac,48691376
.set _fac2,60500196
.elseif ((_i & 32768) == 0)
.set _fac,28111944
.set _fac2,40232015
.elseif ((_i & 65536) == 0)
.set _fac,26577652
.set _fac2,28323882
.endif
.set _w1,(_w1*_fac2)%81788929
.set _w2,(_w2*_fac)%81788929
.set _w3,(_w3*_fac)%81788929
.set _i,_i+1
.endr
.size poly_roots_forward,.-poly_roots_forward
.p2align 6
.globl poly_roots_inverse
.type poly_roots_inverse,@object
poly_roots_inverse:
.set _i,0
.set _w1,41942988
.set _w2,41942988
.set _w3,38830621
.rept 65536
.set _w12,(((_w1*_w2)%81788929)*1557504)%81788929
.long _w1,_w2,_w12,((_w1*81788927)&0xffffffff),((_w2*81788927)&0xffffffff),((_w12*81788927)&0xffffffff)
.if ((_i & 1) == 0)
.set _fac,1883838
.set _fac2,23980934
.elseif ((_i & 2) == 0)
.set _fac,34192649
.set _fac2,1977387
.elseif ((_i & 4) == 0)
.set _fac,65864533
.set _fac2,58967103
.elseif ((_i & 8) == 0)
.set _fac,47472559
.set _fac2,56026958
.elseif ((_i & 16) == 0)
.set _fac,32202660
.set _fac2,74557765
.elseif ((_i & 32) == 0)
.set _fac,46455854
.set _fac2,58488554
.elseif ((_i & 64) == 0)
.set _fac,18299665
.set _fac2,3169619
.elseif ((_i & 128) == 0)
.set _fac,51166265
.set _fac2,20142414
.elseif ((_i & 256) == 0)
.set _fac,46148164
.set _fac2,28119438
.elseif ((_i & 512) == 0)
.set _fac,40005067
.set _fac2,26733415
.elseif ((_i & 1024) == 0)
.set _fac,42538512
.set _fac2,74787290
.elseif ((_i & 2048) == 0)
.set _fac,22507185
.set _fac2,67900511
.elseif ((_i & 4096) == 0)
.set _fac,19881487
.set _fac2,63391377
.elseif ((_i & 8192) == 0)
.set _fac,13191717
.set _fac2,74641937
.elseif ((_i & 16384) == 0)
.set _fac,67317322
.set _fac2,67976842
.elseif ((_i & 32768) == 0)
.set _fac,1064838
.set _fac2,40043517
.elseif ((_i & 65536) == 0)
.set _fac,16759432
.set _fac2,2457972
.endif
.set _w1,(_w1*_fac2)%81788929
.set _w2,(_w2*_fac)%81788929
.set _w3,(_w3*_fac)%81788929
.set _i,_i+1
.endr
.size poly_roots_inverse,.-poly_roots_inverse
.p2align 6
.globl poly_roots_dot
.type poly_roots_dot,@object
poly_roots_dot:
.set _i,0
.set _d0,41942988
.set _d1,42958308
.set _d2,28282409
.set _d3,36011086
.rept 32768
.long _d0,_d1,_d2,_d3
.if ((_i & 1) == 0)
.set _fac,34192649
.elseif ((_i & 2) == 0)
.set _fac,76852948
.elseif ((_i & 4) == 0)
.set _fac,45870503
.elseif ((_i & 8) == 0)
.set _fac,27153147
.elseif ((_i & 16) == 0)
.set _fac,50722843
.elseif ((_i & 32) == 0)
.set _fac,53125215
.elseif ((_i & 64) == 0)
.set _fac,43544278
.elseif ((_i & 128) == 0)
.set _fac,13378268
.elseif ((_i & 256) == 0)
.set _fac,50576854
.elseif ((_i & 512) == 0)
.set _fac,50366248
.elseif ((_i & 1024) == 0)
.set _fac,18491959
.elseif ((_i & 2048) == 0)
.set _fac,52344447
.elseif ((_i & 4096) == 0)
.set _fac,58962465
.elseif ((_i & 8192) == 0)
.set _fac,12062499
.elseif ((_i & 16384) == 0)
.set _fac,19859762
.elseif ((_i & 32768) == 0)
.set _fac,9337084
.endif
.set _d0,(_d0*_fac)%81788929
.set _d1,(_d1*_fac)%81788929
.set _d2,(_d2*_fac)%81788929
.set _d3,(_d3*_fac)%81788929
.set _i,_i+1
.endr
.size poly_roots_dot,.-poly_roots_dot
.popsection
)asm");
class NTT {
mutable int data_origin=0;
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 <int k, bool inverse, bool trivial = false>
__attribute__((always_inline)) inline void transform_const(int i, u32* data, u64x4& wi, const Montgomery_simd& mts) const {
const auto& tw=(inverse?fixed_roots::inverse:fixed_roots::forward).data[unsigned(i)>>(k+2)];
u32x8 w1=(u32x8)_mm256_set1_epi32(tw.w1),w2=(u32x8)_mm256_set1_epi32(tw.w2),w3=(u32x8)_mm256_set1_epi32(tw.w3);
u32x8 n1=(u32x8)_mm256_set1_epi32(tw.n1),n2=(u32x8)_mm256_set1_epi32(tw.n2),n3=(u32x8)_mm256_set1_epi32(tw.n3);
u32x8 root=(u32x8)_mm256_set1_epi32(inverse?38830621:42958308);
u32x8 root_n=(u32x8)_mm256_set1_epi32(inverse?1259306467:3035660828);
if constexpr(trivial){w3=root;n3=root_n;}
for(int j=0;j<(1<<k);j+=8){
u32* p=data+i-data_origin+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){
if constexpr(trivial){
prebut<false,true>(a,c,w1,n1,mts);prebut<false,true>(b,d,w1,n1,mts);
prebut<false,true>(a,b,w2,n2,mts);prebut<false,false>(c,d,w3,n3,mts);
}else{
auto cc=mul_pre(c,w1,n1,mts),bb=mul_pre(b,w2,n2,mts),dd=mul_pre(d,w3,n3,mts);
auto A=a+cc,C=a+mts.mod2-cc,B=bb+dd;
auto D=mul_pre(bb+mts.mod2-dd,root,root_n,mts);
a=A+B;b=A+mts.mod2+mts.mod2-B;c=C+D;d=C+mts.mod2-D;
}
}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,s=c+d,v=a+mts.mod2-b;
auto t=mul_pre(c+mts.mod2-d,root,root_n,mts);
auto sum=u+s;
sum=(u32x8)_mm256_min_epu32((i256)sum,(i256)(sum-mts.mod2-mts.mod2));
a=mts.shrink2(sum);
c=mul_pre(u+mts.mod2+mts.mod2-s,w1,n1,mts);
b=mul_pre(v+t,w2,n2,mts);
d=mul_pre(v+mts.mod2-t,w3,n3,mts);
}
}
store_u32x8(p,a);store_u32x8(p+step,b);store_u32x8(p+2*step,c);store_u32x8(p+3*step,d);
}
}
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 temp;
asm("vpmuludq %3,%2,%1\n\tvpaddq %1,%0,%0"
: "+x"(aux_ans[k][j/4]), "=&x"(temp)
: "x"(ai), "m"(*(const __m256i_u*)(aux_b[k] + n - i + j)));
}
}
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];
if constexpr(L==3){
const auto& tw=fixed_roots::dot.data[i>>3];
for(int q=0;q<4;q++)w_ar[q]=(u32x8)_mm256_set1_epi32(tw[q]);
}else{
w_ar[0]=(u32x8)_mm256_permute4x64_epi64((i256)wi,0x00);
w_ar[1]=(u32x8)_mm256_permute4x64_epi64((i256)wi,0x55);
w_ar[2]=(u32x8)_mm256_permute4x64_epi64((i256)wi,0xaa);
w_ar[3]=(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-data_origin,b+(i+j)*sz-data_origin,a+(i+j)*sz-data_origin,
{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-data_origin,b+(i+j)*sz-data_origin,a+(i+j)*sz-data_origin,
{w_ar[j/2],mts.mod-w_ar[j/2]},mts);
}
if constexpr(L!=3) 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);
}
}
template<int K,bool Inv>
__attribute__((noinline)) void fixed_level(int offset,u32* a,u32* b,u64x4* state) const {
const Montgomery_simd mts;
for(int i=offset;i<offset+128;i+=(1<<(K+2))){
if constexpr(Inv){
if(i==0)transform_const<K,true,true>(i,a,state[K],mts);
else transform_const<K,true>(i,a,state[K],mts);
}else{
auto w=state[K];
if(i==0){transform_const<K,false,true>(i,a,w,mts);transform_const<K,false,true>(i,b,state[K],mts);}
else{transform_const<K,false>(i,a,w,mts);transform_const<K,false>(i,b,state[K],mts);}
}
}
}
template<int LG>
void fixed_fused(int offset,u32* a,u32* b,u64x4* fw,u64x4* iw,u64x4& dot) const {
const Montgomery_simd mts;
if constexpr(LG>7){
constexpr int K=LG-2;
auto w=fw[K];
if(offset==0){transform_const<K,false,true>(offset,a,w,mts);transform_const<K,false,true>(offset,b,fw[K],mts);}
else {transform_const<K,false>(offset,a,w,mts);transform_const<K,false>(offset,b,fw[K],mts);}
for(int j=0;j<4;j++)fixed_fused<K>(offset+(j<<K),a,b,fw,iw,dot);
if(offset==0)transform_const<K,true,true>(offset,a,iw[K],mts);
else transform_const<K,true>(offset,a,iw[K],mts);
}else{
fixed_level<5,false>(offset,a,b,fw);fixed_level<3,false>(offset,a,b,fw);
dot_range<3>(offset,offset+128,a,b,dot);
fixed_level<3,true>(offset,a,b,iw);fixed_level<5,true>(offset,a,b,iw);
}
}
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 init_mapped(const u32* src,int n,u32* dst,u32* last,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(last+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{});
}
template<int Quarter>
void init_quarter(const u32* src,int n,u32* dst,int lg) const {
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)));
if constexpr(Quarter==0)store_u32x8(dst+i,a+b);
if constexpr(Quarter==1)store_u32x8(dst+i,a+mts.mod2-b);
if constexpr(Quarter==2)store_u32x8(dst+i,a+v);
if constexpr(Quarter==3)store_u32x8(dst+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(lg==21)fixed_fused<19>(j<<k,a,b,fw,iw,dot);
else 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 mapped(const u32* A,const u32* B,u32* c) const {
constexpr int lg=21,k=19,L=3,n=1000001,m=1000001;
u32* temp=(u32*)_mm_malloc((2<<19)*4,32);
u32* last=temp;u32* scratch=temp+(1<<19);u32* a=(u32*)(((uintptr_t)c+31)&~uintptr_t(31));
init_mapped(A,n,a,last,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);
for(int j=0;j<4;j++){
if(j==0)init_quarter<0>(B,m,scratch,lg);
if(j==1)init_quarter<1>(B,m,scratch,lg);
if(j==2)init_quarter<2>(B,m,scratch,lg);
if(j==3)init_quarter<3>(B,m,scratch,lg);
data_origin=j<<19;
fixed_fused<19>(data_origin,j==3?last:a+data_origin,scratch,fw,iw,dot);
}
data_origin=0;
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;
alignas(32) u32x8 saved[3]={load_u32x8(a+q-8),load_u32x8(a+2*q-8),load_u32x8(a+3*q-8)};
for(int i=0;i<q;i+=8){
u32x8 x,y,z,t=load_u32x8(last+i);
if(i+8==q){x=saved[0];y=saved[1];z=saved[2];}
else{x=load_u32x8(a+i);y=load_u32x8(a+q+i);z=load_u32x8(a+2*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));
}
_mm_free(temp);
}
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++;
}
auto disjoint=[](const u32* src,const u32* dst){
uintptr_t s=(uintptr_t)src,d=(uintptr_t)dst;
return d+2000001ull*4<=s || s+1000001ull*4<=d;
};
if(n==1000001 && m==1000001 && disjoint(A,c) && disjoint(B,c)){
ntt.mapped(A,B,c);return;
}
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 | 20.654 ms | 11 MB + 672 KB | Accepted | Score: 100 | 显示更多 |