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