提交记录 87697


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 mmmi1k. 测测你的整数矩阵乘法-1k Accepted 100 51.333 ms 4432 KB C++17 8.26 KB
提交时间 评测时间
2026-09-25 03:10:46 2026-09-25 03:10:48
/* int32 GEMM exact mod 2^32, AVX2 (Skylake/Coffee Lake).
 * Pre-broadcast A panel + memory-operand vpmaddwd: no port-5 uop in the inner loop.
 * 3 madd + 3 add per 2k x 8 columns -> ceiling 8 MAC/cycle (vs 5.33 for bcst scheme).
 * GENERATED by gen_k2.py.
 */
#include <immintrin.h>
#include <string.h>
#include <stdint.h>

#define MR 4
#define NCG 1
#define KC 128
#define MC 16
#define NC 512
#define NPAIR (KC / 2)
#define NR (8 * NCG)

#define AVX2 __attribute__((target("avx2")))
#define LD(p) _mm256_loadu_si256((const __m256i *)(p))
#define ST(p, v) _mm256_storeu_si256((__m256i *)(p), (v))

static int32_t Apk[(MC + 4) * NPAIR * 16 + 64] __attribute__((aligned(64)));
static int32_t Bpk[2 * NPAIR * NC + 64] __attribute__((aligned(64)));  /* [jg][m][B0|B1] */
/* MIKRO-BEGIN */
#define AVX2 __attribute__((target("avx2")))
static __m256i LA[128] __attribute__((aligned(32)));
AVX2 static inline void mikro(int npair, const int32_t *ap, const int32_t *bp,
                              int32_t *c, int ldc) {
  (void)npair;
  __asm__ volatile(
    "vpxor %%xmm15, %%xmm15, %%xmm15\n\t"
    "vmovdqa %%ymm15, %%ymm0\n\t"
    "vmovdqa %%ymm15, %%ymm4\n\t"
    "vmovdqa %%ymm15, %%ymm1\n\t"
    "vmovdqa %%ymm15, %%ymm5\n\t"
    "vmovdqa %%ymm15, %%ymm2\n\t"
    "vmovdqa %%ymm15, %%ymm6\n\t"
    "vmovdqa %%ymm15, %%ymm3\n\t"
    "vmovdqa %%ymm15, %%ymm7\n\t"
    "shr $1, %[cnt]\n\t"
    "test %[cnt], %[cnt]\n\t"
    "jle 2f\n\t"
    "1:\n\t"
    "vmovdqu 0(%[bp]), %%ymm8\n\t"
    "vmovdqu 32(%[bp]), %%ymm9\n\t"
    "vpmaddwd 0(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
    "vpmaddwd 0(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 32(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 4096(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm1, %%ymm1\n\t"
    "vpmaddwd 4096(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 4128(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 8192(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm2, %%ymm2\n\t"
    "vpmaddwd 8192(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 8224(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 12288(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm3, %%ymm3\n\t"
    "vpmaddwd 12288(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 12320(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm7, %%ymm7\n\t"
    "vmovdqu 64(%[bp]), %%ymm8\n\t"
    "vmovdqu 96(%[bp]), %%ymm9\n\t"
    "vpmaddwd 64(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
    "vpmaddwd 64(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 96(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 4160(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm1, %%ymm1\n\t"
    "vpmaddwd 4160(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 4192(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 8256(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm2, %%ymm2\n\t"
    "vpmaddwd 8256(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 8288(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 12352(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm3, %%ymm3\n\t"
    "vpmaddwd 12352(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 12384(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm10, %%ymm7, %%ymm7\n\t"
    "add $128, %[ap]\n\t"
    "add $128, %[bp]\n\t"
    "dec %[cnt]\n\t"
    "jnz 1b\n\t"
    "2:\n\t"
    "vmovdqu %%ymm0, 0(%[la])\n\t"
    "vmovdqu %%ymm4, 32(%[la])\n\t"
    "vmovdqu %%ymm1, 64(%[la])\n\t"
    "vmovdqu %%ymm5, 96(%[la])\n\t"
    "vmovdqu %%ymm2, 128(%[la])\n\t"
    "vmovdqu %%ymm6, 160(%[la])\n\t"
    "vmovdqu %%ymm3, 192(%[la])\n\t"
    "vmovdqu %%ymm7, 224(%[la])\n\t"
    ""
    : [ap] "+r" (ap), [bp] "+r" (bp), [cnt] "+r" (npair)
    : [la] "r" (&LA[0])
    : "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm15", "memory");
  { int32_t *cp = c + (size_t)0 * ldc + 0;
    __m256i rs = _mm256_add_epi32(LA[0], _mm256_slli_epi32(LA[1], 16));
    _mm256_storeu_si256((__m256i *)cp, _mm256_add_epi32(_mm256_loadu_si256((const __m256i *)cp), rs)); }
  { int32_t *cp = c + (size_t)1 * ldc + 0;
    __m256i rs = _mm256_add_epi32(LA[2], _mm256_slli_epi32(LA[3], 16));
    _mm256_storeu_si256((__m256i *)cp, _mm256_add_epi32(_mm256_loadu_si256((const __m256i *)cp), rs)); }
  { int32_t *cp = c + (size_t)2 * ldc + 0;
    __m256i rs = _mm256_add_epi32(LA[4], _mm256_slli_epi32(LA[5], 16));
    _mm256_storeu_si256((__m256i *)cp, _mm256_add_epi32(_mm256_loadu_si256((const __m256i *)cp), rs)); }
  { int32_t *cp = c + (size_t)3 * ldc + 0;
    __m256i rs = _mm256_add_epi32(LA[6], _mm256_slli_epi32(LA[7], 16));
    _mm256_storeu_si256((__m256i *)cp, _mm256_add_epi32(_mm256_loadu_si256((const __m256i *)cp), rs)); }
}
/* MIKRO-END */

