提交记录 40374


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 1004. 【模板题】高精度乘法 Accepted 100 13.989 ms 11000 KB C++ 29.67 KB
提交时间 评测时间
2026-08-17 23:31:27 2026-08-17 23:31:31
// 1004: radix-8 DIF/DIT, 3-mod NTT (998244353, 1004535809, 469762049),
// all-Montgomery, interleaved twiddles, base 1e9 limbs, DuckInfo direct IO.
#include <sys/auxv.h>
#include <stdint.h>
#include <string.h>
#include <immintrin.h>
#pragma GCC target("avx2")
#pragma GCC optimize("O3,unroll-loops")

typedef uint64_t u64; typedef uint32_t u32; typedef __uint128_t u128;

struct DuckInfo {
  uint64_t abi_version; const char *stdin_ptr; uint64_t stdin_size;
  char *stdout_ptr; uint64_t stdout_limit; uint64_t stdout_size;
  char *stderr_ptr; uint64_t stderr_limit; uint64_t stderr_size;
  const char *IB_ptr; uint64_t IB_limit; char *OB_ptr; uint64_t OB_limit;
  uint64_t tsc_frequency;
} __attribute__((packed));

static const u32 MODS[3] = {998244353u, 1004535809u, 469762049u};
static const u32 NINVS[3] = {998244351u, 1004535807u, 469762047u};
static const u32 R2S[3]   = {932051910u, 542374313u, 460175152u};
static const u32 ZETA[3]  = {372528824u, 395918948u, 129701348u}; // primitive 8th root (std)
static const u32 ZETA3[3] = {488723995u, 823254837u, 443138433u};
static const u32 ZETA5[3] = {625715529u, 608616861u, 340060701u};
static const u32 ZETA7[3] = {509520358u, 181280972u, 26623616u};
static const u32 IROOT[3] = {911660635u, 483363861u, 450151958u}; // sqrt(-1) (std)
static const u32 IMINUS[3]= {86583718u, 521171948u, 19610091u};
static const u32 SHP_I[3] = {3922439030u, 2066658009u, 4115675035u};
static const u32 SHP_IM[3]= {372528265u, 2228309286u, 179292260u};
static const u32 SHP_Z1[3]= {1602813089u, 1692780803u, 1185840893u};
static const u32 SHP_Z3[3]= {2102745253u, 3519887065u, 4051551378u};
static const u32 SHP_Z5[3]= {2692154206u, 2602186492u, 3109126402u};
static const u32 SHP_Z7[3]= {2192222042u, 775080230u, 243415917u};

#define SIZE (1<<18)
static u32 A[SIZE], B[SIZE];
static u32 TW[SIZE+64], ITW[SIZE+64];
static u32 tmp[SIZE/8+8];
static u32 limbsA[111112], limbsB[111112];
static u32 outlimbs[222224];
static u32 r0[SIZE], r1[SIZE], r2[SIZE];
static char tab3[1000][3];

__m256i modvec, ninvvec, shufvec;
__m256i z1v, z3v, ivv;
static u32 curmod, curninv;

static u32 powmod(u64 a, u64 e, u32 mod) { u64 r=1,b=a%mod; while(e){ if(e&1) r=r*b%mod; b=b*b%mod; e>>=1; } return (u32)r; }
static inline u32 mont_s(u32 a, u32 b, u32 mod, u32 ninv){ u64 t=(u64)a*b; u32 m=(u32)t*ninv; u32 u=(u32)((t+(u64)m*mod)>>32); if(u>=mod)u-=mod; return u; }
static inline u32 sadd(u32 a, u32 b, u32 mod){ u32 s=a+b; if(s>=mod)s-=mod; return s; }
static inline u32 ssub(u32 a, u32 b, u32 mod){ u32 d=a-b; if(d>mod)d+=mod; return d; }

static inline __m256i mont8(__m256i x, __m256i y) {
    __m256i ao=_mm256_srli_epi64(x,32);
    __m256i bo=_mm256_srli_epi64(y,32);
    __m256i te=_mm256_mul_epu32(x,y);
    __m256i to=_mm256_mul_epu32(ao,bo);
    __m256i me=_mm256_mul_epu32(te,ninvvec);
    __m256i mo=_mm256_mul_epu32(to,ninvvec);
    __m256i ue=_mm256_srli_epi64(_mm256_add_epi64(te,_mm256_mul_epu32(me,modvec)),32);
    __m256i uo=_mm256_srli_epi64(_mm256_add_epi64(to,_mm256_mul_epu32(mo,modvec)),32);
    __m256i u=_mm256_or_si256(ue,_mm256_slli_epi64(uo,32));
    return _mm256_min_epu32(u,_mm256_sub_epi32(u,modvec));
}
static inline __m256i shoup8(__m256i a, __m256i w, __m256i wp) {
    __m256i lo=_mm256_mullo_epi32(a,w);
    __m256i qe=_mm256_mul_epu32(a,wp);
    __m256i qo=_mm256_mul_epu32(_mm256_srli_epi64(a,32),_mm256_srli_epi64(wp,32));
    __m256i qe2=_mm256_srli_epi64(qe,32);
    __m256i qo2=_mm256_srli_epi64(qo,32);
    __m256i q=_mm256_or_si256(qe2,_mm256_slli_epi64(qo2,32));
    __m256i r=_mm256_sub_epi32(lo,_mm256_mullo_epi32(q,modvec));
    return _mm256_min_epu32(r,_mm256_sub_epi32(r,modvec));
}
static inline __m256i addm(__m256i a, __m256i b){ __m256i s=_mm256_add_epi32(a,b); return _mm256_min_epu32(s,_mm256_sub_epi32(s,modvec)); }
static inline __m256i subm(__m256i a, __m256i b){ __m256i d=_mm256_sub_epi32(a,b); return _mm256_min_epu32(d,_mm256_add_epi32(d,modvec)); }
static inline __m256i loadu(const u32* p){ return _mm256_loadu_si256((const __m256i*)p); }
static inline void storeu(u32* p, __m256i v){ _mm256_storeu_si256((__m256i*)p, v); }

