Something went wrong. Try again.
sentence embeddings in pure zig: bge-small with an HF-exact tokenizer and an SME matmul
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130//! fp32 matmul on the Arm Scalable Matrix Extension (Apple M4 and later).//!//! Each call computes a 16-row × 64-column block of y = x·Wᵀ: per K step, one//! 16-lane column of x and four 16-lane rows of the packed weight panels feed//! four FMOPA outer products into the ZA tiles za0..za3. Streaming mode is//! entered and left inside the one asm block, so no SVE/ZA state escapes it;//! smstart/smstop zero the vector registers, hence the v0..v31 clobbers.
const std = @import("std");const builtin = @import("builtin");
pub const available = builtin.cpu.arch == .aarch64 and builtin.cpu.has(.aarch64, .sme);
/// f32 lanes per streaming vector the kernel is written for; callers check RDSVLpub const lanes = 16;/// output columns per call: one 16-wide panel per ZA tilepub const cols = 4 * lanes;
/// Accumulate a zero-initialized 16×64 block./// `xa`: K×16 floats, x rows packed K-major. `w0..w3`: four weight panels, each/// K×16 floats. Writes `rows` (≤16) rows of 64 floats to `y`, `y_stride` floats apart.pub fn block(xa: [*]const f32, w: [4][*]const f32, k: usize, y: [*]f32, y_stride: usize, rows: usize) void { std.debug.assert(k > 0 and rows > 0 and rows <= lanes); asm volatile ( \\ smstart \\ ptrue p0.s \\ zero {za} \\ mov x9, %[k] \\ mov x10, %[xa] \\ mov x11, %[w0] \\ mov x13, %[w1] \\ mov x14, %[w2] \\ mov x15, %[w3] \\ 1: \\ ld1w {z0.s}, p0/z, [x10] \\ ld1w {z1.s}, p0/z, [x11] \\ ld1w {z2.s}, p0/z, [x13] \\ ld1w {z3.s}, p0/z, [x14] \\ ld1w {z4.s}, p0/z, [x15] \\ fmopa za0.s, p0/m, p0/m, z0.s, z1.s \\ fmopa za1.s, p0/m, p0/m, z0.s, z2.s \\ fmopa za2.s, p0/m, p0/m, z0.s, z3.s \\ fmopa za3.s, p0/m, p0/m, z0.s, z4.s \\ add x10, x10, #64 \\ add x11, x11, #64 \\ add x13, x13, #64 \\ add x14, x14, #64 \\ add x15, x15, #64 \\ subs x9, x9, #1 \\ b.ne 1b \\ mov w12, #0 \\ mov x16, %[y] \\ 2: \\ st1w {za0h.s[w12, 0]}, p0, [x16] \\ add x17, x16, #64 \\ st1w {za1h.s[w12, 0]}, p0, [x17] \\ add x17, x17, #64 \\ st1w {za2h.s[w12, 0]}, p0, [x17] \\ add x17, x17, #64 \\ st1w {za3h.s[w12, 0]}, p0, [x17] \\ add x16, x16, %[ys] \\ add w12, w12, #1 \\ cmp x12, %[rows] \\ b.ne 2b \\ smstop : : [k] "r" (k), [xa] "r" (xa), [w0] "r" (w[0]), [w1] "r" (w[1]), [w2] "r" (w[2]), [w3] "r" (w[3]), [y] "r" (y), [ys] "r" (y_stride * @sizeOf(f32)), [rows] "r" (rows), : .{ .memory = true, .nzcv = true, .p0 = true, .x9 = true, .x10 = true, .x11 = true, .x12 = true, .x13 = true, .x14 = true, .x15 = true, .x16 = true, .x17 = true, .v0 = true, .v1 = true, .v2 = true, .v3 = true, .v4 = true, .v5 = true, .v6 = true, .v7 = true, .v8 = true, .v9 = true, .v10 = true, .v11 = true, .v12 = true, .v13 = true, .v14 = true, .v15 = true, .v16 = true, .v17 = true, .v18 = true, .v19 = true, .v20 = true, .v21 = true, .v22 = true, .v23 = true, .v24 = true, .v25 = true, .v26 = true, .v27 = true, .v28 = true, .v29 = true, .v30 = true, .v31 = true, });}
/// streaming vector length in bytespub fn svlBytes() usize { return asm volatile ("rdsvl %[out], #1" : [out] "=r" (-> usize), );}