// Duck.ac 1002: exact convolution for coefficients 0..9, degrees <= 1,000,000.
// AVX2 / GCC 9.3, adapted from Qwerty1232: https://duck.ac/submission/28087
// Cleaned from https://duck.ac/submission/48181 (20.505208 ms).
// p = 39 * 2^21 + 1 > 81 * 1,000,001, so one modulus gives exact integers.
// The fast path fuses radix-4 NTT stages in 512-element blocks and reuses c.
#include <immintrin.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 (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 MontgomeryAVX2 {
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};
MontgomeryAVX2() = default;
explicit MontgomeryAVX2(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 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;
}
};
// Assemble read-only roots at build time. C++ constexpr expansion exceeds the
// judge compiler's memory limit. n1/n2/n3 are w1/w2/w3 * n_inv modulo 2^32.
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");
} // namespace fixed_roots
asm(R"asm(
.pushsection .rodata
// Select by the number of trailing one bits in the table index.
.macro next_factor dest, mask, value, rest:vararg
.if ((_i & \mask) == 0)
.set \dest,\value
.else
next_factor \dest,(\mask*2),\rest
.endif
.endm
.p2align 6
.globl poly_roots_forward
.type poly_roots_forward,@object
poly_roots_forward:
.set _i,0
.set _w1,41942988
.set _w2,41942988
.rept 65536
.set _w12,(((_w1*_w2)%81788929)*1557504)%81788929
.long _w1,_w2,_w12,((_w1*81788927)&0xffffffff),((_w2*81788927)&0xffffffff),((_w12*81788927)&0xffffffff)
next_factor _fac, 1, 1977387,49739338,76551861,57685215,25318722,22305379,75160758,77449485,49050524,58847824,69356575,69052175,45043381,68811137,48691376,28111944,26577652
next_factor _fac2, 1, 57807995,1883838,27152551,62819432,22367481,25489457,54748607,18371892,60074596,43336831,16579980,78708963,26101542,51041304,60500196,40232015,28323882
.set _w1,(_w1*_fac2)%81788929
.set _w2,(_w2*_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
.rept 65536
.set _w12,(((_w1*_w2)%81788929)*1557504)%81788929
.long _w1,_w2,_w12,((_w1*81788927)&0xffffffff),((_w2*81788927)&0xffffffff),((_w12*81788927)&0xffffffff)
next_factor _fac, 1, 1883838,34192649,65864533,47472559,32202660,46455854,18299665,51166265,46148164,40005067,42538512,22507185,19881487,13191717,67317322,1064838,16759432
next_factor _fac2, 1, 23980934,1977387,58967103,56026958,74557765,58488554,3169619,20142414,28119438,26733415,74787290,67900511,63391377,74641937,67976842,40043517,2457972
.set _w1,(_w1*_fac2)%81788929
.set _w2,(_w2*_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
next_factor _fac, 1, 34192649,76852948,45870503,27153147,50722843,53125215,43544278,13378268,50576854,50366248,18491959,52344447,58962465,12062499,19859762,9337084
.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
.purgem next_factor
.popsection
)asm");
class NTT {
// Global coefficient offset of the quarter currently in local scratch.
mutable int data_origin = 0;
public:
u32 mod;
private:
static const int LG = 32; // more than enough for u32
Montgomery mt;
MontgomeryAVX2 mts;
u32 w[4], wr[4];
u64x4 wt_init, wrt_init;
u64x4 wd_x4[LG], wrd_x4[LG];
u64x4 wl_init;
u64x4 wld_x4[LG];
public:
NTT(u32 mod) : mod(mod), mt(mod), mts(mod) {
const Montgomery mt = this->mt;
constexpr u32 pr_root = 7; // Primitive root for the fixed modulus 81,788,929.
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 int L0 = 3;
int leaf_log2(int lg) const { return lg % 2 == L0 % 2 ? L0 : L0 + 1; }
// Precomputed w*n_inv lets the product and reduction start independently.
static u32x8 mul_pre(u32x8 a, u32x8 w, u32x8 wn, const MontgomeryAVX2 &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 butterfly_pair(u32x8 &a, u32x8 &b, u32x8 w, u32x8 wn, const MontgomeryAVX2 &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;
}
}
// Fixed-size path: table-indexed roots, no running twiddle dependency.
// For the official input, forward residues stay below 51p < 2^32.
template <int k, bool inverse, bool trivial = false>
__attribute__((always_inline)) inline void transform_fixed(int i, u32 *data, const MontgomeryAVX2 &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;
}
if constexpr (inverse && !trivial && k >= 5) {
// 2x-unrolled nontrivial inverse: two independent butterflies in flight.
const int step = 1 << k;
for (int j = 0; j < step; j += 16) {
u32 *p = data + i - data_origin + j;
u32 *q = p + 8;
auto a = load_u32x8(p), b = load_u32x8(p + step), c = load_u32x8(p + step * 2),
d = load_u32x8(p + step * 3);
auto e = load_u32x8(q), f = load_u32x8(q + step), g = load_u32x8(q + step * 2),
h = load_u32x8(q + step * 3);
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 u2 = e + f, s2 = g + h, v2 = e + mts.mod2 - f;
auto t2 = mul_pre(g + mts.mod2 - h, 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);
auto sum2 = u2 + s2;
sum2 = (u32x8)_mm256_min_epu32((i256)sum2, (i256)(sum2 - mts.mod2 - mts.mod2));
e = mts.shrink2(sum2);
c = mul_pre(u + mts.mod2 + mts.mod2 - s, w1, n1, mts);
g = mul_pre(u2 + mts.mod2 + mts.mod2 - s2, w1, n1, mts);
b = mul_pre(v + t, w2, n2, mts);
f = mul_pre(v2 + t2, w2, n2, mts);
d = mul_pre(v + mts.mod2 - t, w3, n3, mts);
h = mul_pre(v2 + mts.mod2 - t2, w3, n3, mts);
store_u32x8(p, a);
store_u32x8(p + step, b);
store_u32x8(p + 2 * step, c);
store_u32x8(p + 3 * step, d);
store_u32x8(q, e);
store_u32x8(q + step, f);
store_u32x8(q + 2 * step, g);
store_u32x8(q + 3 * step, h);
}
return;
}
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) {
butterfly_pair<false, true>(a, c, w1, n1, mts);
butterfly_pair<false, true>(b, d, w1, n1, mts);
butterfly_pair<false, true>(a, b, w2, n2, mts);
butterfly_pair<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) {
butterfly_pair<true, true>(a, b, w2, n2, mts);
butterfly_pair<true, false>(c, d, w3, n3, mts);
butterfly_pair<true, true>(a, c, w1, n1, mts);
butterfly_pair<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);
}
}
// Paired forward transform: two arrays share one twiddle broadcast set.
template <int k, bool trivial = false>
__attribute__((always_inline)) inline void transform_fixed_pair(int i, u32 *data, u32 *data2,
const MontgomeryAVX2 &mts) const {
const auto &tw = 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(42958308);
u32x8 root_n = (u32x8)_mm256_set1_epi32(3035660828);
if constexpr (trivial) {
w3 = root;
n3 = root_n;
}
const int step = 1 << k;
const u32 *base = data + i - data_origin;
const u32 *base2 = data2 + i - data_origin;
for (int j = 0; j < step; j += 8) {
u32 *p = const_cast<u32 *>(base) + j;
u32 *q = const_cast<u32 *>(base2) + j;
auto a = load_u32x8(p), b = load_u32x8(p + step), c = load_u32x8(p + step * 2),
d = load_u32x8(p + step * 3);
auto e = load_u32x8(q), f = load_u32x8(q + step), g = load_u32x8(q + step * 2),
h = load_u32x8(q + step * 3);
u32x8 A1, B1, C1, D1, A2, B2, C2, D2;
if constexpr (trivial) {
butterfly_pair<false, true>(a, c, w1, n1, mts);
butterfly_pair<false, true>(b, d, w1, n1, mts);
butterfly_pair<false, true>(e, g, w1, n1, mts);
butterfly_pair<false, true>(f, h, w1, n1, mts);
butterfly_pair<false, true>(a, b, w2, n2, mts);
butterfly_pair<false, false>(c, d, w3, n3, mts);
butterfly_pair<false, true>(e, f, w2, n2, mts);
butterfly_pair<false, false>(g, h, 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 gg = mul_pre(g, w1, n1, mts), ff = mul_pre(f, w2, n2, mts), hh = mul_pre(h, w3, n3, mts);
A1 = a + cc, C1 = a + mts.mod2 - cc, B1 = bb + dd;
D1 = mul_pre(bb + mts.mod2 - dd, root, root_n, mts);
A2 = e + gg, C2 = e + mts.mod2 - gg, B2 = ff + hh;
D2 = mul_pre(ff + mts.mod2 - hh, root, root_n, mts);
a = A1 + B1;
b = A1 + mts.mod2 + mts.mod2 - B1;
c = C1 + D1;
d = C1 + mts.mod2 - D1;
e = A2 + B2;
f = A2 + mts.mod2 + mts.mod2 - B2;
g = C2 + D2;
h = C2 + mts.mod2 - D2;
}
store_u32x8(p, a);
store_u32x8(p + step, b);
store_u32x8(p + 2 * step, c);
store_u32x8(p + 3 * step, d);
store_u32x8(q, e);
store_u32x8(q + step, f);
store_u32x8(q + 2 * step, g);
store_u32x8(q + 3 * step, h);
}
}
template <bool inverse, bool trivial = false>
void transform_stage(int k, int i, u32 *data, u64x4 &wi, const MontgomeryAVX2 &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) {
butterfly_pair<false, trivial>(a, c, w1, n1, mts);
butterfly_pair<false, trivial>(b, d, w1, n1, mts);
butterfly_pair<false, trivial>(a, b, w2, n2, mts);
butterfly_pair<false, false>(c, d, w3, n3, mts);
} else {
if constexpr (trivial) {
butterfly_pair<true, true>(a, b, w2, n2, mts);
butterfly_pair<true, false>(c, d, w3, n3, mts);
butterfly_pair<true, true>(a, c, w1, n1, mts);
butterfly_pair<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:
// Generic forward transform; data is 32-byte aligned.
// Lazy residues grow through the stages; the product kernel normalizes them.
void transform_forward(int lg, u32 *data) const {
const MontgomeryAVX2 mts = this->mts;
const int L = leaf_log2(lg);
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_stage<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_stage<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 MontgomeryAVX2 mts = this->mts;
const int L = leaf_log2(lg);
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_stage<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_stage<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 (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:
// Multiply modulo x^(2^L)-w. Normalize lazy inputs; output is below 2p.
// At L=3 each sum <= 128*p*p, so Montgomery reduction gives <4p.
// O3 and the memory-operand multiply/accumulate are performance-critical.
template <int L, int K, bool remove_montgomery_reduction_factor = true>
__attribute__((optimize("O3"))) static void
multiply_leaf(const u32 *a, const u32 *b, u32 *c, const std::array<u32x8, K> &ar_w, const MontgomeryAVX2 &mts) {
static_assert(L >= 3);
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 (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 + 2 <= n; i += 2) {
for (int k = 0; k < K; k++) {
u64x4 ai = (u64x4)_mm256_set1_epi32(aux_a[k][i]);
u64x4 ai1 = (u64x4)_mm256_set1_epi32(aux_a[k][i + 1]);
for (int j = 0; j < n; j += 4) {
u64x4 t0, t1;
asm("vpmuludq %3,%2,%1\n\tvpaddq %1,%0,%0"
: "+x"(aux_ans[k][j / 4]), "=&x"(t0)
: "x"(ai), "m"(*(const __m256i_u *)(aux_b[k] + n - i + j)));
asm("vpmuludq %3,%2,%1\n\tvpaddq %1,%0,%0"
: "+x"(aux_ans[k][j / 4]), "=&x"(t1)
: "x"(ai1), "m"(*(const __m256i_u *)(aux_b[k] + n - i - 1 + j)));
}
}
if (((i + 1) & 7) == 7 && i + 1 >= 15) {
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]);
}
}
}
}
// n is even (L >= 3): the unrolled loop above consumed rows in pairs
// and advanced i past the final pair; nothing remains.
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 multiply_leaves(int lg, const u32 *a, const u32 *b, u32 *c) const {
constexpr int sz = 1 << L;
const MontgomeryAVX2 mts = this->mts;
int cnt = 1 << lg - L;
if (cnt == 1) {
multiply_leaf<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]);
multiply_leaf<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 (L == L0) {
for (int j = 0; j < 8; j += 4) {
multiply_leaf<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) {
multiply_leaf<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:
// Leaf products: normalize lazy inputs and return residues below 2p.
template <bool remove_montgomery_reduction_factor = true>
void multiply_all_leaves(int lg, const u32 *a, const u32 *b, u32 *c) const {
int L = leaf_log2(lg);
if (L == L0) {
multiply_leaves<L0, remove_montgomery_reduction_factor>(lg, a, b, c);
} else {
multiply_leaves<L0 + 1, remove_montgomery_reduction_factor>(lg, a, b, c);
}
}
template <int L> void multiply_range(int begin, int end, u32 *a, u32 *b, u64x4 &wi) const {
const MontgomeryAVX2 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)
multiply_leaf<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)
multiply_leaf<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 convolve_block(int lg, int offset, u32 *a, u32 *b, u64x4 *fw, u64x4 *iw, u64x4 &dot) const {
const MontgomeryAVX2 mts = this->mts;
if (lg > 11) {
int k = lg - 2;
u64x4 w = fw[k];
if (offset == 0) {
transform_stage<false, true>(k, offset, a, w, mts);
transform_stage<false, true>(k, offset, b, fw[k], mts);
} else {
transform_stage<false>(k, offset, a, w, mts);
transform_stage<false>(k, offset, b, fw[k], mts);
}
for (int j = 0; j < 4; j++) convolve_block<L>(k, offset + (j << k), a, b, fw, iw, dot);
if (offset == 0)
transform_stage<true, true>(k, offset, a, iw[k], mts);
else
transform_stage<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_stage<false, true>(k, i, a, w, mts);
transform_stage<false, true>(k, i, b, fw[k], mts);
} else {
transform_stage<false>(k, i, a, w, mts);
transform_stage<false>(k, i, b, fw[k], mts);
}
}
}
multiply_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_stage<true, true>(k, i, a, iw[k], mts);
else
transform_stage<true>(k, i, a, iw[k], mts);
}
}
template <int K, bool Inv> __attribute__((noinline)) void transform_block(int offset, u32 *a, u32 *b) const {
const MontgomeryAVX2 mts;
for (int i = offset; i < offset + 512; i += (1 << (K + 2))) {
if constexpr (Inv) {
if (i == 0)
transform_fixed<K, true, true>(i, a, mts);
else
transform_fixed<K, true>(i, a, mts);
} else {
if (i == 0)
transform_fixed_pair<K, true>(i, a, b, mts);
else
transform_fixed_pair<K>(i, a, b, mts);
}
}
}
template <int LG, bool FWD_A = true, bool FWD_B = true, bool LEAF = true, bool INV = true>
void convolve_fixed(int offset, u32 *a, u32 *b) const {
const MontgomeryAVX2 mts;
if constexpr (LG > 9) {
constexpr int K = LG - 2;
if constexpr (FWD_A && FWD_B) {
if (offset == 0)
transform_fixed_pair<K, true>(offset, a, b, mts);
else
transform_fixed_pair<K>(offset, a, b, mts);
} else {
if constexpr (FWD_A) {
if (offset == 0)
transform_fixed<K, false, true>(offset, a, mts);
else
transform_fixed<K, false>(offset, a, mts);
}
if constexpr (FWD_B) {
if (offset == 0)
transform_fixed<K, false, true>(offset, b, mts);
else
transform_fixed<K, false>(offset, b, mts);
}
}
for (int j = 0; j < 4; j++) convolve_fixed<K, FWD_A, FWD_B, LEAF, INV>(offset + (j << K), a, b);
if constexpr (INV) {
if (offset == 0)
transform_fixed<K, true, true>(offset, a, mts);
else
transform_fixed<K, true>(offset, a, mts);
}
} else {
if constexpr (FWD_A || FWD_B) {
transform_block<7, false>(offset, a, b);
transform_block<5, false>(offset, a, b);
transform_block<3, false>(offset, a, b);
}
if constexpr (LEAF) {
u64x4 unused_root{};
multiply_range<3>(offset, offset + 512, a, b, unused_root);
}
if constexpr (INV) {
transform_block<3, true>(offset, a, b);
transform_block<5, true>(offset, a, b);
transform_block<7, true>(offset, a, b);
}
}
}
void prepare_quarters(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 prepare_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{});
}
// Final inverse butterfly and scaling; the last vector may be partial.
__attribute__((always_inline)) inline void finish_quarters(u32x8 x, u32x8 y, u32x8 z, u32x8 t, u32 *c, int i, int q,
int sz, u32x8 fx, u32x8 froot) const {
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_inputs(const u32 *A, int n, const u32 *B, int m, u32 *c, int lg, u32 *a, u32 *b) const {
prepare_quarters(A, n, a, a + (3 << (lg - 2)), lg);
prepare_quarters(B, m, b, b + (3 << (lg - 2)), 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 = leaf_log2(lg);
for (int j = 0; j < 4; j++) {
if (lg == 21)
convolve_fixed<19>(j << k, a, b);
else if (L == L0)
convolve_block<L0>(k, j << k, a, b, fw, iw, dot);
else
convolve_block<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 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);
finish_quarters(x, y, z, t, c, i, q, sz, fx, froot);
}
}
// First three A quarters live in c; one A quarter and one B quarter use
// 4 MiB of scratch. Inputs remain read-only, and c may be unaligned.
void convolve_reusing_output(const u32 *A, const u32 *B, u32 *c) const {
constexpr int lg = 21, k = 19, L = 3, n = 1000001, m = 1000001;
// All four B quarters are built in one streaming pass (single read of B,
// one shared i-table gather) instead of four re-reads.
u32 *temp = (u32 *)_mm_malloc((5 << 19) * 4, 32);
u32 *last = temp;
u32 *bq = temp + (1 << 19);
u32 *a = (u32 *)(((uintptr_t)c + 31) & ~uintptr_t(31));
prepare_quarters(A, n, a, last, lg);
prepare_quarters(B, m, bq, bq + (3 << 19), lg);
_mm_sfence();
for (int j = 0; j < 4; j++) {
data_origin = j << 19;
convolve_fixed<19>(data_origin, j == 3 ? last : a + data_origin, bq + data_origin);
}
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 froot = (u32x8)_mm256_set1_epi32(mt.mul(f, wr[1]));
int q = 1 << k, sz = n + m - 1;
// c may precede aligned scratch by up to seven coefficients. Writes to
// the next quarter would overwrite these tails before their final read.
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);
}
finish_quarters(x, y, z, t, c, i, q, sz, fx, froot);
}
_mm_free(temp);
}
void convolve_cyclic(int lg, u32 *a, u32 *b) const {
if (lg < 7) {
transform_forward(lg, a);
transform_forward(lg, b);
multiply_all_leaves<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 = leaf_log2(lg);
if (L == L0)
convolve_block<L0>(lg, 0, a, b, fw, iw, dot);
else
convolve_block<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 = 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.convolve_reusing_output(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 Error | Score: N/A | 显示更多 |