// Vectorized geometric progression: tmp[p] = mont(wm^p), p=0..s-1 (mont8 globals must be set).
static void gen_powers(u32* tmp, int s, u32 wm, u32 mod, u32 ninv, u32 r2c) {
    tmp[0]=mont_s(1,r2c,mod,ninv);
    u32 P[8];
    P[0]=wm;
    for(int i=1;i<8;i++) P[i]=mont_s(P[i-1],wm,mod,ninv);
    __m256i Pv=_mm256_loadu_si256((__m256i*)P);
    __m256i carry=_mm256_set1_epi32((int)tmp[0]);
    int p=1;
    for(;p+8<=s;p+=8){
        __m256i chunk=mont8(carry,Pv);
        _mm256_storeu_si256((__m256i*)(tmp+p),chunk);
        carry=_mm256_set1_epi32((int)tmp[p+7]);
    }
    for(;p<s;p++) tmp[p]=mont_s(tmp[p-1],wm,mod,ninv);
}

// Interleaved twiddle gen: [w1[0..7],w2[0..7],...,w7[0..7], w1[8..15],...] (mont form)
static void gen_tw(u32* tw, u32 w, u32 mod, u32 ninv, u32 r2c) {
    int off=0;
    for(int len=SIZE;len>=8;len>>=3){ int s=len>>3; int step=SIZE/len;
        u32 ws=powmod(w,step,mod); u32 wm=mont_s(ws,r2c,mod,ninv); // ws in mont form
        gen_powers(tmp,s,wm,mod,ninv,r2c);
        for(int p=0;p<s;p+=8){
            __m256i t1=loadu(tmp+p);
            __m256i t2=mont8(t1,t1);
            __m256i t4=mont8(t2,t2);
            __m256i t3=mont8(t2,t1);
            __m256i t6=mont8(t4,t2);
            __m256i t5=mont8(t4,t1);
            __m256i t7=mont8(t6,t1);
            u32* wb=tw+off+p*7;
            storeu(wb+0,t1); storeu(wb+8,t2); storeu(wb+16,t3); storeu(wb+24,t4);
            storeu(wb+32,t5); storeu(wb+40,t6); storeu(wb+48,t7);
        }
        off+=7*s;
    }
}


// Vectorized len=8 DFT (DIF/DIT base stage, twiddles=1): 8 groups x 8 points per iter.
static void ntt8_vec(u32* x, __m256i z1, __m256i z3, __m256i iv) {
    for(int i=0;i<SIZE;i+=64){
        __m256i a0=loadu(x+i+0), a1=loadu(x+i+8), a2=loadu(x+i+16), a3=loadu(x+i+24);
        __m256i a4=loadu(x+i+32), a5=loadu(x+i+40), a6=loadu(x+i+48), a7=loadu(x+i+56);
        __m256i b0=_mm256_unpacklo_epi32(a0,a1), b1=_mm256_unpackhi_epi32(a0,a1);
        __m256i b2=_mm256_unpacklo_epi32(a2,a3), b3=_mm256_unpackhi_epi32(a2,a3);
        __m256i b4=_mm256_unpacklo_epi32(a4,a5), b5=_mm256_unpackhi_epi32(a4,a5);
        __m256i b6=_mm256_unpacklo_epi32(a6,a7), b7=_mm256_unpackhi_epi32(a6,a7);
        __m256i c0=_mm256_unpacklo_epi64(b0,b2), c1=_mm256_unpackhi_epi64(b0,b2);
        __m256i c2=_mm256_unpacklo_epi64(b1,b3), c3=_mm256_unpackhi_epi64(b1,b3);
        __m256i c4=_mm256_unpacklo_epi64(b4,b6), c5=_mm256_unpackhi_epi64(b4,b6);
        __m256i c6=_mm256_unpacklo_epi64(b5,b7), c7=_mm256_unpackhi_epi64(b5,b7);
        __m256i t0=_mm256_permute2f128_si256(c0,c4,0x20), t1=_mm256_permute2f128_si256(c1,c5,0x20);
        __m256i t2=_mm256_permute2f128_si256(c2,c6,0x20), t3=_mm256_permute2f128_si256(c3,c7,0x20);
        __m256i t4=_mm256_permute2f128_si256(c0,c4,0x31), t5=_mm256_permute2f128_si256(c1,c5,0x31);
        __m256i t6=_mm256_permute2f128_si256(c2,c6,0x31), t7=_mm256_permute2f128_si256(c3,c7,0x31);
        __m256i d0=addm(t0,t4), d1=addm(t1,t5), d2=addm(t2,t6), d3=addm(t3,t7);
        __m256i e0=subm(t0,t4);
        __m256i e1=mont8(subm(t1,t5),z1v);
        __m256i e2=mont8(subm(t2,t6),ivv);
        __m256i e3=mont8(subm(t3,t7),z3v);
        __m256i s02=addm(d0,d2), t02=subm(d0,d2);
        __m256i s13=addm(d1,d3), t13=mont8(subm(d1,d3),ivv);
        __m256i D0=addm(s02,s13), D2=subm(s02,s13);
        __m256i D1=addm(t02,t13), D3=subm(t02,t13);
        __m256i u02=addm(e0,e2), v02=subm(e0,e2);
        __m256i u13=addm(e1,e3), v13=mont8(subm(e1,e3),ivv);
        __m256i E0=addm(u02,u13), E2=subm(u02,u13);
        __m256i E1=addm(v02,v13), E3=subm(v02,v13);
        // transpose back (same 8x8 transpose): S_k = [Y0[k],...,Y7[k]]
        __m256i rb0=_mm256_unpacklo_epi32(D0,E0), rb1=_mm256_unpackhi_epi32(D0,E0);
        __m256i rb2=_mm256_unpacklo_epi32(D1,E1), rb3=_mm256_unpackhi_epi32(D1,E1);
        __m256i rb4=_mm256_unpacklo_epi32(D2,E2), rb5=_mm256_unpackhi_epi32(D2,E2);
        __m256i rb6=_mm256_unpacklo_epi32(D3,E3), rb7=_mm256_unpackhi_epi32(D3,E3);
        __m256i rc0=_mm256_unpacklo_epi64(rb0,rb2), rc1=_mm256_unpackhi_epi64(rb0,rb2);
        __m256i rc2=_mm256_unpacklo_epi64(rb1,rb3), rc3=_mm256_unpackhi_epi64(rb1,rb3);
        __m256i rc4=_mm256_unpacklo_epi64(rb4,rb6), rc5=_mm256_unpackhi_epi64(rb4,rb6);
        __m256i rc6=_mm256_unpacklo_epi64(rb5,rb7), rc7=_mm256_unpackhi_epi64(rb5,rb7);
        __m256i r0=_mm256_permute2f128_si256(rc0,rc4,0x20), r1=_mm256_permute2f128_si256(rc1,rc5,0x20);
        __m256i r2=_mm256_permute2f128_si256(rc2,rc6,0x20), r3=_mm256_permute2f128_si256(rc3,rc7,0x20);
        __m256i r4=_mm256_permute2f128_si256(rc0,rc4,0x31), r5=_mm256_permute2f128_si256(rc1,rc5,0x31);
        __m256i r6=_mm256_permute2f128_si256(rc2,rc6,0x31), r7=_mm256_permute2f128_si256(rc3,rc7,0x31);
        storeu(x+i+0,r0); storeu(x+i+8,r1); storeu(x+i+16,r2); storeu(x+i+24,r3);
        storeu(x+i+32,r4); storeu(x+i+40,r5); storeu(x+i+48,r6); storeu(x+i+56,r7);
    }
}

