提交记录 36369


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 1004. 【模板题】高精度乘法 Accepted 100 19.817 ms 10868 KB C++ 15.36 KB
提交时间 评测时间
2026-08-15 02:02:28 2026-08-15 02:04:39
// 1004: 8-lane AVX2 NTT (Montgomery), 3 mods, base 1e9, NTT 2^18.
// Big stages (len>=16) vectorized 8-wide; small stages (len=8,4,2) scalar.
// Array is u32 (1MB), viewed as __m256i (8x32) for vector ops -> L3 friendly.
#include <sys/auxv.h>
#include <stdint.h>
#include <string.h>
#include <immintrin.h>
#pragma GCC target("avx2")

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 NINV[3] = {998244351u, 1004535807u, 469762047u};
static const u32 R2[3]   = {932051910u, 542374313u, 460175152u};

#define SIZE (1<<18)
#define VSIZE (SIZE/8)

static __m256i A[VSIZE] __attribute__((aligned(64)));
static __m256i B[VSIZE] __attribute__((aligned(64)));
static u32 TW[SIZE], ITW[SIZE];
static u32 limbsA[111112], limbsB[111112];
static u32 outlimbs[222224];
static u32 r0[SIZE], r1[SIZE], r2[SIZE];

static __m256i modvec, ninvvec, sub1vec;

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 mont_mul(__m256i a, __m256i b) {
  __m256i ao = _mm256_srli_epi64(a, 32);
  __m256i bo = _mm256_srli_epi64(b, 32);
  __m256i te = _mm256_mul_epu32(a, b);
  __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));
  __m256i ge = _mm256_cmpgt_epi32(u, sub1vec);
  return _mm256_sub_epi32(u, _mm256_and_si256(ge, modvec));
}
static inline __m256i vadd(__m256i a, __m256i b) {
  __m256i s = _mm256_add_epi32(a, b);
  __m256i ge = _mm256_cmpgt_epi32(s, sub1vec);
  return _mm256_sub_epi32(s, _mm256_and_si256(ge, modvec));
}
static inline __m256i vsub(__m256i a, __m256i b) {
  __m256i d = _mm256_sub_epi32(a, b);
  __m256i lt = _mm256_cmpgt_epi32(_mm256_setzero_si256(), d);
  return _mm256_add_epi32(d, _mm256_and_si256(lt, modvec));
}
static inline __m256i load_tw(const u32* tw, int j) {
  return _mm256_loadu_si256((const __m256i*)(tw + j));
}

static void fill_stage(u32* dst, int half, u32 wstep, u32 onem, u32 mod, u32 ninv) {
  const int B = 16;
  dst[0] = onem;
  int k = 1;
  for (; k < B && k < half; k++) dst[k] = mont_s(dst[k - 1], wstep, mod, ninv);
  if (half <= B) return;
  u32 wB = wstep;
  for (int t = 1; t < B; t++) wB = mont_s(wB, wstep, mod, ninv);
  __m256i d0 = _mm256_loadu_si256((const __m256i*)dst);        // dst[0..7]
  __m256i d1 = _mm256_loadu_si256((const __m256i*)(dst + 8));  // dst[8..15]
  int base = B;
  for (; base + B <= half; base += B) {
    u32 wbase = mont_s(dst[base - B], wB, mod, ninv);
    __m256i wbv = _mm256_set1_epi32((int)wbase);
    _mm256_storeu_si256((__m256i*)(dst + base), mont_mul(wbv, d0));
    _mm256_storeu_si256((__m256i*)(dst + base + 8), mont_mul(wbv, d1));
  }
  for (; base < half; base++) dst[base] = mont_s(dst[base - B], wB, mod, ninv);
}

// DIF forward big stages (len = SIZE .. 16). Returns twiddle offset.
static int ntt_fwd_big(__m256i* x, const u32* TW) {
  int n = SIZE, off = 0;
  for (int len = n; len >= 16; len >>= 1) {
    int half = len >> 1;
    int hv = half >> 3; // half in vector units
    for (int i = 0; i < n; i += len) {
      __m256i* y = x + (i >> 3);
      int b = 0;
      for (; b + 4 <= hv; b += 4) {
        __m256i u0 = y[b],     v0 = y[b + hv];
        __m256i u1 = y[b + 1], v1 = y[b + 1 + hv];
        __m256i u2 = y[b + 2], v2 = y[b + 2 + hv];
        __m256i u3 = y[b + 3], v3 = y[b + 3 + hv];
        __m256i tw0 = load_tw(TW, off + (b << 3));
        __m256i tw1 = load_tw(TW, off + (b << 3) + 8);
        __m256i tw2 = load_tw(TW, off + (b << 3) + 16);
        __m256i tw3 = load_tw(TW, off + (b << 3) + 24);
        y[b]       = vadd(u0, v0);  y[b + hv]       = mont_mul(vsub(u0, v0), tw0);
        y[b + 1]   = vadd(u1, v1);  y[b + 1 + hv]   = mont_mul(vsub(u1, v1), tw1);
        y[b + 2]   = vadd(u2, v2);  y[b + 2 + hv]   = mont_mul(vsub(u2, v2), tw2);
        y[b + 3]   = vadd(u3, v3);  y[b + 3 + hv]   = mont_mul(vsub(u3, v3), tw3);
      }
      for (; b < hv; b++) {
        __m256i u = y[b], v = y[b + hv], tw = load_tw(TW, off + (b << 3));
        y[b] = vadd(u, v);
        y[b + hv] = mont_mul(vsub(u, v), tw);
      }
    }
    off += half;
  }
  return off;
}