AVX2 static inline void packA(int ic, int kk, int n, const int32_t *A) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  const __m256i i0 = _mm256_set1_epi32(0), i1 = _mm256_set1_epi32(1);
  const __m256i i4 = _mm256_set1_epi32(4), i5 = _mm256_set1_epi32(5);
  for (int r = 0; r < MC; r++) {
    const int32_t *src = A + (size_t)(ic + r) * n + kk;
    int32_t *dst = Apk + (size_t)r * NPAIR * 16;
    for (int m = 0; m < KC / 8; m++) {
      __m256i x = LD(src + 8 * m);
      __m256i al = _mm256_and_si256(x, m16);
      __m256i ah = _mm256_srli_epi32(_mm256_add_epi32(x, half), 16);
      __m256i pl = _mm256_packus_epi32(al, al);
      __m256i ph = _mm256_packus_epi32(ah, ah);
      int32_t *d = dst + 64 * m;
      ST(d + 0, _mm256_permutevar8x32_epi32(pl, i0));
      ST(d + 8, _mm256_permutevar8x32_epi32(ph, i0));
      ST(d + 16, _mm256_permutevar8x32_epi32(pl, i1));
      ST(d + 24, _mm256_permutevar8x32_epi32(ph, i1));
      ST(d + 32, _mm256_permutevar8x32_epi32(pl, i4));
      ST(d + 40, _mm256_permutevar8x32_epi32(ph, i4));
      ST(d + 48, _mm256_permutevar8x32_epi32(pl, i5));
      ST(d + 56, _mm256_permutevar8x32_epi32(ph, i5));
    }
  }
}

AVX2 static inline void packB(int jc, int kk, int n, const int32_t *B) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  for (int m = 0; m < NPAIR; m++) {
    const int32_t *s0 = B + (size_t)(kk + 2 * m) * n + jc;
    const int32_t *s1 = B + (size_t)(kk + 2 * m + 1) * n + jc;
    int32_t *d0 = Bpk + (size_t)m * 16;
    int32_t *d1 = d0 + 8;
    for (int j = 0; j < NC; j += 8) {
      __m256i x0 = LD(s0 + j), x1 = LD(s1 + j);
      __m256i l0 = _mm256_and_si256(x0, m16);
      __m256i l1 = _mm256_and_si256(x1, m16);
      __m256i h0 = _mm256_srli_epi32(_mm256_add_epi32(x0, half), 16);
      __m256i h1 = _mm256_srli_epi32(_mm256_add_epi32(x1, half), 16);
      ST(d0 + (size_t)(j / 8) * NPAIR * 16, _mm256_or_si256(l0, _mm256_slli_epi32(l1, 16)));
      ST(d1 + (size_t)(j / 8) * NPAIR * 16, _mm256_or_si256(h0, _mm256_slli_epi32(h1, 16)));
    }
  }
}

/* correct for any n, just slow */
AVX2 static void generic(int n, const int32_t *A, const int32_t *B, int32_t *C) {
  for (int i = 0; i < n; i++)
    for (int j = 0; j < n; j++) {
      uint32_t s = 0;
      for (int k = 0; k < n; k++) s += (uint32_t)A[(size_t)i * n + k] * (uint32_t)B[(size_t)k * n + j];
      C[(size_t)i * n + j] = (int32_t)s;
    }
}

AVX2 static inline void block_mikro(int ic, int jc, int n, int32_t *C) {
  for (int r = 0; r < MC; r += MR)
    for (int jg = 0; jg < NC; jg += NR)
      mikro(NPAIR, Apk + (size_t)r * NPAIR * 16, Bpk + (size_t)(jg / 8) * NPAIR * 16,
            C + (size_t)(ic + r) * n + jc + jg, n);
}

AVX2 void matrix_multiply(int n, const int32_t *A, const int32_t *B, int32_t *C) {
  if ((n % MC) || (n % NC) || (n % KC) || (MC % MR) || (NC % NR)) { generic(n, A, B, C); return; }
  memset(C, 0, (size_t)n * n * sizeof(int32_t));
  for (int jc = 0; jc < n; jc += NC) {
    for (int kk = 0; kk < n; kk += KC) {
      packB(jc, kk, n, B);
      for (int ic = 0; ic < n; ic += MC) {
        packA(ic, kk, n, A);
        block_mikro(ic, jc, n, C);
      }
    }
  }
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #151.333 ms4 MB + 336 KBAcceptedScore: 100


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