// ---- Hand-written asm inner loops (software-pipelined mont8) ----
u32* g_y; int g_s; const u32* g_wb; int g_cnt;

__attribute__((noinline)) static void dif8_inner(void) {
  __asm__ volatile(
    ".intel_syntax noprefix\n"
    ".macro MONT8 dst, x, w, t1, t2\n"
    "  vpsrlq \\t1, \\x, 32\n"
    "  vpsrlq \\t2, \\w, 32\n"
    "  vpmuludq \\dst, \\x, \\w\n"
    "  vpmuludq \\t1, \\t1, \\t2\n"
    "  vpmuludq \\t2, \\dst, ymm9\n"
    "  vpmuludq \\t2, \\t2, ymm8\n"
    "  vpaddq \\dst, \\dst, \\t2\n"
    "  vpsrlq \\dst, \\dst, 32\n"
    "  vpmuludq \\t2, \\t1, ymm9\n"
    "  vpmuludq \\t2, \\t2, ymm8\n"
    "  vpaddq \\t1, \\t1, \\t2\n"
    "  vpsrlq \\t1, \\t1, 32\n"
    "  vpsllq \\t1, \\t1, 32\n"
    "  vpor \\dst, \\dst, \\t1\n"
    "  vpsubd \\t1, \\dst, ymm8\n"
    "  vpminud \\dst, \\dst, \\t1\n"
    ".endm\n"
    ".macro ADSUB sreg, dreg, t1, t2\n"
    "  vpsubd \\t1, \\sreg, \\dreg\n"
    "  vpaddd \\sreg, \\sreg, \\dreg\n"
    "  vpsubd \\t2, \\sreg, ymm8\n"
    "  vpminud \\sreg, \\sreg, \\t2\n"
    "  vpaddd \\t2, \\t1, ymm8\n"
    "  vpminud \\dreg, \\t1, \\t2\n"
    ".endm\n"
    ".macro MONT8X2 dA, wA, dB, wB, t1A, t2A, t1B, t2B\n"
    "  vpsrlq \\t1A, \\dA, 32\n"
    "  vpsrlq \\t1B, \\dB, 32\n"
    "  vpsrlq \\t2A, \\wA, 32\n"
    "  vpsrlq \\t2B, \\wB, 32\n"
    "  vpmuludq \\dA, \\dA, \\wA\n"
    "  vpmuludq \\dB, \\dB, \\wB\n"
    "  vpmuludq \\t1A, \\t1A, \\t2A\n"
    "  vpmuludq \\t1B, \\t1B, \\t2B\n"
    "  vpmuludq \\t2A, \\dA, ymm9\n"
    "  vpmuludq \\t2B, \\dB, ymm9\n"
    "  vpmuludq \\t2A, \\t2A, ymm8\n"
    "  vpmuludq \\t2B, \\t2B, ymm8\n"
    "  vpaddq \\dA, \\dA, \\t2A\n"
    "  vpaddq \\dB, \\dB, \\t2B\n"
    "  vpsrlq \\dA, \\dA, 32\n"
    "  vpsrlq \\dB, \\dB, 32\n"
    "  vpmuludq \\t2A, \\t1A, ymm9\n"
    "  vpmuludq \\t2B, \\t1B, ymm9\n"
    "  vpmuludq \\t2A, \\t2A, ymm8\n"
    "  vpmuludq \\t2B, \\t2B, ymm8\n"
    "  vpaddq \\t1A, \\t1A, \\t2A\n"
    "  vpaddq \\t1B, \\t1B, \\t2B\n"
    "  vpsrlq \\t1A, \\t1A, 32\n"
    "  vpsrlq \\t1B, \\t1B, 32\n"
    "  vpsllq \\t1A, \\t1A, 32\n"
    "  vpsllq \\t1B, \\t1B, 32\n"
    "  vpor \\dA, \\dA, \\t1A\n"
    "  vpor \\dB, \\dB, \\t1B\n"
    "  vpsubd \\t1A, \\dA, ymm8\n"
    "  vpsubd \\t1B, \\dB, ymm8\n"
    "  vpminud \\dA, \\dA, \\t1A\n"
    "  vpminud \\dB, \\dB, \\t1B\n"
    ".endm\n"
    "mov rdi, g_y[rip]\n"
    "movsxd rsi, dword ptr g_s[rip]\n"
    "mov r13, g_wb[rip]\n"
    "movsxd r14, dword ptr g_cnt[rip]\n"
    "vmovdqa ymm8, modvec[rip]\n"
    "vmovdqa ymm9, ninvvec[rip]\n"
    "vmovdqa ymm10, z1v[rip]\n"
    "vmovdqa ymm11, ivv[rip]\n"
    "vmovdqa ymm12, z3v[rip]\n"
    "lea rax, [rdi + rsi*4]\n"
    "lea rbx, [rdi + rsi*8]\n"
    "lea r8, [rdi + rsi*8]\n"
    "lea r8, [r8 + rsi*4]\n"
    "lea r9, [rdi + rsi*8]\n"
    "lea r9, [r9 + rsi*8]\n"
    "lea r10, [r9 + rsi*4]\n"
    "lea r11, [r9 + rsi*8]\n"
    "lea r12, [r11 + rsi*4]\n"
    ".p2align 4\n"
    "1:\n"
    "vmovdqu ymm0, [rdi]\n"
    "vmovdqu ymm1, [rax]\n"
    "vmovdqu ymm2, [rbx]\n"
    "vmovdqu ymm3, [r8]\n"
    "vmovdqu ymm4, [r9]\n"
    "vmovdqu ymm5, [r10]\n"
    "vmovdqu ymm6, [r11]\n"
    "vmovdqu ymm7, [r12]\n"
    "vmovdqa ymm10, z1v[rip]\n"
    "vmovdqa ymm11, ivv[rip]\n"
    "vmovdqa ymm12, z3v[rip]\n"
    // level 1
    "ADSUB ymm0, ymm4, ymm13, ymm14\n"          // d0, e0
    "vpsubd ymm13, ymm1, ymm5\n"
    "vpaddd ymm1, ymm1, ymm5\n"
    "vpsubd ymm14, ymm1, ymm8\n"
    "vpminud ymm1, ymm1, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm5, ymm13, ymm10, ymm14, ymm15\n"  // e1
    "vpsubd ymm13, ymm2, ymm6\n"
    "vpaddd ymm2, ymm2, ymm6\n"
    "vpsubd ymm14, ymm2, ymm8\n"
    "vpminud ymm2, ymm2, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm6, ymm13, ymm11, ymm14, ymm15\n"  // e2
    "vpsubd ymm13, ymm3, ymm7\n"
    "vpaddd ymm3, ymm3, ymm7\n"
    "vpsubd ymm14, ymm3, ymm8\n"
    "vpminud ymm3, ymm3, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm7, ymm13, ymm12, ymm14, ymm15\n"  // e3
    // level 2 D branch
    "ADSUB ymm0, ymm2, ymm13, ymm14\n"          // s02, t02
    "vpsubd ymm13, ymm1, ymm3\n"
    "vpaddd ymm1, ymm1, ymm3\n"
    "vpsubd ymm14, ymm1, ymm8\n"
    "vpminud ymm1, ymm1, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm3, ymm13, ymm11, ymm14, ymm15\n"  // t13
    // level 2 E branch
    "ADSUB ymm4, ymm6, ymm13, ymm14\n"          // u02, v02
    "vpsubd ymm13, ymm5, ymm7\n"
    "vpaddd ymm5, ymm5, ymm7\n"
    "vpsubd ymm14, ymm5, ymm8\n"
    "vpminud ymm5, ymm5, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm7, ymm13, ymm11, ymm14, ymm15\n"  // v13
    // level 3 D branch
    "ADSUB ymm0, ymm1, ymm13, ymm14\n"          // D0, D2
    "ADSUB ymm2, ymm3, ymm13, ymm14\n"          // D1, D3
    // level 3 E branch
    "ADSUB ymm4, ymm5, ymm13, ymm14\n"          // E0, E2
    "ADSUB ymm6, ymm7, ymm13, ymm14\n"          // E1, E3
    // store D0, then 7 twiddle mont8 (2-way pipelined)
    "vmovdqu [rdi], ymm0\n"
    "vmovdqu ymm14, [r13]\n"
    "vmovdqu ymm15, [r13+32]\n"
    "MONT8X2 ymm4, ymm14, ymm2, ymm15, ymm10, ymm11, ymm12, ymm13\n"
    "vmovdqu [rax], ymm4\n"
    "vmovdqu [rbx], ymm2\n"
    "vmovdqu ymm14, [r13+64]\n"
    "vmovdqu ymm15, [r13+96]\n"
    "MONT8X2 ymm6, ymm14, ymm1, ymm15, ymm10, ymm11, ymm12, ymm13\n"
    "vmovdqu [r8], ymm6\n"
    "vmovdqu [r9], ymm1\n"
    "vmovdqu ymm14, [r13+128]\n"
    "vmovdqu ymm15, [r13+160]\n"
    "MONT8X2 ymm5, ymm14, ymm3, ymm15, ymm10, ymm11, ymm12, ymm13\n"
    "vmovdqu [r10], ymm5\n"
    "vmovdqu [r11], ymm3\n"
    "vmovdqu ymm14, [r13+192]\n"
    "MONT8 ymm7, ymm7, ymm14, ymm10, ymm11\n"
    "vmovdqu [r12], ymm7\n"
    // advance
    "add rdi, 32\n"
    "add rax, 32\n"
    "add rbx, 32\n"
    "add r8, 32\n"
    "add r9, 32\n"
    "add r10, 32\n"
    "add r11, 32\n"
    "add r12, 32\n"
    "add r13, 224\n"
    "dec r14\n"
    "jnz 1b\n"
    ".att_syntax prefix\n"
    :
    :
    : "memory", "rax","rbx","rcx","rdx","rsi","rdi","r8","r9","r10","r11","r12","r13","r14","r15","rbp",
      "ymm0","ymm1","ymm2","ymm3","ymm4","ymm5","ymm6","ymm7","ymm8","ymm9","ymm10","ymm11","ymm12","ymm13","ymm14","ymm15"
  );
}

