提交记录 32024


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 1004. 【模板题】高精度乘法 Accepted 100 85.674 ms 36472 KB C++ 8.66 KB
提交时间 评测时间
2026-08-14 10:23:03 2026-08-14 10:23:29
// 1004 high precision multiply: AVX2 3-mod packed NTT (Montgomery), base 1e9, NTT 2^18
#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 GROOT[3] = {3u, 3u, 3u};
static const u32 NINV[3] = {998244351u, 1004535807u, 469762047u};
static const u32 R2[3]   = {932051910u, 542374313u, 460175152u};

#define SIZE (1<<18)

static __m256i A[SIZE] __attribute__((aligned(64)));
static __m256i B[SIZE] __attribute__((aligned(64)));
static __m256i roots[SIZE] __attribute__((aligned(64)));
static __m256i iroots[SIZE] __attribute__((aligned(64)));
static u32 limbsA[111112], limbsB[111112];

static __m256i modvec, ninvvec, sub1vec, MASK32, ZERO;

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

// a,b in low 32 bits of 4x64 lanes; per-lane Montgomery mul mod modvec
static inline __m256i mont_mul(__m256i a, __m256i b) {
  __m256i t = _mm256_mul_epu32(a, b);
  __m256i m = _mm256_mul_epu32(t, ninvvec);
  m = _mm256_and_si256(m, MASK32);
  __m256i mp = _mm256_mul_epu32(m, modvec);
  __m256i u = _mm256_add_epi64(t, mp);
  u = _mm256_srli_epi64(u, 32);
  u = _mm256_and_si256(u, MASK32);
  __m256i ge = _mm256_cmpgt_epi32(u, sub1vec);
  u = _mm256_sub_epi32(u, _mm256_and_si256(ge, modvec));
  return u;
}

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(ZERO, d);
  return _mm256_add_epi32(d, _mm256_and_si256(lt, modvec));
}

static inline void ntt_fwd(__m256i* x) {
  int n = SIZE;
  for (int len = n; len > 1; len >>= 1) {
    int half = len >> 1;
    int step = n / len;
    for (int i = 0; i < n; i += len) {
      __m256i* y = x + i;
      for (int j = 0; j < half; j++) {
        __m256i u = y[j];
        __m256i v = y[j + half];
        y[j] = vadd(u, v);
        y[j + half] = mont_mul(vsub(u, v), roots[j * step]);
      }
    }
  }
}

static inline void ntt_inv(__m256i* x) {
  int n = SIZE;
  for (int len = 2; len <= n; len <<= 1) {
    int half = len >> 1;
    int step = n / len;
    for (int i = 0; i < n; i += len) {
      __m256i* y = x + i;
      for (int j = 0; j < half; j++) {
        __m256i u = y[j];
        __m256i v = mont_mul(y[j + half], iroots[j * step]);
        y[j] = vadd(u, v);
        y[j + half] = vsub(u, v);
      }
    }
  }
}

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

static inline __m256i pack3(u32 a, u32 b, u32 c) {
  return _mm256_setr_epi32((int)a, 0, (int)b, 0, (int)c, 0, 0, 0);
}