// DIF forward small stages (len=8,4,2) vectorized in-register.
static void ntt_fwd_small_vec(__m256i* x, const u32* TW, int off) {
  { // len=8
    __m256i tw = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)(TW + off)));
    for (int v = 0; v < VSIZE; v++) {
      __m256i u = x[v];
      __m256i us = _mm256_permute4x64_epi64(u, 0x4E);
      __m256i s = vadd(u, us);
      __m256i dd = mont_mul(vsub(us, u), tw);
      x[v] = _mm256_blend_epi32(s, dd, 0xF0);
    }
    off += 4;
  }
  { // len=4
    __m256i tw = _mm256_broadcastsi128_si256(_mm_setr_epi32(0, 0, (int)TW[off], (int)TW[off + 1]));
    for (int v = 0; v < VSIZE; v++) {
      __m256i u = x[v];
      __m256i us = _mm256_shuffle_epi32(u, 0x4E);
      __m256i s = vadd(u, us);
      __m256i dd = mont_mul(vsub(us, u), tw);
      x[v] = _mm256_blend_epi32(s, dd, 0xCC);
    }
    off += 2;
  }
  { // len=2
    for (int v = 0; v < VSIZE; v++) {
      __m256i u = x[v];
      __m256i us = _mm256_shuffle_epi32(u, 0xB1);
      __m256i s = vadd(u, us);
      __m256i d = vsub(us, u);
      x[v] = _mm256_blend_epi32(s, d, 0xAA);
    }
  }
}

// DIT inverse small stages (len=2,4,8) vectorized in-register.
static void ntt_inv_small_vec(__m256i* x, const u32* ITW) {
  int off = 0;
  { // len=2 (same as forward)
    for (int v = 0; v < VSIZE; v++) {
      __m256i u = x[v];
      __m256i us = _mm256_shuffle_epi32(u, 0xB1);
      __m256i s = vadd(u, us);
      __m256i d = vsub(us, u);
      x[v] = _mm256_blend_epi32(s, d, 0xAA);
    }
    off += 1;
  }
  { // len=4
    __m256i tw = _mm256_broadcastsi128_si256(_mm_setr_epi32((int)ITW[off], (int)ITW[off + 1], (int)ITW[off], (int)ITW[off + 1]));
    for (int v = 0; v < VSIZE; v++) {
      __m256i u = x[v];
      __m256i us = _mm256_shuffle_epi32(u, 0x4E);
      __m256i vv = mont_mul(us, tw);
      __m256i s = vadd(u, vv);
      __m256i d = _mm256_shuffle_epi32(vsub(u, vv), 0x4E);
      x[v] = _mm256_blend_epi32(s, d, 0xCC);
    }
    off += 2;
  }
  { // len=8
    __m256i tw = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)(ITW + off)));
    for (int v = 0; v < VSIZE; v++) {
      __m256i u = x[v];
      __m256i us = _mm256_permute4x64_epi64(u, 0x4E);
      __m256i vv = mont_mul(us, tw);
      __m256i s = vadd(u, vv);
      __m256i d = _mm256_permute4x64_epi64(vsub(u, vv), 0x4E);
      x[v] = _mm256_blend_epi32(s, d, 0xF0);
    }
  }
}

// DIT inverse big stages (len=16 .. SIZE). ITW offset starts after small stages (=7).
static void ntt_inv_big(__m256i* x, const u32* ITW) {
  int n = SIZE, off = 7;
  for (int len = 16; len <= n; len <<= 1) {
    int half = len >> 1;
    int hv = half >> 3;
    for (int i = 0; i < n; i += len) {
      __m256i* y = x + (i >> 3);
      int b = 0;
      for (; b + 4 <= hv; b += 4) {
        __m256i u0 = y[b],     v0 = mont_mul(y[b + hv],       load_tw(ITW, off + (b << 3)));
        __m256i u1 = y[b + 1], v1 = mont_mul(y[b + 1 + hv],   load_tw(ITW, off + (b << 3) + 8));
        __m256i u2 = y[b + 2], v2 = mont_mul(y[b + 2 + hv],   load_tw(ITW, off + (b << 3) + 16));
        __m256i u3 = y[b + 3], v3 = mont_mul(y[b + 3 + hv],   load_tw(ITW, off + (b << 3) + 24));
        y[b]       = vadd(u0, v0);  y[b + hv]       = vsub(u0, v0);
        y[b + 1]   = vadd(u1, v1);  y[b + 1 + hv]   = vsub(u1, v1);
        y[b + 2]   = vadd(u2, v2);  y[b + 2 + hv]   = vsub(u2, v2);
        y[b + 3]   = vadd(u3, v3);  y[b + 3 + hv]   = vsub(u3, v3);
      }
      for (; b < hv; b++) {
        __m256i u = y[b], v = mont_mul(y[b + hv], load_tw(ITW, off + (b << 3)));
        y[b] = vadd(u, v);
        y[b + hv] = vsub(u, v);
      }
    }
    off += half;
  }
}

