it's been a little while since i messed around with writing any sort of neural net kernels. lots of people are doing it these days and the tools are pretty advanced for cuda and whatever else we've got going on in gpu land.
and yet people keep making all these dang ARM computers in my cloud. so how do we make two matrices who love each other very much create a third matrix using arm assembly? and, more importantly, can we make it gooooooo (pretty fast, maybe not gpu fast but you know kinda fast)?
anyway we're gonna try to do the thing using neon assembly to start with. nothing too advanced. i'll make another article about using SME/SVE2 after this.
today's weapon of choice is a macbook air with an M2 processor. The theoretical limit of f32 single-threaded matmul for this processor is 112 GFlops (3.5Ghz * 4 FMA/Cycle * 4 lanes (128-bit register / 32 bits) * 2 flops/fma (multiply, accumulate)).
the performance numbers in this post are based on a sweep of shapes between 64x64 and 1024x1024 (with 20 runs each step), using the average value for each kernel. you can find a chart of sweeps at the bottom of this post.
here's the basic idea:
static void mmul_ref(int M, int K, int N, float* A, float* B, float* C) {
for (int m = 0; m < M; m++)
for (int n = 0; n < N; n++) {
float acc = 0.0f;
for (int k = 0; k < K; k++)
acc += A[m * K + k] * B[k * N + n];
C[m * N + n] = acc;
}
}
At -O2 this runs at 2.09 Gflops, about 1.9% of our theoretical maximum..
let's try to write this naive implementation in assembly
// function signature is the same as above
// x0 = M
// x1 = K
// x2 = N
// x3 = *A
// x4 = *B
// x5 = *C
_mmul_naive_asm:
mov w9, #0 // m
1:
mov w10, #0 // n
cmp w9, M
bge 1f
2:
mov w11, #0 // k
cmp w10, N
bge 2f
movi v2.4s, #0
3:
cmp w11, K
bge 3f
// matC[m * N + n] += matA[m * K + k] * matB[k * N + n]
madd w13, w9, K, w11 // matA offset
madd w14, w11, N, w10 // matB offset
ldr s0, [x3, x13, lsl #2] // load from x3 + offA * 4
ldr s1, [x4, x14, lsl #2] // load from x4 + offB * 4
fmadd s2, s0, s1, s2 // FMA acc = acc + (s0 * s1)
add w11, w11, #1 // increment k by 1
b 3b
3:
str s2, [x5]
add x5, x5, #4
add w10, w10, #1 // increment n by 1
b 2b
2:
add w9, w9, #1 // increment m by 1
b 1b
1:
ret
alright we have 1.98 Gflops (or about 1.8% of the theoretical max)... the c compiler is smarter than me.
let's do some simd...
we're going to load 4 values from B into a single 128-bit register and broadcast a single value from A into a 128-bit register and use the fmla instruction to do 4 fma's at the same time across 4 output values. an important caveat here is now N needs to be divisible by 4. we could add some overhang stuff but right now we're in control of the shape of the matrices and i don't feel like doing that.
_mmul_vec_asm:
mov w9, #0 // m
1:
mov w10, #0 // n
cmp w9, M
bge 1f
2:
mov w11, #0 // k
cmp w10, N
bge 2f
// clear output register
movi v2.4s, #0
3:
cmp w11, K
bge 3f
// matC[m * N + n] += matA[m * K + k] * matB[k * N + n]
madd w13, w9, K, w11 // matA offset
madd w14, w11, N, w10 // matB offset
add x15, x3, x13, lsl #2
lsl w14, w14, #2
ldr q1, [x4, x14] // load 4 matB values
ld1r {v0.4s}, [x15] // load a single matA value into all lanes of q0
fmla v2.4s, v1.4s, v0.4s
add w11, w11, #1 // increment k by 1
b 3b
3:
str q2, [x5]
add x5, x5, #16
add w10, w10, #4 // increment n by 4
b 2b
2:
add w9, w9, #1 // increment m by 1
b 1b
1:
ret
alright this gets us to 7.86 Gflops, about 7% of our target.. but about 4x better than where we were at before - progress!
let's try unrolling K a little more too... get a little more compute happening...
_mmul_vec_asm2:
mov w9, #0 // m
1:
mov w10, #0 // n
cmp w9, M
bge 1f
2:
mov w11, #0 // k
cmp w10, N
bge 2f
// clear output accumulators
// We are using 4 here to remove dependencies within K blocks
movi v0.4s, #0
movi v1.4s, #0
movi v2.4s, #0
movi v3.4s, #0
// q0..q3 -> outputs
// q4..q7 -> matB inputs
// q16..q19 -> matA inputs
3:
cmp w11, K
bge 3f
// matC[m * N + n] += matA[m * K + k] * matB[k * N + n]
madd w13, w9, K, w11 // matA offset
madd w14, w11, N, w10 // matB offset
add x15, x3, x13, lsl #2
lsl w14, w14, #2
ldr q4, [x4, x14] // load 4 matB values
ld1r {v16.4s}, [x15], #4 // load a single matA value into all lanes
ld1r {v17.4s}, [x15], #4
ld1r {v18.4s}, [x15], #4
ld1r {v19.4s}, [x15], #4
add w14, w14, N, lsl #2
ldr q5, [x4, x14] // load 4 matB values
add w14, w14, N, lsl #2
ldr q6, [x4, x14] // load 4 matB values
add w14, w14, N, lsl #2
ldr q7, [x4, x14] // load 4 matB values
fmla v0.4s, v4.4s, v16.4s
fmla v1.4s, v5.4s, v17.4s
fmla v2.4s, v6.4s, v18.4s
fmla v3.4s, v7.4s, v19.4s
add w11, w11, #4
b 3b
3:
// reduce to v0.4s
// we will sum v0 and v1 into v0, then v2 and v3 into v2, and finally v0 and v2 into v0.
// this is again to allow multiple floating point instructions to run in parallel
fadd v0.4s, v0.4s, v1.4s
fadd v2.4s, v2.4s, v3.4s
fadd v0.4s, v0.4s, v2.4s
// write 4 outputs...
str q0, [x5]
add x5, x5, #16
add w10, w10, #4 // do 4 N at once
b 2b
2:
add w9, w9, #1
b 1b
1:
ret
ok now we're at 15.95Gflops (14.2%). So we're starting to move.. something you might have noticed here is we're still writing to a single output but we have 4 accumulators. this is so we can have the cpu pipeline the fmla instructions by eliminating dependencies between fmla executions. if we use a single register here we lose 3-5Gflops.
but the problem here is we're not really getting a lot of compute intensity for each load. one of the classic tricks to improve matrix multiplication performance that you might have seen or used if you've ever done any gpu kernels is tiling. if you're not familiar with it, the idea is to work on patches of the output matrix by sliding windows along the input matrices' K inner dimension. this way we can load a bunch of data and then do a whole bunch of compute to make those loads relatively less expensive via amortization. we are also reducing the overall number of redundant loads by reusing data in registers rather than reloading it.
We have 32 128-bit registers at our disposal here for doing float operations.. so we can have an 8x8 output tile using 16 of those registers and use the other 16 for inputs (8x4 @ 4x8). Now we will do 16 loads and 64 fmla operations on each iteration over the K dimension. We will still unroll K by 4 as we did in the last example, but we'll do 8 values in both the M and N dimensions.
_mmul_r_tile_8x8:
// free up fp registers since we will use all 32
stp d8, d9, [sp, #-64]!
stp d10, d11, [sp, #16]
stp d12, d13, [sp, #32]
stp d14, d15, [sp, #48]
// q0..q15 -> outputs (8x8=64 / 4 single-precision = 16 registers)
// q16..q23 -> matA inputs
// q24..q31 -> matB inputs
mov w9, #0 // m
1:
cmp w9, M
bge 1f
mov w10, #0 // n
2:
cmp w10, N
bge 2f
// clear accumulators for this tile
.altmacro
.macro clear_reg index
movi v\index\().4s, #0
.endm
.set i, 0
.rept 16
clear_reg %i
.set i, i+1
.endr
.noaltmacro
mov w11, #0 // k
3:
cmp w11, K
bge 3f
// 8x4 A tile x 4x8 B tile, 64 fmla per K-block.
// 4 is in the K direction, 8 is in the M or N direction.
// matC[m*N+n] += matA[m*K+k] * matB[k*N+n]
madd w14, w11, N, w10 // matB offset
// load 32 matB values
add x14, x4, x14, lsl #2 // matB pointer + offset (bytes)
ldp q24, q25, [x14]
add x14, x14, x2, lsl #2 // move matB pointer by stride N
ldp q26, q27, [x14]
add x14, x14, x2, lsl #2
madd w13, w9, K, w11 // matA offset
add x13, x3, x13, lsl #2 // matA pointer + offset (bytes)
ldr q16, [x13]
add x13, x13, x1, lsl #2 // move matA pointer by stride K
ldr q17, [x13]
add x13, x13, x1, lsl #2
ldr q18, [x13]
add x13, x13, x1, lsl #2
ldr q19, [x13]
add x13, x13, x1, lsl #2
ldr q20, [x13]
add x13, x13, x1, lsl #2
ldr q21, [x13]
add x13, x13, x1, lsl #2
ldr q22, [x13]
add x13, x13, x1, lsl #2
ldr q23, [x13]
add x13, x13, x1, lsl #2
// template for 16 fmlas... makes life a little easier.
.macro fmla16 lane, a0, a1
fmla v0.4s, \a0\().4s, v16.s[\lane]
fmla v1.4s, \a1\().4s, v16.s[\lane]
fmla v2.4s, \a0\().4s, v17.s[\lane]
fmla v3.4s, \a1\().4s, v17.s[\lane]
fmla v4.4s, \a0\().4s, v18.s[\lane]
fmla v5.4s, \a1\().4s, v18.s[\lane]
fmla v6.4s, \a0\().4s, v19.s[\lane]
fmla v7.4s, \a1\().4s, v19.s[\lane]
fmla v8.4s, \a0\().4s, v20.s[\lane]
fmla v9.4s, \a1\().4s, v20.s[\lane]
fmla v10.4s, \a0\().4s, v21.s[\lane]
fmla v11.4s, \a1\().4s, v21.s[\lane]
fmla v12.4s, \a0\().4s, v22.s[\lane]
fmla v13.4s, \a1\().4s, v22.s[\lane]
fmla v14.4s, \a0\().4s, v23.s[\lane]
fmla v15.4s, \a1\().4s, v23.s[\lane]
.endm
// k=0
fmla16 0, v24, v25
ldp q28, q29, [x14]
add x14, x14, x2, lsl #2
//k=1
fmla16 1, v26, v27
ldp q30, q31, [x14]
//k=2
fmla16 2, v28, v29
//k=3
fmla16 3, v30, v31
add w11, w11, #4
b 3b
3:
// write accumulators to memory
madd w14, w9, N, w10 // matC offset
add x14, x5, x14, lsl #2
stp q0, q1, [x14]
add x14, x14, x2, lsl #2
stp q2, q3, [x14]
add x14, x14, x2, lsl #2
stp q4, q5, [x14]
add x14, x14, x2, lsl #2
stp q6, q7, [x14]
add x14, x14, x2, lsl #2
stp q8, q9, [x14]
add x14, x14, x2, lsl #2
stp q10, q11, [x14]
add x14, x14, x2, lsl #2
stp q12, q13, [x14]
add x14, x14, x2, lsl #2
stp q14, q15, [x14]
add w10, w10, #8
b 2b
2:
add w9, w9, #8
b 1b
1:
// pop the registers back from the stack
ldp d10, d11, [sp, #16]
ldp d12, d13, [sp, #32]
ldp d14, d15, [sp, #48]
ldp d8, d9, [sp], #64
ret
So this gets us to 96.65 Gflops, about 86.3% of our target!
Some details here.. first, hell yea we unrolled that by a lot...
Second in these vector instructions you'll notice we are indexing into the registers loaded in A. This is because we loaded values in the K dimension for A when we did the 128-bit load (whereas for B we were loading values along the N dimension). So as we iterate along K we need to index into the registers holding A data and increment which registers we're using for the B data.
third interesting thing we have some of our B data loads interleaved with flmas. this is to do something like double buffering but maybe in a lazier (as in i had to type less to do this) way. it buys us about 5 Gflops.
Because our fmlas with dependencies are pretty far apart we are also not doing the trick from the previous one where we used multiple accumulators.
if you've done this on a gpu before you might also notice that there's no intermediate load step into workgroup shared memory. first of all there's no work group so there's no workgroup shared memory. secondly with this tile size we end up being compute bound on the M2, so there's no need for us to explicitly shuffle memory around.
but why stop at 97 gflops when we can go for 200 gflops? FP16 time.
this one is going to be an 8x16 tile. we'll use the same number of registers but be able to double the number of values we compute.
_mmul_r_tile_8x8_f16:
// free up fp registers since we will use all 32
stp d8, d9, [sp, #-64]!
stp d10, d11, [sp, #16]
stp d12, d13, [sp, #32]
stp d14, d15, [sp, #48]
// q16..q31 -> outputs (16x8=128 / 8 half-precision = 16 registers)
// q0..q7 -> matA inputs
// q8..q15 -> matB inputs
mov w9, #0 // m
1:
cmp w9, M
bge 1f
mov w10, #0 // n
2:
cmp w10, N
bge 2f
// clear accumulators for this tile
.altmacro
.macro clear_reg index
movi v\index\().8h, #0
.endm
.set i, 0
.rept 16
clear_reg %i + 16
.set i, i+1
.endr
.noaltmacro
mov w11, #0 // k
3:
cmp w11, K
bge 3f
// 8x4 A tile x 4x8 B tile, 64 fmla per K-block.
// 4 is in the K direction, 8 is in the M or N direction.
// matC[m*N+n] += matA[m*K+k] * matB[k*N+n]
madd w14, w11, N, w10 // matB offset
// load 32 matB values
add x14, x4, x14, lsl #1 // matB pointer + offset (bytes)
ldp q8, q9, [x14]
add x14, x14, x2, lsl #1 // move matB pointer by stride N
ldp q10, q11, [x14]
add x14, x14, x2, lsl #1
// load all the needed values from A (8x8)
madd w13, w9, K, w11 // matA offset
add x13, x3, x13, lsl #1 // matA pointer + offset (bytes)
ldr q0, [x13]
add x13, x13, x1, lsl #1 // move matA pointer by stride K
ldr q1, [x13]
add x13, x13, x1, lsl #1
ldr q2, [x13]
add x13, x13, x1, lsl #1
ldr q3, [x13]
add x13, x13, x1, lsl #1
ldr q4, [x13]
add x13, x13, x1, lsl #1
ldr q5, [x13]
add x13, x13, x1, lsl #1
ldr q6, [x13]
add x13, x13, x1, lsl #1
ldr q7, [x13]
.macro fmla16 lane, a0, a1
fmla v16.8h, \a0\().8h, v0.h[\lane]
fmla v17.8h, \a1\().8h, v0.h[\lane]
fmla v18.8h, \a0\().8h, v1.h[\lane]
fmla v19.8h, \a1\().8h, v1.h[\lane]
fmla v20.8h, \a0\().8h, v2.h[\lane]
fmla v21.8h, \a1\().8h, v2.h[\lane]
fmla v22.8h, \a0\().8h, v3.h[\lane]
fmla v23.8h, \a1\().8h, v3.h[\lane]
fmla v24.8h, \a0\().8h, v4.h[\lane]
fmla v25.8h, \a1\().8h, v4.h[\lane]
fmla v26.8h, \a0\().8h, v5.h[\lane]
fmla v27.8h, \a1\().8h, v5.h[\lane]
fmla v28.8h, \a0\().8h, v6.h[\lane]
fmla v29.8h, \a1\().8h, v6.h[\lane]
fmla v30.8h, \a0\().8h, v7.h[\lane]
fmla v31.8h, \a1\().8h, v7.h[\lane]
.endm
// k=0
fmla16 0, v8, v9
ldp q12, q13, [x14]
add x14, x14, x2, lsl #1
//k=1
fmla16 1, v10, v11
ldp q14, q15, [x14]
add x14, x14, x2, lsl #1
//k=2
fmla16 2, v12, v13
ldp q8, q9, [x14]
add x14, x14, x2, lsl #1
//k=3
fmla16 3, v14, v15
ldp q10, q11, [x14]
add x14, x14, x2, lsl #1
//k=4
fmla16 4, v8, v9
ldp q12, q13, [x14]
add x14, x14, x2, lsl #1
//k=5
fmla16 5, v10, v11
ldp q14, q15, [x14]
//k=6
fmla16 6, v12, v13
//k=7
fmla16 7, v14, v15
add w11, w11, #8
b 3b
3:
// write accumulators to memory
madd w14, w9, N, w10 // matC offset
add x14, x5, x14, lsl #1
stp q16, q17, [x14]
add x14, x14, x2, lsl #1
stp q18, q19, [x14]
add x14, x14, x2, lsl #1
stp q20, q21, [x14]
add x14, x14, x2, lsl #1
stp q22, q23, [x14]
add x14, x14, x2, lsl #1
stp q24, q25, [x14]
add x14, x14, x2, lsl #1
stp q26, q27, [x14]
add x14, x14, x2, lsl #1
stp q28, q29, [x14]
add x14, x14, x2, lsl #1
stp q30, q31, [x14]
add w10, w10, #16 // advance N by 16
b 2b
2:
add w9, w9, #8 // advance M by 8
b 1b
1:
// pop the registers back from the stack
ldp d10, d11, [sp, #16]
ldp d12, d13, [sp, #32]
ldp d14, d15, [sp, #48]
ldp d8, d9, [sp], #64
ret
now we're hitting 203.65 gflops, about 91% of the theoretical 224 FP16 GFlops the M2 processor can do on a single core. Something we're doing a little more of in this is interleaving loads with compute to do some double buffering. This also helps recycle registers because we do not have enough registers for a K block size of 8 (although we are sort of forced into that size by loading 128 bits of K from matrix A). You'll notice too that we swapped the input and accumulator registers, this is because the .h[lane] notation does not work above v15.
anyway this was fun... next time we'll do some integer matmul and mess around with sme or sve2 or whatever...
Epilogue: Maybe we do need some cache blocking, after all.
So I started poking at Graviton instances, using C9g which apparently has a 2MB L2 cache (vs the M2's 16MB L2) and saw some pretty steep performance drop off when running the sweep...
the dashed lines are the Graviton instance.. you can see that as the matrix sizes grow the performance decreases. So, we need to load some data into an intermediate scratch buffer first. For this we're going to make a single scratch block to pull data out of B into a block where loads can be contiguous (vs the current stride of N). We can block in M (for mat A) and K as well, but N (mat B) is probably going to be the most fruitful for us to start with because for each M row we are loading the entire matrix B, whereas we only load an 8xK strip from matrix A (which, at the tested sizes, is small enough to stay resident in L1...)
mmul_c_tile_r_tile_8x16_f16:
// preserve link register because we'll be calling out to malloc/free
stp x29, x30, [sp, #-16]!
// free up some registers
stp x19, x20, [sp, #-64]!
stp x21, x22, [sp, #16]
stp x23, x24, [sp, #32]
stp x25, x26, [sp, #48]
// free up fp registers since we will use all 32
stp d8, d9, [sp, #-64]!
stp d10, d11, [sp, #16]
stp d12, d13, [sp, #32]
stp d14, d15, [sp, #48]
mov x19, x0
mov x20, x1
mov x21, x2
mov x22, x3
mov x23, x4
mov x24, x5
mov x25, x6
mul w0, w6, w1 // BN * K
lsl w0, w0, #1 // * 2 (f16)
uxtw x0, w0
bl MALLOC // alloc scratch memory
cmp x0, #0
mov x8, x0
mov x0, #1
beq alloc_fail
mov x0, x19
mov x1, x20
mov x2, x21
mov x3, x22
mov x4, x23
mov x5, x24
mov x6, x25
// q16..q31 -> outputs (16x8=128 / 8 half-precision = 16 registers)
// q0..q7 -> matA inputs
// q8..q15 -> matB inputs
// x0 -> M
// x1 -> K
// x2 -> N
// x3 -> pA
// x4 -> pB
// x5 -> pC
// x6 -> BN
// x8 -> pScratch
mov w17, #0 // bn
4: // outer loop blocking N
cmp w17, N
bge 4f
// alright we want to load data from mat B into the scratch memory
// we are going to load each tile contiguously so that they are packed nicely
mov w10, #0 // bn_l
mov x15, x8 // scratch write pointer we can advance
5: // BN
cmp w10, w6
bge 5f
mov w11, #0 // k_l
6:
cmp w11, K
bge 6f
// Offset = ((bn + bn_l) + k_l * N) * 2 (bytes)
add w12, w17, w10
madd w12, w11, N, w12
uxtw x12, w12
lsl x12, x12, #1
add x12, x12, pMatB
ldp q0, q1, [x12]
stp q0, q1, [x15], #32 // wrote 32 bytes, advance x15 pointer
add w11, w11, #1 // advance k_l by 1
b 6b
6:
add w10, w10, #16 // tile is 16-wide
b 5b
5:
// perform compute on loaded scratch
mov w9, #0 // m
1:
cmp w9, M
bge 1f
mov x15, x8 // scratch read pointer we can advance
mov w10, #0 // n
2:
cmp w10, w6 // n < BN?
bge 2f
// clear accumulators for this tile
.altmacro
.macro clear_reg index
movi v\index\().8h, #0
.endm
.set i, 0
.rept 16
clear_reg %i + 16
.set i, i+1
.endr
.noaltmacro
mov w11, #0 // k
3:
cmp w11, K
bge 3f
// 8x4 A tile x 4x8 B tile, 64 fmla per K-block.
// 4 is in the K direction, 8 is in the M or N direction.
// matC[m*N+n] += matA[m*K+k] * matB[k*N+n]
ldp q8, q9, [x15], #32
ldp q10, q11, [x15], #32
// load all the needed values from A (8x8)
madd w13, w9, K, w11 // matA offset
add x13, x3, x13, lsl #1 // matA pointer + offset (bytes)
ldr q0, [x13]
add x13, x13, x1, lsl #1 // move matA pointer by stride K
ldr q1, [x13]
add x13, x13, x1, lsl #1
ldr q2, [x13]
add x13, x13, x1, lsl #1
ldr q3, [x13]
add x13, x13, x1, lsl #1
ldr q4, [x13]
add x13, x13, x1, lsl #1
ldr q5, [x13]
add x13, x13, x1, lsl #1
ldr q6, [x13]
add x13, x13, x1, lsl #1
ldr q7, [x13]
.macro fmla16 lane, a0, a1
fmla v16.8h, \a0\().8h, v0.h[\lane]
fmla v17.8h, \a1\().8h, v0.h[\lane]
fmla v18.8h, \a0\().8h, v1.h[\lane]
fmla v19.8h, \a1\().8h, v1.h[\lane]
fmla v20.8h, \a0\().8h, v2.h[\lane]
fmla v21.8h, \a1\().8h, v2.h[\lane]
fmla v22.8h, \a0\().8h, v3.h[\lane]
fmla v23.8h, \a1\().8h, v3.h[\lane]
fmla v24.8h, \a0\().8h, v4.h[\lane]
fmla v25.8h, \a1\().8h, v4.h[\lane]
fmla v26.8h, \a0\().8h, v5.h[\lane]
fmla v27.8h, \a1\().8h, v5.h[\lane]
fmla v28.8h, \a0\().8h, v6.h[\lane]
fmla v29.8h, \a1\().8h, v6.h[\lane]
fmla v30.8h, \a0\().8h, v7.h[\lane]
fmla v31.8h, \a1\().8h, v7.h[\lane]
.endm
// k=0
fmla16 0, v8, v9
ldp q12, q13, [x15], #32
//k=1
fmla16 1, v10, v11
ldp q14, q15, [x15], #32
//k=2
fmla16 2, v12, v13
ldp q8, q9, [x15], #32
//k=3
fmla16 3, v14, v15
ldp q10, q11, [x15], #32
//k=4
fmla16 4, v8, v9
ldp q12, q13, [x15], #32
//k=5
fmla16 5, v10, v11
ldp q14, q15, [x15], #32
//k=6
fmla16 6, v12, v13
//k=7
fmla16 7, v14, v15
add w11, w11, #8
b 3b // advance k by 8
3:
// write accumulators to memory
// matC offset
// m * N + (bn + n)
add w14, w10, w17 // bn + n
madd w14, w9, N, w14
add x14, x5, x14, lsl #1
stp q16, q17, [x14]
add x14, x14, x2, lsl #1
stp q18, q19, [x14]
add x14, x14, x2, lsl #1
stp q20, q21, [x14]
add x14, x14, x2, lsl #1
stp q22, q23, [x14]
add x14, x14, x2, lsl #1
stp q24, q25, [x14]
add x14, x14, x2, lsl #1
stp q26, q27, [x14]
add x14, x14, x2, lsl #1
stp q28, q29, [x14]
add x14, x14, x2, lsl #1
stp q30, q31, [x14]
add w10, w10, #16 // advance n by 16
b 2b
2:
add w9, w9, #8 // advance m by 8
b 1b
1:
add w17, w17, w6 // advance bn by BN
b 4b
4:
mov x0, x8 // free the scratch memory
bl FREE
mov x0, #0
alloc_fail:
// pop the registers back from the stack
ldp d10, d11, [sp, #16]
ldp d12, d13, [sp, #32]
ldp d14, d15, [sp, #48]
ldp d8, d9, [sp], #64
ldp x21, x22, [sp, #16]
ldp x23, x24, [sp, #32]
ldp x25, x26, [sp, #48]
ldp x19, x20, [sp], #64
ldp x29, x30, [sp], #16
ret
This is starting to look a little more like how you would do it on a gpu...
anyway this has given us a nice boost in performance at larger matrix sizes.. For these tests I used BN=128, so the scratch is holding 8 tiles worth of data with their full K dimension.
(the dashed line is non-blocked)



Top comments (0)