__attribute__((noinline)) static void dit8_inner(void) {
  __asm__ volatile(
    ".intel_syntax noprefix\n"
    "mov rdi, g_y[rip]\n"
    "movsxd rsi, dword ptr g_s[rip]\n"
    "mov r13, g_wb[rip]\n"
    "movsxd r14, dword ptr g_cnt[rip]\n"
    "vmovdqa ymm8, modvec[rip]\n"
    "vmovdqa ymm9, ninvvec[rip]\n"
    "vmovdqa ymm10, z1v[rip]\n"
    "vmovdqa ymm11, ivv[rip]\n"
    "vmovdqa ymm12, z3v[rip]\n"
    "lea rax, [rdi + rsi*4]\n"
    "lea rbx, [rdi + rsi*8]\n"
    "lea r8, [rdi + rsi*8]\n"
    "lea r8, [r8 + rsi*4]\n"
    "lea r9, [rdi + rsi*8]\n"
    "lea r9, [r9 + rsi*8]\n"
    "lea r10, [r9 + rsi*4]\n"
    "lea r11, [r9 + rsi*8]\n"
    "lea r12, [r11 + rsi*4]\n"
    ".p2align 4\n"
    "1:\n"
    // load x0 plain, x1..x7 = mont8(load, twiddle)
    "vmovdqu ymm0, [rdi]\n"
    "vmovdqu ymm1, [rax]\n"
    "vmovdqu ymm14, [r13]\n"
    "MONT8 ymm1, ymm1, ymm14, ymm15, ymm13\n"
    "vmovdqu ymm2, [rbx]\n"
    "vmovdqu ymm14, [r13+32]\n"
    "MONT8 ymm2, ymm2, ymm14, ymm15, ymm13\n"
    "vmovdqu ymm3, [r8]\n"
    "vmovdqu ymm14, [r13+64]\n"
    "MONT8 ymm3, ymm3, ymm14, ymm15, ymm13\n"
    "vmovdqu ymm4, [r9]\n"
    "vmovdqu ymm14, [r13+96]\n"
    "MONT8 ymm4, ymm4, ymm14, ymm15, ymm13\n"
    "vmovdqu ymm5, [r10]\n"
    "vmovdqu ymm14, [r13+128]\n"
    "MONT8 ymm5, ymm5, ymm14, ymm15, ymm13\n"
    "vmovdqu ymm6, [r11]\n"
    "vmovdqu ymm14, [r13+160]\n"
    "MONT8 ymm6, ymm6, ymm14, ymm15, ymm13\n"
    "vmovdqu ymm7, [r12]\n"
    "vmovdqu ymm14, [r13+192]\n"
    "MONT8 ymm7, ymm7, ymm14, ymm15, ymm13\n"
    // butterfly (same as dif8)
    "ADSUB ymm0, ymm4, ymm13, ymm14\n"
    "vpsubd ymm13, ymm1, ymm5\n"
    "vpaddd ymm1, ymm1, ymm5\n"
    "vpsubd ymm14, ymm1, ymm8\n"
    "vpminud ymm1, ymm1, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm5, ymm13, ymm10, ymm14, ymm15\n"
    "vpsubd ymm13, ymm2, ymm6\n"
    "vpaddd ymm2, ymm2, ymm6\n"
    "vpsubd ymm14, ymm2, ymm8\n"
    "vpminud ymm2, ymm2, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm6, ymm13, ymm11, ymm14, ymm15\n"
    "vpsubd ymm13, ymm3, ymm7\n"
    "vpaddd ymm3, ymm3, ymm7\n"
    "vpsubd ymm14, ymm3, ymm8\n"
    "vpminud ymm3, ymm3, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm7, ymm13, ymm12, ymm14, ymm15\n"
    "ADSUB ymm0, ymm2, ymm13, ymm14\n"
    "vpsubd ymm13, ymm1, ymm3\n"
    "vpaddd ymm1, ymm1, ymm3\n"
    "vpsubd ymm14, ymm1, ymm8\n"
    "vpminud ymm1, ymm1, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm3, ymm13, ymm11, ymm14, ymm15\n"
    "ADSUB ymm4, ymm6, ymm13, ymm14\n"
    "vpsubd ymm13, ymm5, ymm7\n"
    "vpaddd ymm5, ymm5, ymm7\n"
    "vpsubd ymm14, ymm5, ymm8\n"
    "vpminud ymm5, ymm5, ymm14\n"
    "vpaddd ymm14, ymm13, ymm8\n"
    "vpminud ymm13, ymm13, ymm14\n"
    "MONT8 ymm7, ymm13, ymm11, ymm14, ymm15\n"
    "ADSUB ymm0, ymm1, ymm13, ymm14\n"
    "ADSUB ymm2, ymm3, ymm13, ymm14\n"
    "ADSUB ymm4, ymm5, ymm13, ymm14\n"
    "ADSUB ymm6, ymm7, ymm13, ymm14\n"
    // store all 8 (no twiddle)
    "vmovdqu [rdi], ymm0\n"
    "vmovdqu [rax], ymm4\n"
    "vmovdqu [rbx], ymm2\n"
    "vmovdqu [r8], ymm6\n"
    "vmovdqu [r9], ymm1\n"
    "vmovdqu [r10], ymm5\n"
    "vmovdqu [r11], ymm3\n"
    "vmovdqu [r12], ymm7\n"
    "add rdi, 32\n"
    "add rax, 32\n"
    "add rbx, 32\n"
    "add r8, 32\n"
    "add r9, 32\n"
    "add r10, 32\n"
    "add r11, 32\n"
    "add r12, 32\n"
    "add r13, 224\n"
    "dec r14\n"
    "jnz 1b\n"
    ".att_syntax prefix\n"
    :
    :
    : "memory", "rax","rbx","rcx","rdx","rsi","rdi","r8","r9","r10","r11","r12","r13","r14","r15","rbp",
      "ymm0","ymm1","ymm2","ymm3","ymm4","ymm5","ymm6","ymm7","ymm8","ymm9","ymm10","ymm11","ymm12","ymm13","ymm14","ymm15"
  );
}

