提交记录 87838


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 mmms4k. 测测你的短整数矩阵乘法-4k Accepted 100 954.976 ms 207004 KB C++17 9.30 KB
提交时间 评测时间
2026-09-25 03:30:12 2026-09-25 03:30:15
// mmms* -- C = A*B, short, exact mod 2^16, Strassen on top of a packed base-case GEMM.
// The identity used is exact in Z/2^16, so recursion is valid for arbitrary short data.
#include <immintrin.h>
#include <stddef.h>
#include <stdlib.h>
#pragma GCC push_options
#pragma GCC target("avx2")

typedef short TE;

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

#define LDV(p) _mm256_loadu_si256((const __m256i *)(p))
#define STV(p, v) _mm256_storeu_si256((__m256i *)(p), (v))

static TE Apanel[MR * KC * 16 + 64] __attribute__((aligned(64)));
static TE Bpanel[(size_t)NC * KC + 128] __attribute__((aligned(64)));

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

// ---------------- micro kernel: MR x NR, A panel pre-replicated 16x ----------------
#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"
#define MICRO_CLOB "ymm12", "ymm13", "ymm14", "ymm15", "cc", "memory"

template <int OP>   // 0 = store, 1 = add, 2 = sub
static inline void micro(long k, const TE *ap, const TE *bp, TE *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)) {                                                \
    TE *p0 = C + (size_t)(r) * ldc;                                \
    if (OP == 0) { STV(p0, v0); STV(p0 + 16, v1); }                \
    else if (OP == 1) { STV(p0, _mm256_add_epi16(LDV(p0), v0)); STV(p0 + 16, _mm256_add_epi16(LDV(p0 + 16), v1)); } \
    else { STV(p0, _mm256_sub_epi16(LDV(p0), v0)); STV(p0 + 16, _mm256_sub_epi16(LDV(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
}

// ---------------- base-case GEMM: C = A*B (lda/ldb arbitrary), n<=NC ----------------
static void leaf(TE *C, int ldc, const TE *A, int lda, const TE *B, int ldb, int n) {
  const int nbc = n - (n % NR);
  for (int jc = 0; jc < nbc; jc += NC) {
    const int nc = (nbc - jc < NC) ? (nbc - jc) : NC;
    for (int pc = 0; pc < n; pc += KC) {
      const int kc = (n - pc < KC) ? (n - pc) : KC;
      for (int jr = 0, jb = 0; jr < nc; jr += NR, jb++) {
        TE *q = Bpanel + (size_t)jb * kc * NR;
        for (int k = 0; k < kc; k++) {
          const TE *src = B + (size_t)(pc + k) * ldb + jc + jr;
          for (int j = 0; j < NR; j++) q[j] = src[j];
          q += NR;
        }
      }
      const int firstpc = (pc == 0);
      for (int ic = 0; ic < n; ic += MR) {
        const int rows = (n - ic < MR) ? (n - ic) : MR;
        TE *p = Apanel;
        for (int k = 0; k < kc; k++) {
          for (int r = 0; r < rows; r++) { STV(p, b16(A + (size_t)(ic + r) * lda + pc + k)); p += 16; }
          for (int r = rows; r < MR; r++) { STV(p, _mm256_setzero_si256()); p += 16; }
        }
        const TE *b0 = Bpanel;
        TE *cp = C + (size_t)ic * ldc + jc;
        for (int jr = 0; jr < nc; jr += NR) {
          if (firstpc) micro<0>(kc, Apanel, b0, cp + jr, ldc, rows);
          else micro<1>(kc, Apanel, b0, cp + jr, ldc, rows);
          b0 += (size_t)kc * NR;
        }
      }
    }
  }
  if (nbc < n) {   // ragged columns (not reachable for our power-of-two sizes)
    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 * lda + k] * (unsigned)(unsigned short)B[(size_t)k * ldb + j];
        C[(size_t)i * ldc + j] = (TE)acc;
      }
  }
}

// ---------------- Strassen driver ----------------
static TE *g_pool;
static size_t g_off;

// form the 5 A-sums and 5 B-sums of an h-split block
static void mk5(TE *AS, TE *BS, const TE *A, int lda, const TE *B, int ldb, int h) {
  const size_t q = (size_t)h * h;
  const TE *a11 = A, *a12 = A + h, *a21 = A + (size_t)h * lda, *a22 = a21 + h;
  const TE *b11 = B, *b12 = B + h, *b21 = B + (size_t)h * ldb, *b22 = b21 + h;
  for (int i = 0; i < h; i++) {
    const TE *r11 = a11 + (size_t)i * lda, *r12 = a12 + (size_t)i * lda;
    const TE *r21 = a21 + (size_t)i * lda, *r22 = a22 + (size_t)i * lda;
    const TE *s11 = b11 + (size_t)i * ldb, *s12 = b12 + (size_t)i * ldb;
    const TE *s21 = b21 + (size_t)i * ldb, *s22 = b22 + (size_t)i * ldb;
    TE *o0 = AS + (size_t)i * h, *o1 = o0 + q, *o2 = o1 + q, *o3 = o2 + q, *o4 = o3 + q;
    TE *p0 = BS + (size_t)i * h, *p1 = p0 + q, *p2 = p1 + q, *p3 = p2 + q, *p4 = p3 + q;
    int j = 0;
    for (; j + 16 <= h; j += 16) {
      __m256i v11 = LDV(r11 + j), v12 = LDV(r12 + j), v21 = LDV(r21 + j), v22 = LDV(r22 + j);
      __m256i w11 = LDV(s11 + j), w12 = LDV(s12 + j), w21 = LDV(s21 + j), w22 = LDV(s22 + j);
      STV(o0 + j, _mm256_add_epi16(v11, v22));      // A11+A22
      STV(o1 + j, _mm256_add_epi16(v21, v22));      // A21+A22
      STV(o2 + j, _mm256_add_epi16(v11, v12));      // A11+A12
      STV(o3 + j, _mm256_sub_epi16(v21, v11));      // A21-A11
      STV(o4 + j, _mm256_sub_epi16(v12, v22));      // A12-A22
      STV(p0 + j, _mm256_add_epi16(w11, w22));      // B11+B22
      STV(p1 + j, _mm256_sub_epi16(w12, w22));      // B12-B22
      STV(p2 + j, _mm256_sub_epi16(w21, w11));      // B21-B11
      STV(p3 + j, _mm256_add_epi16(w11, w12));      // B11+B12
      STV(p4 + j, _mm256_add_epi16(w21, w22));      // B21+B22
    }
    for (; j < h; j++) {
      TE x11 = r11[j], x12 = r12[j], x21 = r21[j], x22 = r22[j];
      TE y11 = s11[j], y12 = s12[j], y21 = s21[j], y22 = s22[j];
      o0[j] = x11 + x22; o1[j] = x21 + x22; o2[j] = x11 + x12;
      o3[j] = x21 - x11; o4[j] = x12 - x22;
      p0[j] = y11 + y22; p1[j] = y12 - y22; p2[j] = y21 - y11;
      p3[j] = y11 + y12; p4[j] = y21 + y22;
    }
  }
}

static void combine(TE *C, int ldc, const TE *M1, const TE *M2, const TE *M3,
                    const TE *M4, const TE *M5, const TE *M6, const TE *M7, int h) {
  const size_t q = (size_t)h * h;
  for (int i = 0; i < h; i++) {
    const TE *m1 = M1 + (size_t)i * h, *m2 = M2 + (size_t)i * h, *m3 = M3 + (size_t)i * h;
    const TE *m4 = M4 + (size_t)i * h, *m5 = M5 + (size_t)i * h, *m6 = M6 + (size_t)i * h;
    const TE *m7 = M7 + (size_t)i * h;
    TE *c11 = C + (size_t)i * ldc, *c12 = c11 + h;
    TE *c21 = c11 + (size_t)h * ldc, *c22 = c21 + h;
    int j = 0;
    for (; j + 16 <= h; j += 16) {
      __m256i a1 = LDV(m1 + j), a2 = LDV(m2 + j), a3 = LDV(m3 + j), a4 = LDV(m4 + j);
      __m256i a5 = LDV(m5 + j), a6 = LDV(m6 + j), a7 = LDV(m7 + j);
      // C11 = M1+M4-M5+M7   C12 = M3+M5   C21 = M2+M4   C22 = M1-M2+M3+M6
      STV(c11 + j, _mm256_add_epi16(_mm256_sub_epi16(_mm256_add_epi16(a1, a4), a5), a7));
      STV(c12 + j, _mm256_add_epi16(a3, a5));
      STV(c21 + j, _mm256_add_epi16(a2, a4));
      STV(c22 + j, _mm256_add_epi16(_mm256_sub_epi16(_mm256_add_epi16(a1, a3), a2), a6));
    }
    for (; j < h; j++) {
      c11[j] = m1[j] + m4[j] - m5[j] + m7[j];
      c12[j] = m3[j] + m5[j];
      c21[j] = m2[j] + m4[j];
      c22[j] = m1[j] - m2[j] + m3[j] + m6[j];
    }
  }
  (void)q;
}

static void mm(TE *C, int ldc, const TE *A, int lda, const TE *B, int ldb, int n) {
  if (n <= CUT) { leaf(C, ldc, A, lda, B, ldb, n); return; }
  const int h = n >> 1;
  const size_t q = (size_t)h * h;
  TE *AS = g_pool + g_off;
  TE *BS = AS + 5 * q;
  TE *M = BS + 5 * q;
  g_off += 17 * q;
  mk5(AS, BS, A, lda, B, ldb, h);
  const TE *A11 = A, *A22 = A + (size_t)h * lda + h;
  const TE *B11 = B, *B22 = B + (size_t)h * ldb + h;
  mm(M + 0 * q, h, AS + 0 * q, h, BS + 0 * q, h, h);   // M1 = (A11+A22)(B11+B22)
  mm(M + 1 * q, h, AS + 1 * q, h, B11, ldb, h);        // M2 = (A21+A22)B11
  mm(M + 2 * q, h, A11, lda, BS + 1 * q, h, h);        // M3 = A11(B12-B22)
  mm(M + 3 * q, h, A22, lda, BS + 2 * q, h, h);        // M4 = A22(B21-B11)
  mm(M + 4 * q, h, AS + 2 * q, h, B22, ldb, h);        // M5 = (A11+A12)B22
  mm(M + 5 * q, h, AS + 3 * q, h, BS + 3 * q, h, h);   // M6 = (A21-A11)(B11+B12)
  mm(M + 6 * q, h, AS + 4 * q, h, BS + 4 * q, h, h);   // M7 = (A12-A22)(B21+B22)
  combine(C, ldc, M, M + q, M + 2 * q, M + 3 * q, M + 4 * q, M + 5 * q, M + 6 * q, h);
  g_off -= 17 * q;
}

void matrix_multiply(int n, const short *A, const short *B, short *C) {
  size_t need = 64;
  for (int m = n; m > CUT; m >>= 1) need += 17 * (size_t)(m >> 1) * (m >> 1);
  if (need > 64) {
    void *p = 0;
    if (posix_memalign(&p, 4096, need * sizeof(TE))) return;
    g_pool = (TE *)p;
    g_off = 0;
    mm(C, n, A, n, B, n, n);
    free(p);
  } else {
    leaf(C, n, A, n, B, n, n);
  }
}
#pragma GCC pop_options

CompilationN/AN/ACompile OKScore: N/A

Testcase #1954.976 ms202 MB + 156 KBAcceptedScore: 100


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