static char tab3[1000][3];
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;
  }
}

#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;

  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;
    }
  }
  {
    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;
    }
  }

  for (int mi = 0; mi < 3; mi++) {
    u32 mod = MODS[mi], ninv = NINV[mi], r2c = R2[mi];
    modvec = _mm256_set1_epi32((int)mod);
    ninvvec = _mm256_set1_epi32((int)ninv);
    sub1vec = _mm256_set1_epi32((int)(mod - 1));

    u32 w = powmod(3u, (mod - 1) / SIZE, mod);
    u32 iw = powmod(w, mod - 2, mod);
    u32 onem = mont_s(1, r2c, mod, ninv);

    {
      int off = 0;
      for (int len = SIZE; len > 1; len >>= 1) {
        int half = len >> 1;
        int step = SIZE / len;
        u32 wstep = mont_s(powmod(w, step, mod), r2c, mod, ninv);
        fill_stage(TW + off, half, wstep, onem, mod, ninv);
        off += half;
      }
    }
    {
      int off = 0;
      for (int len = 2; len <= SIZE; len <<= 1) {
        int half = len >> 1;
        int step = SIZE / len;
        u32 wstep = mont_s(powmod(iw, step, mod), r2c, mod, ninv);
        fill_stage(ITW + off, half, wstep, onem, mod, ninv);
        off += half;
      }
    }

    __m256i r2b = _mm256_set1_epi32((int)r2c);
    __m256i zero = _mm256_setzero_si256();
    {
      int vlim = (na + 7) >> 3;
      for (int v = 0; v < vlim; v++) {
        u32 x[8];
        for (int k = 0; k < 8; k++) {
          int idx = v * 8 + k;
          u32 val = (idx < na) ? limbsA[idx] : 0u;
          while (val >= mod) val -= mod;
          x[k] = val;
        }
        A[v] = mont_mul(_mm256_setr_epi32((int)x[0],(int)x[1],(int)x[2],(int)x[3],(int)x[4],(int)x[5],(int)x[6],(int)x[7]), r2b);
      }
      for (int v = vlim; v < VSIZE; v++) A[v] = zero;
    }
    {
      int vlim = (nb + 7) >> 3;
      for (int v = 0; v < vlim; v++) {
        u32 x[8];
        for (int k = 0; k < 8; k++) {
          int idx = v * 8 + k;
          u32 val = (idx < nb) ? limbsB[idx] : 0u;
          while (val >= mod) val -= mod;
          x[k] = val;
        }
        B[v] = mont_mul(_mm256_setr_epi32((int)x[0],(int)x[1],(int)x[2],(int)x[3],(int)x[4],(int)x[5],(int)x[6],(int)x[7]), r2b);
      }
      for (int v = vlim; v < VSIZE; v++) B[v] = zero;
    }

    int off = ntt_fwd_big(A, TW);
    ntt_fwd_big(B, TW);
    ntt_fwd_small_vec(A, TW, off);
    ntt_fwd_small_vec(B, TW, off);

    for (int v = 0; v < VSIZE; v++) A[v] = mont_mul(A[v], B[v]);

    ntt_inv_small_vec(A, ITW);
    ntt_inv_big(A, ITW);

    u32 ninvn = powmod(SIZE, mod - 2, mod);
    __m256i ninvnb = _mm256_set1_epi32((int)ninvn);
    for (int v = 0; v < VSIZE; v++) A[v] = mont_mul(A[v], ninvnb);
    u32* dst = (mi == 0) ? r0 : (mi == 1) ? r1 : r2;
    for (int v = 0; v < VSIZE; v++) _mm256_storeu_si256((__m256i*)(dst + v * 8), A[v]);
  }

  {
    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;              // < 2^60, fits u64
    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];
      // a0 < m0 < m1, so a0 % m1 == a0
      u64 t1 = a1 - a0 + m1; if (t1 >= m1) t1 -= m1;
      t1 = t1 * INV01 % m1;
      // a0 % m2 (a0 < m0 < 3*m2)
      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 #119.817 ms10 MB + 628 KBAcceptedScore: 100


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