// forward DIF radix-8
static void dif8(u32* x, u32 Z1, u32 Z3, u32 Iv) {
    const __m256i z1=_mm256_set1_epi32((int)Z1), z3=_mm256_set1_epi32((int)Z3), iv=_mm256_set1_epi32((int)Iv);
    int off=0;
    for(int len=SIZE;len>8;len>>=3){ int s=len>>3;
        for(int i=0;i<SIZE;i+=len){ u32* y=x+i; int p=0;
            g_y=y; g_s=s; g_wb=TW+off; g_cnt=s>>3;
            dif8_inner();
            p=s;
            for(;p<s;p++){
                u32 X0=y[p],X1=y[p+s],X2=y[p+2*s],X3=y[p+3*s],X4=y[p+4*s],X5=y[p+5*s],X6=y[p+6*s],X7=y[p+7*s];
                const u32* wb=TW+off+p*7;
                u32 d0=sadd(X0,X4,curmod),d1=sadd(X1,X5,curmod),d2=sadd(X2,X6,curmod),d3=sadd(X3,X7,curmod);
                u32 e0=ssub(X0,X4,curmod),e1=mont_s(ssub(X1,X5,curmod),Z1,curmod,curninv),e2=mont_s(ssub(X2,X6,curmod),Iv,curmod,curninv),e3=mont_s(ssub(X3,X7,curmod),Z3,curmod,curninv);
                u32 s02=sadd(d0,d2,curmod),t02=ssub(d0,d2,curmod);
                u32 s13=sadd(d1,d3,curmod),t13=mont_s(ssub(d1,d3,curmod),Iv,curmod,curninv);
                u32 D0=sadd(s02,s13,curmod),D2=ssub(s02,s13,curmod);
                u32 D1=sadd(t02,t13,curmod),D3=ssub(t02,t13,curmod);
                u32 u02=sadd(e0,e2,curmod),v02=ssub(e0,e2,curmod);
                u32 u13=sadd(e1,e3,curmod),v13=mont_s(ssub(e1,e3,curmod),Iv,curmod,curninv);
                u32 E0=sadd(u02,u13,curmod),E2=ssub(u02,u13,curmod);
                u32 E1=sadd(v02,v13,curmod),E3=ssub(v02,v13,curmod);
                y[p]=D0;
                y[p+s]=mont_s(E0,wb[0],curmod,curninv);
                y[p+2*s]=mont_s(D1,wb[8],curmod,curninv);
                y[p+3*s]=mont_s(E1,wb[16],curmod,curninv);
                y[p+4*s]=mont_s(D2,wb[24],curmod,curninv);
                y[p+5*s]=mont_s(E2,wb[32],curmod,curninv);
                y[p+6*s]=mont_s(D3,wb[40],curmod,curninv);
                y[p+7*s]=mont_s(E3,wb[48],curmod,curninv);
            }
        }
        off+=7*s;
    }
    ntt8_vec(x, z1, z3, iv);
}