#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 > sa) {
      const char* st = pos - 9; if (st < sa) st = sa;
      u32 v = 0;
      for (const char* q = st; q < pos; q++) v = v * 10 + (u32)(*q - '0');
      limbsA[na++] = v;
      pos = st;
    }
  }
  {
    const char* pos = b_end;
    while (pos > sb) {
      const char* st = pos - 9; if (st < sb) st = sb;
      u32 v = 0;
      for (const char* q = st; q < pos; q++) v = v * 10 + (u32)(*q - '0');
      limbsB[nb++] = v;
      pos = st;
    }
  }

  const u32 m0 = MODS[0], m1 = MODS[1], m2 = MODS[2];
  MASK32 = _mm256_set1_epi64x(0xFFFFFFFFLL);
  ZERO = _mm256_setzero_si256();
  modvec = pack3(m0, m1, m2);
  ninvvec = pack3(NINV[0], NINV[1], NINV[2]);
  sub1vec = pack3(m0 - 1, m1 - 1, m2 - 1);

  // build roots (Montgomery form, packed)
  {
    u32 w0 = powmod(GROOT[0], (MODS[0] - 1) / SIZE, MODS[0]);
    u32 w1 = powmod(GROOT[1], (MODS[1] - 1) / SIZE, MODS[1]);
    u32 w2 = powmod(GROOT[2], (MODS[2] - 1) / SIZE, MODS[2]);
    u32 iw0 = powmod(w0, MODS[0] - 2, MODS[0]);
    u32 iw1 = powmod(w1, MODS[1] - 2, MODS[1]);
    u32 iw2 = powmod(w2, MODS[2] - 2, MODS[2]);
    __m256i r2vec = pack3(R2[0], R2[1], R2[2]);
    __m256i wv = mont_mul(pack3(w0, w1, w2), r2vec);
    __m256i iwv = mont_mul(pack3(iw0, iw1, iw2), r2vec);
    __m256i one = mont_mul(pack3(1, 1, 1), r2vec);
    __m256i cw = one;
    for (int i = 0; i < SIZE; i++) { roots[i] = cw; cw = mont_mul(cw, wv); }
    __m256i ci = one;
    for (int i = 0; i < SIZE; i++) { iroots[i] = ci; ci = mont_mul(ci, iwv); }
  }

  // fill A, B with packed residues (Montgomery form)
  {
    __m256i r2vec = pack3(R2[0], R2[1], R2[2]);
    for (int i = 0; i < na; i++) {
      u32 v = limbsA[i];
      u32 a0 = v; if (a0 >= m0) a0 -= m0;
      u32 a1 = v; // v < 1e9 < m1
      u32 a2 = v; while (a2 >= m2) a2 -= m2;
      A[i] = mont_mul(pack3(a0, a1, a2), r2vec);
    }
    for (int i = na; i < SIZE; i++) A[i] = ZERO;
    for (int i = 0; i < nb; i++) {
      u32 v = limbsB[i];
      u32 b0 = v; if (b0 >= m0) b0 -= m0;
      u32 b1 = v;
      u32 b2 = v; while (b2 >= m2) b2 -= m2;
      B[i] = mont_mul(pack3(b0, b1, b2), r2vec);
    }
    for (int i = nb; i < SIZE; i++) B[i] = ZERO;
  }

  ntt_fwd(A);
  ntt_fwd(B);
  for (int i = 0; i < SIZE; i++) A[i] = mont_mul(A[i], B[i]);
  ntt_inv(A);

  // scale by n^-1 (normal form) -> residues in normal form
  {
    u32 n0 = powmod(SIZE, MODS[0] - 2, MODS[0]);
    u32 n1 = powmod(SIZE, MODS[1] - 2, MODS[1]);
    u32 n2 = powmod(SIZE, MODS[2] - 2, MODS[2]);
    __m256i ninvN = pack3(n0, n1, n2);
    for (int i = 0; i < SIZE; i++) A[i] = mont_mul(A[i], ninvN);
  }

  // CRT reconstruction (Garner), carry in base 1e9, then output
  {
    u64 M0 = m0, M1 = m1;
    u64 INV01 = powmod(M0 % m1, m1 - 2, m1);
    u128 M01 = (u128)M0 * M1;
    u64 INV012 = powmod((u64)(M01 % m2), m2 - 2, m2);

    int outlen = na + nb - 1;
    static u32 outlimbs[222224];
    u128 carry = 0;
    int nout = 0;
    for (int i = 0; i < outlen; i++) {
      __m256i v = A[i];
      u32 r0 = (u32)_mm256_extract_epi32(v, 0);
      u32 r1 = (u32)_mm256_extract_epi32(v, 2);
      u32 r2 = (u32)_mm256_extract_epi32(v, 4);
      u128 c = r0;
      u64 t1 = (u64)(r1 - (u32)(c % m1) + m1) % m1;
      t1 = (u64)((u128)t1 * INV01 % m1);
      c += (u128)t1 * M0;
      u64 t2 = (u64)(r2 - (u32)(c % m2) + m2) % m2;
      t2 = (u64)((u128)t2 * INV012 % m2);
      c += (u128)t2 * M01;
      c += carry;
      outlimbs[i] = (u32)(c % 1000000000u);
      carry = c / 1000000000u;
    }
    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 #185.674 ms35 MB + 632 KBAcceptedScore: 100


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