提交记录 48177


用户 题目 状态 得分 用时 内存 语言 代码长度
iMMIQ 1002. 测测你的多项式乘法 Accepted 100 20.512 ms 11936 KB C++17 41.31 KB
提交时间 评测时间
2026-09-16 10:50:52 2026-09-16 10:50:58
// 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+2048;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>11){
            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<9,false>(offset,a,b,fw);fixed_level<7,false>(offset,a,b,fw);
            fixed_level<5,false>(offset,a,b,fw);fixed_level<3,false>(offset,a,b,fw);
            dot_range<3>(offset,offset+2048,a,b,dot);
            fixed_level<3,true>(offset,a,b,iw);fixed_level<5,true>(offset,a,b,iw);
            fixed_level<7,true>(offset,a,b,iw);fixed_level<9,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);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #120.512 ms11 MB + 672 KBAcceptedScore: 100


Judge Duck Online | 评测鸭在线
Server Time: 2026-09-17 17:49:15 | Loaded in 1 ms | Server Status
个人娱乐项目,仅供学习交流使用 | 捐赠