/* 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);
}
}
}
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 51.333 ms | 4 MB + 336 KB | Accepted | Score: 100 | 显示更多 |