提交记录 87328


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 mmms2k. 测测你的短整数矩阵乘法-2k Accepted 100 148.457 ms 8340 KB C++17 5.29 KB
提交时间 评测时间
2026-09-25 02:21:48 2026-09-25 02:21:50
// mmms* -- C = A*B for short matrices, exact mod 2^16.
// Packed GEMM, MR=6 x NR=32, hand-scheduled asm micro-kernel, vpmullw+vpaddw in 16-bit lanes.
// MS_REP=1: A panel pre-replicated 16x -> the A operand is a fused 32-byte load (no port-5 use).
// MS_REP=0: vpbroadcastw from the raw A panel.
#include <immintrin.h>
#include <stddef.h>
#include <stdlib.h>
#pragma GCC push_options
#pragma GCC target("avx2")

#ifndef MS_REP
#define MS_REP 1
#endif
#define MR 6
#define NR 32
#ifndef KC
#define KC 64
#endif
#ifndef NC
#define NC 1024
#endif

static short Apack_rep[MR * KC * 16 + 64] __attribute__((aligned(64)));
static short Apack_raw[MR * KC + 64] __attribute__((aligned(64)));
static short Bpack[(size_t)NC * KC + 128] __attribute__((aligned(64)));

static inline __m256i b16(const short *p) {
  __m256i r; __asm__("vpbroadcastw %1, %0" : "=x"(r) : "m"(*p)); return r;
}

#if MS_REP
// One fused-uop load per (row, k); A panel holds 16 copies (32B) of each value.
#define KSTEP(off, C0, C1) \
  "vpmullw " #off "(%[a]), %%ymm14, %%ymm13\n\t" \
  "vpaddw %%ymm13, %" C0 ", %" C0 "\n\t" \
  "vpmullw " #off "(%[a]), %%ymm15, %%ymm13\n\t" \
  "vpaddw %%ymm13, %" C1 ", %" C1 "\n\t"
#define MICRO_BODY \
  "1:\n\t" \
  "vmovdqa 0(%[b]), %%ymm14\n\t" \
  "vmovdqa 32(%[b]), %%ymm15\n\t" \
  KSTEP(0,   "0",  "1") KSTEP(32,  "2",  "3") KSTEP(64,  "4",  "5") \
  KSTEP(96,  "6",  "7") KSTEP(128, "8",  "9") KSTEP(160, "10", "11") \
  "addq $192, %[a]\n\t" \
  "addq $64, %[b]\n\t" \
  "decq %[k]\n\t" \
  "jnz 1b\n\t"
#else
#define KSTEP(off, C0, C1) \
  "vpmullw %%ymm13, %%ymm14, %%ymm12\n\t" \
  "vpaddw %%ymm12, %" C0 ", %" C0 "\n\t" \
  "vpmullw %%ymm13, %%ymm15, %%ymm12\n\t" \
  "vpaddw %%ymm12, %" C1 ", %" C1 "\n\t"
#define MKB(off) "vpbroadcastw " #off "(%[a]), %%ymm13\n\t"
#define MICRO_BODY \
  "1:\n\t" \
  "vmovdqa 0(%[b]), %%ymm14\n\t" \
  "vmovdqa 32(%[b]), %%ymm15\n\t" \
  MKB(0)  KSTEP(0,  "0",  "1")  MKB(2)  KSTEP(0,  "2",  "3") \
  MKB(4)  KSTEP(0,  "4",  "5")  MKB(6)  KSTEP(0,  "6",  "7") \
  MKB(8)  KSTEP(0,  "8",  "9")  MKB(10) KSTEP(0,  "10", "11") \
  "addq $12, %[a]\n\t" \
  "addq $64, %[b]\n\t" \
  "decq %[k]\n\t" \
  "jnz 1b\n\t"
#endif

#define MICRO_CLOB "ymm12","ymm13","ymm14","ymm15","cc","memory"