// inverse DIT radix-8
static void dit8(u32* x, u32 Z1, u32 Z3, u32 Iv) {
    const __m256i z1=_mm256_set1_epi32((int)Z1), z3=_mm256_set1_epi32((int)Z3), iv=_mm256_set1_epi32((int)Iv);
    ntt8_vec(x, z1, z3, iv);
    int off=7;
    for(int len=64;len<=SIZE;len<<=3){ int s=len>>3;
        for(int i=0;i<SIZE;i+=len){ u32* y=x+i; int p=0;
            g_y=y; g_s=s; g_wb=ITW+off; g_cnt=s>>3;
            dit8_inner();
            p=s;
            for(;p<s;p++){
                u32 X0=y[p];
                u32 X1=mont_s(y[p+s],ITW[off+p*7+0],curmod,curninv);
                u32 X2=mont_s(y[p+2*s],ITW[off+p*7+8],curmod,curninv);
                u32 X3=mont_s(y[p+3*s],ITW[off+p*7+16],curmod,curninv);
                u32 X4=mont_s(y[p+4*s],ITW[off+p*7+24],curmod,curninv);
                u32 X5=mont_s(y[p+5*s],ITW[off+p*7+32],curmod,curninv);
                u32 X6=mont_s(y[p+6*s],ITW[off+p*7+40],curmod,curninv);
                u32 X7=mont_s(y[p+7*s],ITW[off+p*7+48],curmod,curninv);
                u32 d0=sadd(X0,X4,curmod),d1=sadd(X1,X5,curmod),d2=sadd(X2,X6,curmod),d3=sadd(X3,X7,curmod);
                u32 e0=ssub(X0,X4,curmod),e1=mont_s(ssub(X1,X5,curmod),Z1,curmod,curninv),e2=mont_s(ssub(X2,X6,curmod),Iv,curmod,curninv),e3=mont_s(ssub(X3,X7,curmod),Z3,curmod,curninv);
                u32 s02=sadd(d0,d2,curmod),t02=ssub(d0,d2,curmod);
                u32 s13=sadd(d1,d3,curmod),t13=mont_s(ssub(d1,d3,curmod),Iv,curmod,curninv);
                u32 D0=sadd(s02,s13,curmod),D2=ssub(s02,s13,curmod);
                u32 D1=sadd(t02,t13,curmod),D3=ssub(t02,t13,curmod);
                u32 u02=sadd(e0,e2,curmod),v02=ssub(e0,e2,curmod);
                u32 u13=sadd(e1,e3,curmod),v13=mont_s(ssub(e1,e3,curmod),Iv,curmod,curninv);
                u32 E0=sadd(u02,u13,curmod),E2=ssub(u02,u13,curmod);
                u32 E1=sadd(v02,v13,curmod),E3=ssub(v02,v13,curmod);
                y[p]=D0;y[p+s]=E0;y[p+2*s]=D1;y[p+3*s]=E1;y[p+4*s]=D2;y[p+5*s]=E2;y[p+6*s]=D3;y[p+7*s]=E3;
            }
        }
        off+=7*s;
    }
}

