This is a continuation of my previous article on using ARM NEON to do matrix multiplication.
Here's a fun one, we're going to use BFMMLA to do 2x4 @ 4x2 matmul on bf16 inputs. The result is a 2x2 matrix in F32 accumulators in a single float register. We'll break this one up a little more than the last post because it's a little more involved.
#if defined(__APPLE__)
#define MALLOC _malloc
#define FREE _free
#else
#define MALLOC malloc
#define FREE free
#endif
.global _mmul_8x8_bf16_mma
.extern MALLOC
.extern FREE
_mmul_8x8_bf16_mma:
// 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, #-80]!
stp x21, x22, [sp, #16]
stp x23, x24, [sp, #32]
stp x25, x26, [sp, #48]
stp x27, x28, [sp, #64]
// 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
mov x26, x7
mul w0, w6, w1 // BN * K
lsl w0, w0, #1 // * 2 (f16)
uxtw x0, w0
bl MALLOC // alloc scratch memory
cmp x0, #0
mov x27, x0
mov x0, #1
beq alloc_fail
mul w0, w26, w20 // BM * K
lsl w0, w0, #1
uxtw x0, w0
bl MALLOC
cmp x0, #0
mov x28, 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
mov x7, x26
mov x25, x27 // B scratch pointer in x25
mov x26, x28 // A scratch pointer in x26
// 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
// x7 -> BM
// x20..x24 -> counters
// x25 -> B scratch
// x26 -> A scratch
mostly housekeeping and prep here. We are allocating two scratch slabs for inputs from the A and B matrices. The reason we are doing this is because we need to swizzle the inputs to be in the right shape for a 2x4 or 4x2 matrix such that we get a sensible output result. We will swizzle as we load data rather than as we perform compute so that we don't have to do it multiple times.
mov x20, #0 // bm
.M_block: // outer loop blocking M
cmp w20, M
bge .M_block_x
mov x15, x26 // A scratch pointer we can advance
mov w21, #0 // bm_l
.M_block_load:
cmp w21, w7
bge .M_block_load_x
mov w22, #0 // bm_k_l
.M_block_load_k:
cmp w22, K
bge .M_block_load_k_x
// Why are we loading 8 M rows at a time? Because we load 8 N columns at a time due to vector width
// Why are we reordering the K blocks? Loading 8 K rows in the N block load phase may cause register
// exhaustion (but mostly I haven't tried it yet).. we make up for this in the compute phase
// offset = (bm + bm_l) * K + bm_k_l
add w13, w20, w21
madd w13, w13, K, w22
uxtw x13, w13
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]
zip1 v8.2d, v0.2d, v1.2d // [m0k0, m0k1, m0k2, ..., m1k3]
zip2 v9.2d, v0.2d, v1.2d // [m0k4, m0k5, m0k6, ..., m1k7]
zip1 v10.2d, v2.2d, v3.2d // [m2k0, m2k1, m2k2, ... m3k3]
zip2 v11.2d, v2.2d, v3.2d // [m2k4, ...]
zip1 v12.2d, v4.2d, v5.2d // [m4k0, ...]
zip2 v13.2d, v4.2d, v5.2d // [m4k4, ...]
zip1 v14.2d, v6.2d, v7.2d // [m6k0, ..., m7k3]
zip2 v15.2d, v6.2d, v7.2d // [m6k4, ..., m7k7]
// reorder k blocks such that there are 4 k's for all 8 of the m's
stp q8, q10, [x15], #32
stp q12, q14, [x15], #32
stp q9, q11, [x15], #32
stp q13, q15, [x15], #32
add w22, w22, #8 // loading 8 K at a time
b .M_block_load_k
.M_block_load_k_x:
add w21, w21, #8 // loading 8 M rows at a time
b .M_block_load
.M_block_load_x:
Here we are loading data from the A matrix input and storing it into A scratch.. Here we want to reorder the input so that we have groups of 2 M rows and 4 K columns arranged contiguously. This is so that we can put these 8 values into a single 128-bit register representing a 2x4 matrix to use with bfmmla. Because of the load width we actually end up loading 8 values in the K dimension from mat A at once. So we put K values 4-7 after we have all 0-3 values because K will be our compute inner loop, so we'll iterate through all of K, then iterate through the N block, then we will move to the next M block. Our loads therefore store all of these values contiguously such that we can just do contiguous 128-bit wide ldr ahead of the compute phase.
Ordering the data in this way is useful for another reason. Because we are not storing all values for a single row contiguously we are able to take advantage of ILP more easily. Our accumulator for a given output value is not going to be used for all input K values at once, it's going to just do a group of 4, then we have to get through the rest of the output values, then we circle back and do the next group of 4. There is no dependency between any of the ops and a single accumulator register.
mov w21, #0 // bn
.N_block: // outer loop blocking N
cmp w21, N
bge .N_block_x
// 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 w22, #0 // bn_l
mov x15, x25 // B scratch write pointer we can advance
.N_block_load: // BN
cmp w22, w6
bge .N_block_load_x
mov w23, #0 // k_l
.N_block_load_k:
cmp w23, K
bge .N_block_load_k_x
// Offset = ((bn + bn_l) + k_l * N) * 2 (bytes)
add w12, w21, w22
madd w12, w23, N, w12
uxtw x12, w12
lsl x12, x12, #1
add x12, x12, pMatB
ldr q0, [x12]
add x12, x12, x2, lsl #1
ldr q1, [x12]
add x12, x12, x2, lsl #1
ldr q2, [x12]
add x12, x12, x2, lsl #1
ldr q3, [x12]
// Now we have 4x16 values from B, we need to reshape these into 4x2 vectors for
// bfmmla
// [N0{k0, k1, k2, k3}, N1{k0, k1, k2, k3}, ...]
zip1 v4.8h, v0.8h, v1.8h // [k0n0, k1n0, k0n1, k1n1, ..., k1n3]
zip2 v5.8h, v0.8h, v1.8h // [k0n4, k1n4, k0n5, k1n5, ..., k1n7]
zip1 v6.8h, v2.8h, v3.8h // [k2n0, k3n0, k2n1, k3n1, ..., k3n3]
zip2 v7.8h, v2.8h, v3.8h // [k2n4, k3n4, k2n5, k3n5, ..., k3n7]
zip1 v8.4s, v4.4s, v6.4s // [k0n0, k1n0, k2n0, ..., k3n1]
zip2 v9.4s, v4.4s, v6.4s // [k0n2, k1n2, k2n2, ..., k3n3]
zip1 v10.4s, v5.4s, v7.4s // [k0n4, k1n4, k2n4, ..., k3n5]
zip2 v11.4s, v5.4s, v7.4s // [k0n6, k1n6, k2n6, ..., k3n7]
stp q8, q9, [x15], #32 // wrote 32 bytes, advance x15 pointer
stp q10, q11, [x15], #32 // wrote 32 bytes, advance x15 pointer
add w23, w23, #4 // advance k_l by 4
b .N_block_load_k
.N_block_load_k_x:
add w22, w22, #8 // 8-wide
b .N_block_load
.N_block_load_x:
For the N block where we load from Matrix B, it's the same idea but our swizzle is a little bit different because the orientation for this matrix is different. Again we want to have 4 values from each column interleaved in pairs so that we can do a 2x4 @ 4x2 bfmmla.
// perform compute on loaded scratch
mov w22, #0 // m
.M_inner:
cmp w22, w7 // m < BM
bge .M_inner_x
mov x15, x25 // B scratch read pointer we can advance
mov w23, #0 // n
.N_inner:
cmp w23, w6 // n < BN?
bge .N_inner_x
// create an A scratch pointer we can advance in x14
mul w13, w22,w1
add x14, x26, w13, uxtw #1
// 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 w24, #0 // k
.K_inner:
cmp w24, K
bge .K_inner_x
ldp q0, q1, [x14], #32
ldp q8, q9, [x15], #32
ldp q10, q11, [x15], #32
// K0..K3
bfmmla v16.4s, v0.8h, v8.8h // M0+1, N0+1
bfmmla v17.4s, v0.8h, v9.8h // M0+1, N2+3
bfmmla v18.4s, v0.8h, v10.8h // M0+1, N4+5
bfmmla v19.4s, v0.8h, v11.8h // M0+1, N6+7
ldp q12, q13, [x15], #32
ldp q2, q3, [x14], #32
bfmmla v20.4s, v1.8h, v8.8h // M2+3, N0+1
bfmmla v21.4s, v1.8h, v9.8h // M2+3, N2+3
bfmmla v22.4s, v1.8h, v10.8h // M2+3, N4+5
bfmmla v23.4s, v1.8h, v11.8h // M2+3, N6+7
ldp q14, q15, [x15], #32
ldp q4, q5, [x14], #32
bfmmla v24.4s, v2.8h, v8.8h // M4+5, N0+1
bfmmla v25.4s, v2.8h, v9.8h // M4+5, N2+3
bfmmla v26.4s, v2.8h, v10.8h // M4+5, N4+5
bfmmla v27.4s, v2.8h, v11.8h // M4+5, N6+7
bfmmla v28.4s, v3.8h, v8.8h // M6+7, N0+1
bfmmla v29.4s, v3.8h, v9.8h // M6+7, N2+3
bfmmla v30.4s, v3.8h, v10.8h // M6+7, N4+5
bfmmla v31.4s, v3.8h, v11.8h // M6+7, N6+7
ldp q6, q7, [x14], #32
// K4..K7
bfmmla v16.4s, v4.8h, v12.8h // M0+1, N0+1
bfmmla v17.4s, v4.8h, v13.8h // M0+1, N2+3
bfmmla v18.4s, v4.8h, v14.8h // M0+1, N4+5
bfmmla v19.4s, v4.8h, v15.8h // M0+1, N6+7
bfmmla v20.4s, v5.8h, v12.8h // M2+3, N0+1
bfmmla v21.4s, v5.8h, v13.8h // M2+3, N2+3
bfmmla v22.4s, v5.8h, v14.8h // M2+3, N4+5
bfmmla v23.4s, v5.8h, v15.8h // M2+3, N6+7
bfmmla v24.4s, v6.8h, v12.8h // M4+5, N0+1
bfmmla v25.4s, v6.8h, v13.8h // M4+5, N2+3
bfmmla v26.4s, v6.8h, v14.8h // M4+5, N4+5
bfmmla v27.4s, v6.8h, v15.8h // M4+5, N6+7
bfmmla v28.4s, v7.8h, v12.8h // M6+7, N0+1
bfmmla v29.4s, v7.8h, v13.8h // M6+7, N2+3
bfmmla v30.4s, v7.8h, v14.8h // M6+7, N4+5
bfmmla v31.4s, v7.8h, v15.8h // M6+7, N6+7
add w24, w24, #8 // advance k by 8
b .K_inner
.K_inner_x:
Here we are doing the actual compute loop. We do 16 bfmmlas at once, then we do another 16 on the same group of accumulators. So we're moving through the K inner dimension 8 elements at a time and as you can see for both scratch pointers we can simply advance the pointer by 32 bytes every load.
// we need to swizzle the memory order here because all our registers are holding
// 2x2 matrices and we're working on an 8x8 tile
zip1 v0.2d, v16.2d, v17.2d // m0n0, m0n1, m0n2, m0n3
zip2 v1.2d, v16.2d, v17.2d // m1n0, m1n1, m1n2, m1n3
zip1 v2.2d, v18.2d, v19.2d // m0n4, m0n5, m0n6, m0n7
zip2 v3.2d, v18.2d, v19.2d // .. etc
zip1 v4.2d, v20.2d, v21.2d
zip2 v5.2d, v20.2d, v21.2d
zip1 v6.2d, v22.2d, v23.2d
zip2 v7.2d, v22.2d, v23.2d
zip1 v8.2d, v24.2d, v25.2d
zip2 v9.2d, v24.2d, v25.2d
zip1 v10.2d, v26.2d, v27.2d
zip2 v11.2d, v26.2d, v27.2d
zip1 v12.2d, v28.2d, v29.2d
zip2 v13.2d, v28.2d, v29.2d
zip1 v14.2d, v30.2d, v31.2d
zip2 v15.2d, v30.2d, v31.2d
// matC offset
// (m + bn) * N + (bn + n)
add w13, w20, w22
add w12, w21, w23 // bn + n
madd w12, w13, N, w12
uxtw x12, w12
// write accumulators to memory
add x12, x5, x12, lsl #2
stp q0, q2, [x12]
add x12, x12, x2, lsl #2
stp q1, q3, [x12]
add x12, x12, x2, lsl #2
stp q4, q6, [x12]
add x12, x12, x2, lsl #2
stp q5, q7, [x12]
add x12, x12, x2, lsl #2
stp q8, q10, [x12]
add x12, x12, x2, lsl #2
stp q9, q11, [x12]
add x12, x12, x2, lsl #2
stp q12, q14, [x12]
add x12, x12, x2, lsl #2
stp q13, q15, [x12]
add w23, w23, #8 // advance n by 8
b .N_inner
.N_inner_x:
add w22, w22, #8 // advance m by 8
b .M_inner
.M_inner_x:
add w21, w21, w6 // advance bn by BN
b .N_block
.N_block_x:
add w20, w20, w7 // advance bm by BM
b .M_block
.M_block_x:
mov x0, x25 // free the scratch memory
bl FREE
mov x0, x26
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 x27, x28, [sp, #64]
ldp x19, x20, [sp], #80
ldp x29, x30, [sp], #16
ret
Here's the end of it. The most interesting part is writing the accumulators. At this point the accumulator registers are holding 2x2 matrices that need to be swizzled to be in the right format for the 8x8 tile we're working on.
On Graviton 5, this gets us to about 280 Gflops on a single core. The performance is sensitive to the size of the M and N block size. A BN=BM=64 size nets us about 250 Gflops, whereas BN=128, BM=256 we see 280. The performance on my M2 Mac is much worse, I am guessing the instruction is not actually well supported by the M2 processor.

Top comments (0)