提交记录 86708


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 mmmc4k. 测测你的8位整数矩阵乘法-4k Accepted 100 1.653 s 16680 KB C++17 4.36 KB
提交时间 评测时间
2026-09-25 01:06:20 2026-09-25 01:06:24
#include <immintrin.h>
#include <string.h>
#include <stdint.h>

#ifndef MC
#define MC 64
#endif
#ifndef KC
#define KC 256
#endif
#ifndef NC
#define NC 512
#endif

#define MR 4
#define NR 32
#define APAD 8
#define AST (KC + APAD)

static int16_t Abuf[MC * AST + 64];
static int16_t Bbuf[KC * NC + 64];

#define AVX2 __attribute__((target("avx2")))

/* 1-uop broadcast of an int16 from memory (port5); _mm256_set1_epi16 makes gcc
   emit movzwl+vmovd+vpbroadcastw (3 uops + 2 extra p5 uops) instead. */
AVX2 static inline __m256i bcastw(const int16_t *p) {
  __m256i v;
  __asm__("vpbroadcastw %1, %0" : "=x"(v) : "m"(*p));
  return v;
}

AVX2 static inline void mikro(int kc, const int16_t *ap, const int16_t *bp, int8_t *c, int ldc) {
  __m256i c00 = _mm256_setzero_si256(), c01 = _mm256_setzero_si256();
  __m256i c10 = _mm256_setzero_si256(), c11 = _mm256_setzero_si256();
  __m256i c20 = _mm256_setzero_si256(), c21 = _mm256_setzero_si256();
  __m256i c30 = _mm256_setzero_si256(), c31 = _mm256_setzero_si256();
  const int16_t *a = ap;
  for (int k = 0; k < kc; k++) {
    __m256i b0 = _mm256_loadu_si256((const __m256i *)(bp));
    __m256i b1 = _mm256_loadu_si256((const __m256i *)(bp + 16));
    bp += NC;
    __m256i a0 = bcastw(a);
    __m256i a1 = bcastw(a + AST);
    __m256i a2 = bcastw(a + 2 * AST);
    __m256i a3 = bcastw(a + 3 * AST);
    a++;
    c00 = _mm256_add_epi16(c00, _mm256_mullo_epi16(a0, b0));
    c01 = _mm256_add_epi16(c01, _mm256_mullo_epi16(a0, b1));
    c10 = _mm256_add_epi16(c10, _mm256_mullo_epi16(a1, b0));
    c11 = _mm256_add_epi16(c11, _mm256_mullo_epi16(a1, b1));
    c20 = _mm256_add_epi16(c20, _mm256_mullo_epi16(a2, b0));
    c21 = _mm256_add_epi16(c21, _mm256_mullo_epi16(a2, b1));
    c30 = _mm256_add_epi16(c30, _mm256_mullo_epi16(a3, b0));
    c31 = _mm256_add_epi16(c31, _mm256_mullo_epi16(a3, b1));
  }
  const __m256i msk = _mm256_set1_epi16(0x00FF);
  __m256i pa = _mm256_permute4x64_epi64(_mm256_packus_epi16(_mm256_and_si256(c00, msk), _mm256_and_si256(c01, msk)), 0xD8);
  __m256i pb = _mm256_permute4x64_epi64(_mm256_packus_epi16(_mm256_and_si256(c10, msk), _mm256_and_si256(c11, msk)), 0xD8);
  __m256i pc = _mm256_permute4x64_epi64(_mm256_packus_epi16(_mm256_and_si256(c20, msk), _mm256_and_si256(c21, msk)), 0xD8);
  __m256i pd = _mm256_permute4x64_epi64(_mm256_packus_epi16(_mm256_and_si256(c30, msk), _mm256_and_si256(c31, msk)), 0xD8);
  int8_t *p;
  p = c;           _mm256_storeu_si256((__m256i *)p, _mm256_add_epi8(_mm256_loadu_si256((const __m256i *)p), pa));
  p = c + ldc;     _mm256_storeu_si256((__m256i *)p, _mm256_add_epi8(_mm256_loadu_si256((const __m256i *)p), pb));
  p = c + 2 * ldc; _mm256_storeu_si256((__m256i *)p, _mm256_add_epi8(_mm256_loadu_si256((const __m256i *)p), pc));
  p = c + 3 * ldc; _mm256_storeu_si256((__m256i *)p, _mm256_add_epi8(_mm256_loadu_si256((const __m256i *)p), pd));
}

AVX2 static inline void widen(const int8_t *src, int16_t *dst, int kc) {
  int k = 0;
  for (; k + 16 <= kc; k += 16)
    _mm256_storeu_si256((__m256i *)(dst + k), _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i *)(src + k))));
  for (; k < kc; k++) dst[k] = src[k];
}

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

AVX2 void matrix_multiply(int n, const int8_t *A, const int8_t *B, int8_t *C) {
  if ((n % MC) || (n % NC) || (n % KC) || (NC % NR) || (MC % MR)) { generic(n, A, B, C); return; }
  memset(C, 0, (size_t)n * n);
  for (int jc = 0; jc < n; jc += NC) {
    int nc = n - jc; if (nc > NC) nc = NC;
    for (int kk = 0; kk < n; kk += KC) {
      int kc = n - kk; if (kc > KC) kc = KC;
      for (int k = 0; k < kc; k++)
        widen(B + (size_t)(kk + k) * n + jc, Bbuf + (size_t)k * NC, nc);
      for (int ic = 0; ic < n; ic += MC) {
        int mc = n - ic; if (mc > MC) mc = MC;
        for (int r = 0; r < mc; r++)
          widen(A + (size_t)(ic + r) * n + kk, Abuf + (size_t)r * AST, kc);
        for (int jr = 0; jr < nc; jr += NR) {
          const int16_t *bp = Bbuf + jr;
          if (nc - jr >= NR)
            for (int i = 0; i + MR <= mc; i += MR)
              mikro(kc, Abuf + (size_t)i * AST, bp, C + (size_t)(ic + i) * n + jc + jr, n);
        }
      }
    }
  }
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #11.653 s16 MB + 296 KBAcceptedScore: 100


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