static inline void micro(long k, const short *ap, const short *bp, short *C, int ldc, int rows) {
  __m256i c00=_mm256_setzero_si256(), c01=c00, c10=c00, c11=c00, c20=c00, c21=c00;
  __m256i c30=c00, c31=c00, c40=c00, c41=c00, c50=c00, c51=c00;
  __asm__ volatile(
    MICRO_BODY
    : "+x"(c00),"+x"(c01),"+x"(c10),"+x"(c11),"+x"(c20),"+x"(c21),
      "+x"(c30),"+x"(c31),"+x"(c40),"+x"(c41),"+x"(c50),"+x"(c51),
      [a] "+r"(ap), [b] "+r"(bp), [k] "+r"(k)
    :
    : MICRO_CLOB);
#define ST(r, v0, v1) if (rows > (r)) { \
    short *p0 = C + (size_t)(r) * ldc; \
    _mm256_storeu_si256((__m256i *)p0, _mm256_add_epi16(_mm256_loadu_si256((const __m256i *)p0), v0)); \
    _mm256_storeu_si256((__m256i *)(p0 + 16), _mm256_add_epi16(_mm256_loadu_si256((const __m256i *)(p0 + 16)), v1)); }
  ST(0, c00, c01) ST(1, c10, c11) ST(2, c20, c21)
  ST(3, c30, c31) ST(4, c40, c41) ST(5, c50, c51)
#undef ST
}

void matrix_multiply(int n, const short *A, const short *B, short *C) {
  const int nbr = n;                     // partial row blocks go through the vector path
  const int nbc = n - (n % NR);
  for (size_t i = 0; i < (size_t)n * n; i++) C[i] = 0;

  for (int jc = 0; jc < nbc; jc += NC) {
    int nc = (nbc - jc < NC) ? (nbc - jc) : NC;
    for (int pc = 0; pc < n; pc += KC) {
      int kc = (n - pc < KC) ? (n - pc) : KC;
      for (int jr = 0, jb = 0; jr < nc; jr += NR, jb++) {
        short *q = Bpack + (size_t)jb * kc * NR;
        for (int k = 0; k < kc; k++) {
          const short *src = B + (size_t)(pc + k) * n + jc + jr;
          for (int j = 0; j < NR; j++) q[j] = src[j];
          q += NR;
        }
      }
      for (int ic = 0; ic < nbr; ic += MR) {
        int rows = (nbr - ic < MR) ? (nbr - ic) : MR;
#if MS_REP
        short *p = Apack_rep;
        for (int k = 0; k < kc; k++) {
          for (int r = 0; r < rows; r++) {
            __m256i v = b16(A + (size_t)(ic + r) * n + pc + k);
            _mm256_store_si256((__m256i *)p, v);
            p += 16;
          }
          for (int r = rows; r < MR; r++) {
            _mm256_store_si256((__m256i *)p, _mm256_setzero_si256());
            p += 16;
          }
        }
        const short *b0 = Bpack;
        short *cp = C + (size_t)ic * n + jc;
        for (int jr = 0; jr < nc; jr += NR) {
          micro(kc, Apack_rep, b0, cp + jr, n, rows);
          b0 += (size_t)kc * NR;
        }
#else
        short *p = Apack_raw;
        for (int k = 0; k < kc; k++) {
          for (int r = 0; r < rows; r++) p[r] = A[(size_t)(ic + r) * n + pc + k];
          for (int r = rows; r < MR; r++) p[r] = 0;
          p += MR;
        }
        const short *b0 = Bpack;
        short *cp = C + (size_t)ic * n + jc;
        for (int jr = 0; jr < nc; jr += NR) {
          micro(kc, Apack_raw, b0, cp + jr, n, rows);
          b0 += (size_t)kc * NR;
        }
#endif
      }
    }
  }
  for (int i = 0; i < n; i++)
    for (int j = nbc; j < n; j++) {
      unsigned acc = 0;
      for (int k = 0; k < n; k++)
        acc += (unsigned)(unsigned short)A[(size_t)i * n + k] * (unsigned)(unsigned short)B[(size_t)k * n + j];
      C[(size_t)i * n + j] = (short)acc;
    }
}
#pragma GCC pop_options

CompilationN/AN/ACompile OKScore: N/A

Testcase #1148.457 ms8 MB + 148 KBAcceptedScore: 100


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