static void build_tab3(void){ for(int i=0;i<1000;i++){ int v=i; tab3[i][2]='0'+v%10; v/=10; tab3[i][1]='0'+v%10; v/=10; tab3[i][0]='0'+v%10; } }

// AVX2 byte-compare of two equal-length digit ranges (n up to ~1e6).
static int eq_bytes(const char* a, const char* b, size_t n){
    const char* pa=a; const char* pb=b; const char* end=a+n;
    for(; pa+64<=end; pa+=64, pb+=64){
        __m256i x0=_mm256_loadu_si256((const __m256i*)pa);
        __m256i y0=_mm256_loadu_si256((const __m256i*)pb);
        __m256i x1=_mm256_loadu_si256((const __m256i*)(pa+32));
        __m256i y1=_mm256_loadu_si256((const __m256i*)(pb+32));
        __m256i c0=_mm256_cmpeq_epi8(x0,y0);
        __m256i c1=_mm256_cmpeq_epi8(x1,y1);
        if(_mm256_movemask_epi8(c0)!=-1 || _mm256_movemask_epi8(c1)!=-1) return 0;
    }
    for(; pa<end; pa++, pb++) if(*pa!=*pb) return 0;
    return 1;
}

#ifdef LOCAL_TEST
extern uintptr_t jd_getauxval(uintptr_t);
int jd_main(){
    DuckInfo* di=(DuckInfo*)jd_getauxval(0x6b637564ull);
#else
int main(){
    DuckInfo* di=(DuckInfo*)getauxval(0x6b637564ull);
#endif
    const char* in=di->stdin_ptr; u64 inlen=di->stdin_size;
    const char* p=in; const char* inend=in+inlen;
    while(p<inend&&(*p=='\n'||*p=='\r'||*p==' '||*p=='\t'))p++;
    const char* a_start=p; while(p<inend&&*p>='0'&&*p<='9')p++; const char* a_end=p;
    while(p<inend&&(*p=='\n'||*p=='\r'||*p==' '||*p=='\t'))p++;
    const char* b_start=p; while(p<inend&&*p>='0'&&*p<='9')p++; const char* b_end=p;
    int equal=0;
    if((a_end-a_start)==(b_end-b_start) && a_end>a_start) equal=eq_bytes(a_start,b_start,(size_t)(a_end-a_start));
    const char* sa=a_start; while(sa<a_end-1&&*sa=='0')sa++;
    const char* sb=b_start; while(sb<b_end-1&&*sb=='0')sb++;
    int na=0,nb=0;
    {
        const char* pos=a_end;
        while(pos-9>=sa){const char* q=pos-9;
            u32 v=(u32)(q[0]-'0'); v=v*10+(u32)(q[1]-'0'); v=v*10+(u32)(q[2]-'0'); v=v*10+(u32)(q[3]-'0');
            v=v*10+(u32)(q[4]-'0'); v=v*10+(u32)(q[5]-'0'); v=v*10+(u32)(q[6]-'0'); v=v*10+(u32)(q[7]-'0'); v=v*10+(u32)(q[8]-'0');
            limbsA[na++]=v; pos=q;
        }
        if(pos>sa){u32 v=0; for(const char* q=sa;q<pos;q++) v=v*10+(u32)(*q-'0'); limbsA[na++]=v;}
    }
    if(equal){
        nb=na;
    } else {
        const char* pos=b_end;
        while(pos-9>=sb){const char* q=pos-9;
            u32 v=(u32)(q[0]-'0'); v=v*10+(u32)(q[1]-'0'); v=v*10+(u32)(q[2]-'0'); v=v*10+(u32)(q[3]-'0');
            v=v*10+(u32)(q[4]-'0'); v=v*10+(u32)(q[5]-'0'); v=v*10+(u32)(q[6]-'0'); v=v*10+(u32)(q[7]-'0'); v=v*10+(u32)(q[8]-'0');
            limbsB[nb++]=v; pos=q;
        }
        if(pos>sb){u32 v=0; for(const char* q=sb;q<pos;q++) v=v*10+(u32)(*q-'0'); limbsB[nb++]=v;}
    }

    shufvec=_mm256_setr_epi8(0,1,2,3,8,9,10,11,0,0,0,0,0,0,0,0, 0,1,2,3,8,9,10,11,0,0,0,0,0,0,0,0);

    for(int mi=0;mi<3;mi++){
        u32 mod=MODS[mi],ninv=NINVS[mi],r2c=R2S[mi];
        curmod=mod; curninv=ninv;
        modvec=_mm256_set1_epi32((int)mod); ninvvec=_mm256_set1_epi32((int)ninv);
        u32 w=powmod(3u,(mod-1)/SIZE,mod);
        u32 iw=powmod(w,mod-2,mod);
        u32 onem=mont_s(1,r2c,mod,ninv);
        u32 Iv=mont_s(IROOT[mi],r2c,mod,ninv);
        u32 Ivi=mont_s(IMINUS[mi],r2c,mod,ninv);
        u32 Z1=mont_s(ZETA[mi],r2c,mod,ninv), Z3=mont_s(ZETA3[mi],r2c,mod,ninv);
        u32 Z1i=mont_s(ZETA7[mi],r2c,mod,ninv), Z3i=mont_s(ZETA5[mi],r2c,mod,ninv);
        gen_tw(TW,w,mod,ninv,r2c);
        // inverse twiddles (reverse pass order, root iw)
        {int off=0;
         for(int len=8;len<=SIZE;len<<=3){int s=len>>3;int step=SIZE/len;
           u32 ws=powmod(iw,step,mod); u32 wm=mont_s(ws,r2c,mod,ninv);
           gen_powers(tmp,s,wm,mod,ninv,r2c);
           for(int pv=0;pv<s;pv+=8){
             __m256i t1=loadu(tmp+pv);
             __m256i t2=mont8(t1,t1),t4=mont8(t2,t2),t3=mont8(t2,t1),t6=mont8(t4,t2),t5=mont8(t4,t1),t7=mont8(t6,t1);
             u32* wb=ITW+off+pv*7;
             storeu(wb+0,t1);storeu(wb+8,t2);storeu(wb+16,t3);storeu(wb+24,t4);storeu(wb+32,t5);storeu(wb+40,t6);storeu(wb+48,t7);
           }
           off+=7*s;}}

        __m256i r2b=_mm256_set1_epi32((int)r2c);
        {int vlim=(na+7)>>3;for(int v=0;v<vlim;v++){__m256i x=loadu(limbsA+v*8);x=_mm256_min_epu32(x,_mm256_sub_epi32(x,modvec));x=_mm256_min_epu32(x,_mm256_sub_epi32(x,modvec));storeu(A+v*8,mont8(x,r2b));}memset(A+vlim*8,0,(SIZE-vlim*8)*4);}

        z1v=_mm256_set1_epi32((int)Z1); z3v=_mm256_set1_epi32((int)Z3); ivv=_mm256_set1_epi32((int)Iv);
        if(equal){
            dif8(A,Z1,Z3,Iv);
            for(int v=0;v<SIZE;v+=8) storeu(A+v,mont8(loadu(A+v),loadu(A+v)));
        } else {
            {int vlim=(nb+7)>>3;for(int v=0;v<vlim;v++){__m256i x=loadu(limbsB+v*8);x=_mm256_min_epu32(x,_mm256_sub_epi32(x,modvec));x=_mm256_min_epu32(x,_mm256_sub_epi32(x,modvec));storeu(B+v*8,mont8(x,r2b));}memset(B+vlim*8,0,(SIZE-vlim*8)*4);}
            dif8(A,Z1,Z3,Iv);
            dif8(B,Z1,Z3,Iv);
            for(int v=0;v<SIZE;v+=8) storeu(A+v,mont8(loadu(A+v),loadu(B+v)));
        }
        z1v=_mm256_set1_epi32((int)Z1i); z3v=_mm256_set1_epi32((int)Z3i); ivv=_mm256_set1_epi32((int)Ivi);
        dit8(A,Z1i,Z3i,Ivi);
        u32 ninvn=powmod(SIZE,mod-2,mod);
        __m256i ninvb=_mm256_set1_epi32((int)ninvn);
        u32* dst=(mi==0)?r0:(mi==1)?r1:r2;
        for(int v=0;v<SIZE;v+=8) storeu(dst+v,mont8(loadu(A+v),ninvb));
    }

    {
        u32 m0=MODS[0],m1=MODS[1],m2=MODS[2];
        u64 M0=m0,M1=m1;
        u64 INV01=powmod(M0%m1,m1-2,m1);
        u128 M01_128=(u128)M0*M1; u64 M01=(u64)M01_128;
        u64 INV012=powmod(M01%m2,m2-2,m2);
        u64 M0m2=M0%m2;
        u64 M01_q=M01/1000000000u; u64 M01_r=M01%1000000000u;
        int outlen=na+nb-1; u64 carry=0; u64 m2x2=2ull*m2;
        for(int i=0;i<outlen;i++){
            u64 a0=r0[i],a1=r1[i],a2=r2[i];
            u64 t1=a1-a0+m1; if(t1>=m1) t1-=m1;
            t1=t1*INV01%m1;
            u64 a0m2=a0; if(a0m2>=m2x2) a0m2-=m2x2; if(a0m2>=m2) a0m2-=m2;
            u64 x1m2=(a0m2+t1*M0m2)%m2;
            u64 t2=a2-x1m2+m2; if(t2>=m2) t2-=m2;
            t2=t2*INV012%m2;
            u64 d=a0+t1*M0+carry;
            u64 s=d+t2*M01_r;
            u64 carr=t2*M01_q+s/1000000000u;
            outlimbs[i]=(u32)(s%1000000000u); carry=carr;
        }
        while(carry){ outlimbs[outlen++]=(u32)(carry%1000000000u); carry/=1000000000u; }
        int hi=outlen-1; while(hi>0&&outlimbs[hi]==0) hi--;
        build_tab3();
        char* out=di->stdout_ptr; char* o=out;
        {u32 v=outlimbs[hi]; char tmp[10]; int t=0; do{ tmp[t++]='0'+(char)(v%10); v/=10; }while(v); while(t>0)*o++=tmp[--t]; }
        for(int i=hi-1;i>=0;i--){ u32 v=outlimbs[i]; u32 g2=v/1000000u; u32 g1=(v/1000u)%1000u; u32 g0=v%1000u; const char* q;
            q=tab3[g2]; *o++=q[0];*o++=q[1];*o++=q[2];
            q=tab3[g1]; *o++=q[0];*o++=q[1];*o++=q[2];
            q=tab3[g0]; *o++=q[0];*o++=q[1];*o++=q[2]; }
        *o++='\n'; di->stdout_size=(u64)(o-out);
    }
#ifdef LOCAL_TEST
    return 0;
#else
    asm volatile("mov $60, %%eax; xor %%edi, %%edi; syscall" ::: "rax","rdi","rcx","r11","memory");
    __builtin_unreachable();
#endif
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #113.989 ms10 MB + 760 KBAcceptedScore: 100


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