#version 450 // Register-blocked mul_mm for Mali (#9584/#9715). The naive shared-memory tile // (mul_mm_tiled) is slower than scalar on Mali-G715 — staging + per-K-tile // barriers cost more than the bandwidth they save, because the kernel still // does 1 MAC per shared load (same arithmetic intensity as scalar, which Mali's // UMA cache already serves well). The real cooperative-matrix-free win is // REGISTER BLOCKING: a 64x64 workgroup tile where each of the 256 threads keeps // a 4x4 accumulator, so each shared element feeds 4 MACs (8 reg-loads → 16 MACs // per k step) — 8x the arithmetic intensity of scalar. C[M,N]=A[M,K]*B[K,N] f32. #define TM 64 // workgroup output rows #define TN 64 // workgroup output cols #define TK 16 // K step staged into shared memory #define WM 4 // per-thread output rows #define WN 4 // per-thread output cols // local = (TN/WN) x (TM/WM) = 16 x 16 = 256 threads layout(local_size_x = 16, local_size_y = 16) in; layout(std430, binding = 0) readonly buffer ABuf { float A[]; }; layout(std430, binding = 1) readonly buffer BBuf { float B[]; }; layout(std430, binding = 2) writeonly buffer CBuf { float C[]; }; layout(push_constant) uniform Push { uint M, N, K; } p; shared float As[TM][TK]; // 64 x 16 shared float Bs[TK][TN]; // 16 x 64 void main() { uint tx = gl_LocalInvocationID.x; // 0..15 uint ty = gl_LocalInvocationID.y; // 0..15 uint tid = ty * 16u + tx; // 0..255 uint rowBase = gl_WorkGroupID.y * TM; uint colBase = gl_WorkGroupID.x * TN; float acc[WM][WN]; for (int i = 0; i < WM; i++) for (int j = 0; j < WN; j++) acc[i][j] = 0.0; uint nTiles = (p.K + TK - 1u) / TK; for (uint t = 0; t < nTiles; ++t) { uint kBase = t * TK; // Collaboratively stage the A tile (TM*TK = 1024 = 256 threads * 4). for (uint l = 0; l < 4u; ++l) { uint idx = tid + l * 256u; uint r = idx / TK, k = idx % TK; uint gr = rowBase + r, gk = kBase + k; As[r][k] = (gr < p.M && gk < p.K) ? A[gr * p.K + gk] : 0.0; } // ...and the B tile (TK*TN = 1024). for (uint l = 0; l < 4u; ++l) { uint idx = tid + l * 256u; uint k = idx / TN, c = idx % TN; uint gk = kBase + k, gc = colBase + c; Bs[k][c] = (gk < p.K && gc < p.N) ? B[gk * p.N + gc] : 0.0; } barrier(); for (uint k = 0; k < TK; ++k) { float aReg[WM], bReg[WN]; for (int i = 0; i < WM; i++) aReg[i] = As[ty * WM + uint(i)][k]; for (int j = 0; j < WN; j++) bReg[j] = Bs[k][tx * WN + uint(j)]; for (int i = 0; i < WM; i++) for (int j = 0; j < WN; j++) acc[i][j] += aReg[i] * bReg[j]; } barrier(); } for (int i = 0; i < WM; i++) for (int j = 0; j < WN; j++) { uint gr = rowBase + ty * WM + uint(i); uint gc = colBase + tx * WN + uint(j); if (gr < p.M && gc < p.N) C[gr * p.N + gc] = acc[i][j]; } }