Skip to content

Instantly share code, notes, and snippets.

@0xekez
Created August 21, 2025 15:02
Show Gist options
  • Select an option

  • Save 0xekez/c94ba3d5b43df10d17c98581e91280e3 to your computer and use it in GitHub Desktop.

Select an option

Save 0xekez/c94ba3d5b43df10d17c98581e91280e3 to your computer and use it in GitHub Desktop.
// A, B, C are input in column-major form.
kernel void matmul(
constant uint& n,
constant uint& k,
constant uint& m,
constant float& alpha,
constant float& beta,
const device float* A,
const device float* B,
device float* C,
uint2 tpos [[thread_position_in_grid]]
) {
uint2 c_origin = tpos*4;
uint2 a_origin = uint2(c_origin.x,0);
uint2 b_origin = uint2(0,c_origin.y);
metal::float4x4 acc = {0.};
metal::float4x4 a4;
metal::float4x4 b4;
for (uint l = 0; l<k/4; l++) {
uint2 a_pos = a_origin + uint2(0,l*4);
uint2 b_pos = b_origin + uint2(l*4,0);
for (uint i = 0; i < 4; i++) {
uint2 ap = a_pos + uint2(0,i);
uint2 bp = b_pos + uint2(0,i);
if (ap.x+4<=n)
a4[i] = reinterpret_cast<const device float4*>(A+ap.x+ap.y*n)[0];
else
for (uint j=0; j<4; j++)
a4[i][j] = (ap.x+j<n) ? A[(ap.x+j)+ap.y*n] : 0.;
b4[i] = bp.y<m ? reinterpret_cast<const device float4*>(B+bp.x+bp.y*k)[0] : float4(0);
}
acc += a4*b4;
}
for (uint l=(k/4)*4; l<k; l++)
for (uint i=0;i<4;i++) {
uint row = c_origin.x+i;
if (row >= n) continue;
float a = A[row+l*n];
for (uint j=0; j<4; j++) {
uint col = c_origin.y+j;
if (col >= m) continue;
acc[j][i] += a*B[l+col*k];
}
}
for (uint i=0; i<4 ;i++)
for (uint j=0; j<4; j++) {
uint2 c_pos = c_origin + uint2(i,j);
if (c_pos.x<n&&c_pos.y<m)
C[c_pos.x+c_pos.y*n] = alpha*acc[j][i] + beta*C[c_pos.x+c_pos.y*n];
}
}
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment