// ===== REFERENCES =====
// [1] duck.ac 用户 saffah_cc_v41_agg1(本账号),提交 #105354 <https://duck.ac/submission/105354>
// (文件 work/fs1_sub5.cpp,14.390660 ms = 本题现役最好件)——**本文件的正文基座**。
// 本件在其上只做"思路"里那一处改动(累加面的表示与布局),其余一字未动。
// [2] 本账号提交 #105282 / #104948 / #104302 / #103973 / #103927 / #103535 / #103476 /
// #103032 / #102993 —— int16 `vpmaddwd` + 一层 Strassen(n=1024→512) 引擎谱系
// (quantize / pack_a / pack_b(tile-4) / mikro / mikro6 / combine / MC=192)。
// [3] 兄弟 lane `f1a_` 的判题机取证:对手 #119584 的 −1.341% = "搬动的字节数"这一轴
// (int8 混合线),⇒ **本题的杠杆 = 字节**;本件即沿该轴把**累加面的字节减半**。
// [4] 兄弟 lane `f1c_` 的精度证书(work/f1c_tmp/f1c_prec.c,**同账号车队内部件**):本题误差预算充裕
// (现役只用 ~10% tau)⇒ 允许在"面"这一级引入有界舍入。
// [5] **无外部引用**:本文件正文 = 本账号自有件 #105354 的逐字节副本 + 本席自作的五处机械改动
// (int16 累加面的 12 位移位表示 + 行对交错布局),未复制任何其他账号的代码;
// 上面 [3][4] 引的是**同账号车队姊妹 lane 的实测结论/精度数据(思想与数据,非代码)**。
// ======================
// ===== 思路 =====
// 【单变量:把 7 个 Strassen 累加面从 int32 换成**带移位的 int16**,并让 mikro 的存数成为整条 64 B line】
//
// 一、机理(字节账)
// 面里的值 |M_i| 实测 ~7e7 量级(题面保证 A/B 均匀分布),而最终容差折算到量化单位是
// 3.93e5 ⇒ 面只需要约 19 位有效。把面右移 12 位(四舍五入)存成 int16 ⇒ 面字节数减半:
// * combine 读:7 面 × 512² × 4 B = 7.34 MB → 3.67 MB
// * prod 首触:5 MB → 2.5 MB(判题机 0.204 Mcyc/MB)
// 精度:每元素舍入误差 U(±2048),7 面求和 σ = 3128 量化单位 = 0.0030(容差 0.375)⇒ +0.8% 预算。
// 溢出:|M_i|max ≈ 4.3σ = 6.8e7 ⇒ /4096 = 1.66e4 < 32767 ✓(2 倍余量;vpackssdw 仍带饱和兜底)。
//
// 二、★ 布局(本件的关键;第一版只改面宽不改布局,判题机实测打平)
// int16 面若沿用"行距 1024 B"的朴素布局,mikro 每次每行只写 32 B = **半条 line**:
// * NT 存半行 ⇒ mikro 相位 **+11 Mcyc**(灾难);WB 存半行 ⇒ RFO×2 ⇒ mikro +0.7 Mcyc,
// 正好吃掉 combine 省下的 0.65 Mcyc(判题机单臂冷跑 48.715 vs 基座 48.665,打平)。
// ⇒ 改布局:**面的最小块 = 64 B = 同一列块下相邻两行各 16 个 int16**:
// flat(p,c,h,k) = ((p*32 + c)*2 + h)*16 + k p=行对 0..255, c=列块 0..31, h=0/1
// ⇒ mikro 的 6 行 = 3 个行对 ⇒ 6 条 32 B 存数落在 {0,32, 2048,2080, 4096,4128},
// 相邻两条合成一条完整 64 B line(与基座 int32 版同形:base 也是两两相邻的 32 B 存数)✓
// ⇒ combine 读同一块:行 i 取 h=(i&1),同一条 line 的另一半正是相邻输出行 i±1 的数据,
// 在相邻迭代被读到(每迭代只搬 ~11 KB)⇒ 仍命中 L1/L2 ⇒ DRAM 流量 = 每 32 列 1 条 line ✓
//
// 三、实现(五处,全部机械)
// 1. `prod[5][H*H+6144] int32` → `prod[5][H*H] int16`(5 MB → 2.5 MB)
// 2. mikro6 尾声:12 条 vmovntdq(每行 2×32 B) → 每行 vpaddd(+2048 舍入) + vpsrad $12 + vpackssdw
// + vpermq $0xD8 + 1 条 vmovntdq(32 B) ⇒ 6 条存数、位移 {0,32,2048,2080,4096,4128}
// 3. `mikro`(4x16 intrinsics) 尾声同样处理,位移表 f1b_off4(**元素单位** {0,16,1024,1040})
// 4. product()/combine() 的面指针 int32_t* → int16_t*,tile 基址改为 ((ic+ir)>>1)*32+(jr>>4)
// 5. combine:每趟 16 列(1 ymm int16/面),vpmovsxwd 展宽到 int32 后做同样的和差树;
// 比例常数 = 2^12/1024^2 = 1/256;**i 循环倒序**(读先于覆写),i<2 时读 2 KB 预拷贝缓冲
//
// 四、闸门
// 本地 f1_drv(双精度参考、tol=3n²·FLT_EPS=0.375)三个种子:
// maxdiff = 0.038566@(621,491) / 0.037501@(504,288) / 0.035487@(688,500)(全 PASS)
// = 与"仅改面宽"版**逐位相同**(布局改动不改任何数值);确定性:同进程两次调用 C1vsC2 = 0。
// 判题机 FNV chk = 0x790a298efefcc6b3(与 i16p 版逐位同)。
//
// 五、判题机实测(立即体验单臂**冷跑**,同窗口交叉配对;基座自身重复性 ±0.02%)
// base 48.646 / 48.663 Mcyc
// i16q 48.217 / 48.232 Mcyc ⇒ **−0.88%(−429 Kcyc,~20σ)**
// 对照:同布局改 WB 存数(面留 L3)48.422 / 48.447 ⇒ 比 NT 版差 0.43% ⇒ 采用 NT。
// 多臂暖台架相位表(同二进制):combine 相位 −260 Kcyc;总时间在该台架上 ≈0
// (本族已知:改动"MB 级 scratch 的 L3 驻留"时多臂台架不可读总时间符号,只读相位)。
// ⇒ 提交预测 ≈ 14.390660 × (48221/48646) = **14.264 ms**(严支线 14.111243 ⇒ 本件单独**不够**,
// 需与 f1c_ 的 int8 前缀段叠加;故本件未提交,按派单规则报协调者)。
// ================
#pragma GCC optimize("O3,unroll-loops")
#pragma GCC target("avx2,bmi,bmi2,popcnt,lzcnt,fma")
#include <immintrin.h>
#include <stdint.h>
#include <stddef.h>
#include <string.h>
#define N 1024
#define H 512
#define KC H
#define MC 192
#define NR 16
#define MR 4
#define NP (KC / 2)
/* qa/qb 不再占 .bss:它们在 matrix_multiply 里指向 C 的缓冲区(见【思路】)。
C 只在最后一步 combine() 才被写,qa/qb 的最后使用者在最后一个 product() 里 ⇒ 生命周期不重叠。 */
static int16_t *qa, *qb;
alignas(4096) static int16_t ap[MC*KC], bp_storage[H*KC+2048];
static int16_t *const bp = bp_storage + 1536;
alignas(4096) static int16_t prod[5][H*H];
static const int32_t f1b_rndv[8] = {1<<11,1<<11,1<<11,1<<11,1<<11,1<<11,1<<11,1<<11};
static int16_t f1b_r0[2][2*H];
static const size_t f1b_off6[6] = {0,32,2048,2080,4096,4128}; /* bytes, used by the asm kernel */
static const size_t f1b_off4[4] = {0,16,1024,1040}; /* ELEMENTS (int16): c+f1b_off4[R] */ /* plane row 0 of the two C-resident planes (pre-copy) */
static inline __m256i ldu(const void *p) {
__m256i v;
__asm__("vmovdqu %1,%0" : "=x"(v) : "m"(*(const __m256i *)p));
return v;
}
static inline __m256i bc32(const void *p) {
__m256i v;
__asm__("vpbroadcastd %1,%0" : "=x"(v) : "m"(*(const int *)p));
return v;
}
static void quantize(const float *src, int16_t *dst) {
const __m256 scale = _mm256_set1_ps(1024.0f);
for (int i = 0; i < N*N; i += 16) {
__m256i x = _mm256_cvtps_epi32(_mm256_mul_ps(_mm256_load_ps(src+i), scale));
__m256i y = _mm256_cvtps_epi32(_mm256_mul_ps(_mm256_load_ps(src+i+8), scale));
_mm256_store_si256((__m256i *)(dst+i),
_mm256_permute4x64_epi64(_mm256_packs_epi32(x,y),0xD8));
}
}
struct Mat { const int16_t *x; const int16_t *y; int sign; };
static inline __m256i vec(const int16_t *x, const int16_t *y, int sign) {
__m256i v = _mm256_load_si256((const __m256i *)x);
if (y) {
__m256i w = _mm256_load_si256((const __m256i *)y);
v = sign > 0 ? _mm256_add_epi16(v,w) : _mm256_sub_epi16(v,w);
}
return v;
}
static void pack_a(Mat a, int ic, int mc) {
for (int r = 0; r < mc; r++) {
const int16_t *x = a.x + (size_t)(ic+r)*N;
const int16_t *y = a.y ? a.y + (size_t)(ic+r)*N : nullptr;
int16_t *d = ap + (size_t)r*KC;
for (int k = 0; k < KC; k += 16)
_mm256_store_si256((__m256i *)(d+k), vec(x+k,y?y+k:nullptr,a.sign));
}
}
static void pack_b(Mat b) {
for (int pb0=0; pb0<NP; pb0+=4) {
for (int jb=0; jb<H; jb+=NR) {
int16_t *dblk = bp + (size_t)(jb/NR)*KC*NR;
for (int p=pb0; p<pb0+4; ++p) {
const int16_t *x0 = b.x + (size_t)(2*p)*N + jb;
const int16_t *x1 = x0 + N;
const int16_t *y0 = b.y ? b.y + (size_t)(2*p)*N + jb : nullptr;
const int16_t *y1 = y0 ? y0 + N : nullptr;
__m256i x = vec(x0,y0,b.sign), y = vec(x1,y1,b.sign);
__m256i lo = _mm256_unpacklo_epi16(x,y), hi = _mm256_unpackhi_epi16(x,y);
int16_t *d = dblk + (size_t)p*NR*2;
_mm256_store_si256((__m256i *)d, _mm256_permute2x128_si256(lo,hi,0x20));
_mm256_store_si256((__m256i *)(d+16), _mm256_permute2x128_si256(lo,hi,0x31));
}
}
}
}
static inline void mikro(const int16_t *a, const int16_t *b, int16_t *c) {
const __m256i rnd = _mm256_set1_epi32(1<<11);
__m256i c00,c01,c10,c11,c20,c21,c30,c31;
c00=c01=c10=c11=c20=c21=c30=c31=_mm256_setzero_si256();
const char *p0=(const char *)a, *p1=(const char *)(a+KC),
*p2=(const char *)(a+2*KC), *p3=(const char *)(a+3*KC);
const char *pb=(const char *)b;
for (int p=0; p<NP; p++) {
__m256i b0=ldu(pb), b1=ldu(pb+32), av;
av=bc32(p0); c00=_mm256_add_epi32(c00,_mm256_madd_epi16(av,b0));
c01=_mm256_add_epi32(c01,_mm256_madd_epi16(av,b1));
av=bc32(p1); c10=_mm256_add_epi32(c10,_mm256_madd_epi16(av,b0));
c11=_mm256_add_epi32(c11,_mm256_madd_epi16(av,b1));
av=bc32(p2); c20=_mm256_add_epi32(c20,_mm256_madd_epi16(av,b0));
c21=_mm256_add_epi32(c21,_mm256_madd_epi16(av,b1));
av=bc32(p3); c30=_mm256_add_epi32(c30,_mm256_madd_epi16(av,b0));
c31=_mm256_add_epi32(c31,_mm256_madd_epi16(av,b1));
p0+=4;p1+=4;p2+=4;p3+=4;pb+=64;
}
#define ST(R,X,Y) do {\
_mm256_stream_si256((__m256i *)(c+f1b_off4[R]),\
_mm256_permute4x64_epi64(_mm256_packs_epi32(\
_mm256_srai_epi32(_mm256_add_epi32(X,rnd),12),\
_mm256_srai_epi32(_mm256_add_epi32(Y,rnd),12)),0xD8));\
} while(0)
ST(0,c00,c01); ST(1,c10,c11); ST(2,c20,c21); ST(3,c30,c31);
#undef ST
}
/* 6 rows x 16 cols over all KC k -- same panel layout, 12 accumulators (gcc 9.3 keeps them
in ymm: 12 acc + 2 b + 1 broadcast = 15 <= 16). Measured on the judge (work/f1_mikro4.cpp,
same binary, floor 0.10%): 48.54 vs 49.11 cyc per 1024 MAC for the 4x16 shape = -1.16%. */
static inline void mikro6(const int16_t *a, const int16_t *b, int16_t *c) {
const char *pa=(const char *)a, *pb2=(const char *)b;
int k=NP-4;
__asm__ __volatile__(
"vmovdqu 0(%[b]), %%ymm14\n\t"
"vmovdqu 32(%[b]), %%ymm15\n\t"
"vpbroadcastd 0(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm0\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm1\n\t"
"vpbroadcastd 1024(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm2\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm3\n\t"
"vpbroadcastd 2048(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm4\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm5\n\t"
"vpbroadcastd 3072(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm6\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm7\n\t"
"vpbroadcastd 4096(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm8\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm9\n\t"
"vpbroadcastd 5120(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm10\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm11\n\t"
"vmovdqu 64(%[b]), %%ymm14\n\t"
"vmovdqu 96(%[b]), %%ymm15\n\t"
"vpbroadcastd 4(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 1028(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 2052(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 3076(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm13, %%ymm7, %%ymm7\n\t"
"vpbroadcastd 4100(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm8, %%ymm8\n\t"
"vpaddd %%ymm13, %%ymm9, %%ymm9\n\t"
"vpbroadcastd 5124(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm10, %%ymm10\n\t"
"vpaddd %%ymm13, %%ymm11, %%ymm11\n\t"
"vmovdqu 128(%[b]), %%ymm14\n\t"
"vmovdqu 160(%[b]), %%ymm15\n\t"
"vpbroadcastd 8(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 1032(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 2056(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 3080(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm13, %%ymm7, %%ymm7\n\t"
"vpbroadcastd 4104(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm8, %%ymm8\n\t"
"vpaddd %%ymm13, %%ymm9, %%ymm9\n\t"
"vpbroadcastd 5128(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm10, %%ymm10\n\t"
"vpaddd %%ymm13, %%ymm11, %%ymm11\n\t"
"vmovdqu 192(%[b]), %%ymm14\n\t"
"vmovdqu 224(%[b]), %%ymm15\n\t"
"vpbroadcastd 12(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 1036(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 2060(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 3084(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm13, %%ymm7, %%ymm7\n\t"
"vpbroadcastd 4108(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm8, %%ymm8\n\t"
"vpaddd %%ymm13, %%ymm9, %%ymm9\n\t"
"vpbroadcastd 5132(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm10, %%ymm10\n\t"
"vpaddd %%ymm13, %%ymm11, %%ymm11\n\t"
"add $16, %[a]\n\t"
"add $256, %[b]\n\t"
".p2align 5\n\t"
"1:\n\t"
"vmovdqu 0(%[b]), %%ymm14\n\t"
"vmovdqu 32(%[b]), %%ymm15\n\t"
"vpbroadcastd 0(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 1024(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 2048(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 3072(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm13, %%ymm7, %%ymm7\n\t"
"vpbroadcastd 4096(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm8, %%ymm8\n\t"
"vpaddd %%ymm13, %%ymm9, %%ymm9\n\t"
"vpbroadcastd 5120(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm10, %%ymm10\n\t"
"vpaddd %%ymm13, %%ymm11, %%ymm11\n\t"
"vmovdqu 64(%[b]), %%ymm14\n\t"
"vmovdqu 96(%[b]), %%ymm15\n\t"
"vpbroadcastd 4(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 1028(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 2052(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 3076(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm13, %%ymm7, %%ymm7\n\t"
"vpbroadcastd 4100(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm8, %%ymm8\n\t"
"vpaddd %%ymm13, %%ymm9, %%ymm9\n\t"
"vpbroadcastd 5124(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm10, %%ymm10\n\t"
"vpaddd %%ymm13, %%ymm11, %%ymm11\n\t"
"vmovdqu 128(%[b]), %%ymm14\n\t"
"vmovdqu 160(%[b]), %%ymm15\n\t"
"vpbroadcastd 8(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 1032(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 2056(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 3080(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm13, %%ymm7, %%ymm7\n\t"
"vpbroadcastd 4104(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm8, %%ymm8\n\t"
"vpaddd %%ymm13, %%ymm9, %%ymm9\n\t"
"vpbroadcastd 5128(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm10, %%ymm10\n\t"
"vpaddd %%ymm13, %%ymm11, %%ymm11\n\t"
"vmovdqu 192(%[b]), %%ymm14\n\t"
"vmovdqu 224(%[b]), %%ymm15\n\t"
"vpbroadcastd 12(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm0, %%ymm0\n\t"
"vpaddd %%ymm13, %%ymm1, %%ymm1\n\t"
"vpbroadcastd 1036(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
"vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
"vpbroadcastd 2060(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
"vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
"vpbroadcastd 3084(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm6, %%ymm6\n\t"
"vpaddd %%ymm13, %%ymm7, %%ymm7\n\t"
"vpbroadcastd 4108(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm8, %%ymm8\n\t"
"vpaddd %%ymm13, %%ymm9, %%ymm9\n\t"
"vpbroadcastd 5132(%[a]), %%ymm13\n\t"
"vpmaddwd %%ymm14, %%ymm13, %%ymm12\n\t"
"vpmaddwd %%ymm15, %%ymm13, %%ymm13\n\t"
"vpaddd %%ymm12, %%ymm10, %%ymm10\n\t"
"vpaddd %%ymm13, %%ymm11, %%ymm11\n\t"
"add $16, %[a]\n\t"
"add $256, %[b]\n\t"
"sub $4, %[k]\n\t"
"jnz 1b\n\t"
"vmovdqu (%[r]), %%ymm15\n\t"
"vpaddd %%ymm15, %%ymm0, %%ymm12\n\t"
"vpsrad $12, %%ymm12, %%ymm12\n\t"
"vpaddd %%ymm15, %%ymm1, %%ymm13\n\t"
"vpsrad $12, %%ymm13, %%ymm13\n\t"
"vpackssdw %%ymm13, %%ymm12, %%ymm12\n\t"
"vpermq $216, %%ymm12, %%ymm12\n\t"
"vmovntdq %%ymm12, 0(%[c])\n\t"
"vpaddd %%ymm15, %%ymm2, %%ymm12\n\t"
"vpsrad $12, %%ymm12, %%ymm12\n\t"
"vpaddd %%ymm15, %%ymm3, %%ymm13\n\t"
"vpsrad $12, %%ymm13, %%ymm13\n\t"
"vpackssdw %%ymm13, %%ymm12, %%ymm12\n\t"
"vpermq $216, %%ymm12, %%ymm12\n\t"
"vmovntdq %%ymm12, 32(%[c])\n\t"
"vpaddd %%ymm15, %%ymm4, %%ymm12\n\t"
"vpsrad $12, %%ymm12, %%ymm12\n\t"
"vpaddd %%ymm15, %%ymm5, %%ymm13\n\t"
"vpsrad $12, %%ymm13, %%ymm13\n\t"
"vpackssdw %%ymm13, %%ymm12, %%ymm12\n\t"
"vpermq $216, %%ymm12, %%ymm12\n\t"
"vmovntdq %%ymm12, 2048(%[c])\n\t"
"vpaddd %%ymm15, %%ymm6, %%ymm12\n\t"
"vpsrad $12, %%ymm12, %%ymm12\n\t"
"vpaddd %%ymm15, %%ymm7, %%ymm13\n\t"
"vpsrad $12, %%ymm13, %%ymm13\n\t"
"vpackssdw %%ymm13, %%ymm12, %%ymm12\n\t"
"vpermq $216, %%ymm12, %%ymm12\n\t"
"vmovntdq %%ymm12, 2080(%[c])\n\t"
"vpaddd %%ymm15, %%ymm8, %%ymm12\n\t"
"vpsrad $12, %%ymm12, %%ymm12\n\t"
"vpaddd %%ymm15, %%ymm9, %%ymm13\n\t"
"vpsrad $12, %%ymm13, %%ymm13\n\t"
"vpackssdw %%ymm13, %%ymm12, %%ymm12\n\t"
"vpermq $216, %%ymm12, %%ymm12\n\t"
"vmovntdq %%ymm12, 4096(%[c])\n\t"
"vpaddd %%ymm15, %%ymm10, %%ymm12\n\t"
"vpsrad $12, %%ymm12, %%ymm12\n\t"
"vpaddd %%ymm15, %%ymm11, %%ymm13\n\t"
"vpsrad $12, %%ymm13, %%ymm13\n\t"
"vpackssdw %%ymm13, %%ymm12, %%ymm12\n\t"
"vpermq $216, %%ymm12, %%ymm12\n\t"
"vmovntdq %%ymm12, 4128(%[c])\n\t" : [a] "+&r"(pa), [b] "+&r"(pb2), [k] "+&r"(k)
: [c] "r"(c), [r] "r"(f1b_rndv)
: "memory", "cc", "ymm0","ymm1","ymm2","ymm3","ymm4","ymm5","ymm6","ymm7","ymm8","ymm9","ymm10","ymm11","ymm12","ymm13","ymm14","ymm15");
}
static void product(Mat a, Mat b, int16_t *out) {
pack_b(b);
for (int ic=0; ic<H; ic+=MC) {
int mc = H-ic < MC ? H-ic : MC;
pack_a(a,ic,mc);
for (int jr=0; jr<H; jr+=NR) {
const int16_t *pb = bp + (size_t)(jr/NR)*KC*NR;
int ir=0;
int six_rows = mc;
while (six_rows % 6 || (mc-six_rows) % 4) six_rows -= 2;
for (; ir<six_rows; ir+=6)
mikro6(ap+(size_t)ir*KC, pb, out+(size_t)(((ic+ir)>>1)*32+(jr>>4))*32);
for (; ir+4<=mc; ir+=4)
mikro(ap+(size_t)ir*KC, pb, out+(size_t)(((ic+ir)>>1)*32+(jr>>4))*32);
}
}
}
/* 32 output columns per iteration (was 8) => 28 independent loads in flight instead of 7.
Pure loop restructuring: the arithmetic and every NT store address are unchanged. */
/* 16 output columns per pass: one ymm of int16 per plane, widened with vpmovsxwd. */
#define CBH(HI, COFF) \
do { \
const __m128i h1 = HI ? _mm256_extracti128_si256(p1,1) : _mm256_castsi256_si128(p1); \
const __m128i h2 = HI ? _mm256_extracti128_si256(p2,1) : _mm256_castsi256_si128(p2); \
const __m128i h3 = HI ? _mm256_extracti128_si256(p3,1) : _mm256_castsi256_si128(p3); \
const __m128i h4 = HI ? _mm256_extracti128_si256(p4,1) : _mm256_castsi256_si128(p4); \
const __m128i h5 = HI ? _mm256_extracti128_si256(p5,1) : _mm256_castsi256_si128(p5); \
const __m128i h6 = HI ? _mm256_extracti128_si256(p6,1) : _mm256_castsi256_si128(p6); \
const __m128i h7 = HI ? _mm256_extracti128_si256(p7,1) : _mm256_castsi256_si128(p7); \
__m256i m1=_mm256_cvtepi16_epi32(h1), m2=_mm256_cvtepi16_epi32(h2); \
__m256i m3=_mm256_cvtepi16_epi32(h3), m4=_mm256_cvtepi16_epi32(h4); \
__m256i m5=_mm256_cvtepi16_epi32(h5), m6=_mm256_cvtepi16_epi32(h6); \
__m256i m7=_mm256_cvtepi16_epi32(h7); \
__m256i c11=_mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(m1,m4),m5),m7); \
__m256i c12=_mm256_add_epi32(m3,m5); \
__m256i c21=_mm256_add_epi32(m2,m4); \
__m256i c22=_mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(m1,m3),m2),m6); \
_mm256_stream_ps(C+(size_t)i*N+(COFF), _mm256_mul_ps(_mm256_cvtepi32_ps(c11),inv)); \
_mm256_stream_ps(C+(size_t)i*N+(COFF)+H, _mm256_mul_ps(_mm256_cvtepi32_ps(c12),inv)); \
_mm256_stream_ps(C+(size_t)(i+H)*N+(COFF), _mm256_mul_ps(_mm256_cvtepi32_ps(c21),inv)); \
_mm256_stream_ps(C+(size_t)(i+H)*N+(COFF)+H, _mm256_mul_ps(_mm256_cvtepi32_ps(c22),inv)); \
} while (0)
static void combine(float *C) {
_mm_sfence(); /* the only fence needed: before reading the NT-written planes */
const __m256 inv = _mm256_set1_ps(4096.0f/(1024.0f*1024.0f));
/* plane row 0 of the two C-resident planes (M2 in qa, M4 in qb) is overwritten by
this loop's own output on iteration i==0; save it first. */
memcpy(f1b_r0[0], qa, (size_t)H*4); /* plane rows 0,1 = the p=0 line pair */
memcpy(f1b_r0[1], qb, (size_t)H*4);
/* DESCENDING: plane row i lives at C bytes i*1024, so it is destroyed by the write of
C row floor(i/4) <= i (M2) / row H+floor(i/4) (M4); descending reads row i first. */
for (int i=H-1; i>=0; i--) {
const size_t rp = (size_t)(i>>1)*1024 + (size_t)(i&1)*16;
const int16_t *r2 = (i<2) ? f1b_r0[0]+(size_t)(i&1)*16 : (const int16_t *)qa+rp;
const int16_t *r4 = (i<2) ? f1b_r0[1]+(size_t)(i&1)*16 : (const int16_t *)qb+rp;
const int16_t *r1=prod[0]+rp, *r3=prod[1]+rp, *r5=prod[2]+rp,
*r6=prod[3]+rp, *r7=prod[4]+rp;
for (int j=0; j<H; j+=16) {
const __m256i p1=_mm256_load_si256((const __m256i *)(r1+j*2));
const __m256i p2=_mm256_load_si256((const __m256i *)(r2+j*2));
const __m256i p3=_mm256_load_si256((const __m256i *)(r3+j*2));
const __m256i p4=_mm256_load_si256((const __m256i *)(r4+j*2));
const __m256i p5=_mm256_load_si256((const __m256i *)(r5+j*2));
const __m256i p6=_mm256_load_si256((const __m256i *)(r6+j*2));
const __m256i p7=_mm256_load_si256((const __m256i *)(r7+j*2));
CBH(0, j);
CBH(1, j+8);
}
}
_mm_sfence();
}
void matrix_multiply(int n, const float *A, const float *B, float *C) {
(void)n;
qa = (int16_t *)C; qb = (int16_t *)C + (size_t)N*N; /* ★ 复用 C 的缓冲区当 quant 的输出 */
quantize(A,qa);quantize(B,qb);
const size_t off=(size_t)H*N;
Mat a11{qa,nullptr,0},a12{qa+H,nullptr,0},
a21{qa+off,nullptr,0},a22{qa+off+H,nullptr,0};
Mat b11{qb,nullptr,0},b12{qb+H,nullptr,0},
b21{qb+off,nullptr,0},b22{qb+off+H,nullptr,0};
product({a11.x,a22.x,+1},{b11.x,b22.x,+1},prod[0]);
product(a11,{b12.x,b22.x,-1},prod[1]);
product({a11.x,a12.x,+1},b22,prod[2]);
product({a21.x,a11.x,-1},{b11.x,b12.x,+1},prod[3]);
product({a12.x,a22.x,-1},{b21.x,b22.x,+1},prod[4]);
product({a21.x,a22.x,+1},b11,(int16_t *)qa);
product(a22,{b21.x,b11.x,-1},(int16_t *)qb);
combine(C);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 14.271 ms | 7 MB + 212 KB | Accepted | Score: 100 | 显示更多 |