// 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")
#define MS_REP 1
#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
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 18.208 ms | 2 MB + 148 KB | Accepted | Score: 100 | 显示更多 |