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.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295//! Integer matmul for AVX2 hosts without AVX-512 BF16 (e.g. AMD Zen 3).//!//! Weights are int8 with one scale per output column; activations are//! quantized per row to 12 bits. Both are held as int16 and packed as K pairs,//! so VPMADDWD multiplies 16 pairs per instruction into int32 lanes, twice the//! MACs of an fp32 FMA. 12-bit activations are the widest that cannot//! overflow: 2047 * 127 * 2 per pair, times 768 pairs (K = 1536), stays under//! 2^31. Accumulation is exact; the only error is the two roundings.//!//! Like bf16.zig, the kernel needs LLVM's assembler, `emulate` defines the//! numerics in plain Zig, and `selfTest` gates the kernel on the running CPU.
const std = @import("std");const builtin = @import("builtin");
pub const compiled = builtin.cpu.arch == .x86_64 and builtin.zig_backend == .stage2_llvm and std.Target.x86.featureSetHas(builtin.cpu.features, .avx2);
/// rows per microkernel call: 12 int32 accumulators + 2 weight + 2 scratch/// registers fill AVX2's 16pub const rows = 6;/// output columns per call and per packed panelpub const cols = 16;
pub const weight_max = 127;pub const act_max = 2047;/// K beyond this could overflow the int32 accumulatorspub const max_in = 2 * ((1 << 31) / (2 * weight_max * act_max) - 1);
pub fn pairs(in: usize) usize { return (in + 1) / 2;}
fn pack16(lo: i16, hi: i16) u32 { return @as(u32, @as(u16, @bitCast(lo))) | (@as(u32, @as(u16, @bitCast(hi))) << 16);}
fn quant(v: f32, inv_scale: f32, limit: i32) i16 { const q: i32 = @intFromFloat(@round(v * inv_scale)); return @intCast(std.math.clamp(q, -limit, limit));}
/// Quantize row-major W[out][in] per output column and pack panels of 16/// columns, K-pair-major. `scales[o]` receives the dequantization scale.pub fn packWeights(dst: []u32, scales: []f32, w: []const f32, in: usize, out: usize) void { const kp = pairs(in); std.debug.assert(out % cols == 0 and dst.len == out * kp and scales.len == out and in <= max_in); for (0..out) |o| { var m: f32 = 0; for (w[o * in ..][0..in]) |v| m = @max(m, @abs(v)); scales[o] = if (m == 0) 1 else m / weight_max; } for (0..out / cols) |p| for (0..kp) |kk| for (0..cols) |j| { const o = p * cols + j; const row = w[o * in ..][0..in]; const inv = 1 / scales[o]; const k = 2 * kk; dst[(p * kp + kk) * cols + j] = pack16(quant(row[k], inv, weight_max), if (k + 1 < in) quant(row[k + 1], inv, weight_max) else 0); };}
/// Quantize up to 6 activation rows (one scale each) and pack K-pair-major:/// dst[kk * 6 + i]. Missing rows are zero with scale 0. Each row is quantized/// contiguously in 16-lane vectors; a pair is then two adjacent int16s, so/// packing is a transpose of u32s.pub fn packRows(dst: []u32, scales: *[rows]f32, x: []const f32, in: usize, m: usize) void { const kp = pairs(in); std.debug.assert(m <= rows and dst.len >= kp * rows and in <= max_row); var q: [rows][max_row]i16 align(64) = undefined; for (0..rows) |i| { if (i >= m) { scales[i] = 0; @memset(q[i][0 .. 2 * kp], 0); continue; } const row = x[i * in ..][0..in]; scales[i] = quantizeRow(row, q[i][0 .. 2 * kp]); } for (0..kp) |kk| for (0..rows) |i| { dst[kk * rows + i] = @as(u32, @as(u16, @bitCast(q[i][2 * kk]))) | (@as(u32, @as(u16, @bitCast(q[i][2 * kk + 1]))) << 16); };}
/// widest activation row packRows acceptspub const max_row = 4096;
const F = @Vector(16, f32);
/// q = round(row / scale) with scale = max|row| / act_max, zero-padded to/// q.len; returns the scale (0 for an all-zero row)fn quantizeRow(row: []const f32, q: []i16) f32 { var mv: F = @splat(0); var k: usize = 0; while (k + 16 <= row.len) : (k += 16) mv = @max(mv, @abs(@as(F, row[k..][0..16].*))); var mx = @reduce(.Max, mv); while (k < row.len) : (k += 1) mx = @max(mx, @abs(row[k])); @memset(q[row.len..], 0); if (mx == 0) { @memset(q[0..row.len], 0); return 0; } const inv: F = @splat(act_max / mx); // adding and subtracting 1.5 * 2^23 rounds to nearest even for |v| < 2^22 const magic: F = @splat(12582912.0); k = 0; while (k + 16 <= row.len) : (k += 16) { const r = (@as(F, row[k..][0..16].*) * inv + magic) - magic; const c = @min(@max(r, @as(F, @splat(-act_max))), @as(F, @splat(act_max))); const qi: @Vector(16, i16) = @intFromFloat(c); q[k..][0..16].* = qi; } while (k < row.len) : (k += 1) q[k] = quant(row[k], act_max / mx, act_max); return mx / act_max;}
fn low16(v: u32) i32 { return @as(i16, @bitCast(@as(u16, @truncate(v))));}
fn high16(v: u32) i32 { return @as(i16, @bitCast(@as(u16, @truncate(v >> 16))));}
/// out[i][c] = Σ_kk (x pair · w pair), exact int32pub fn emulate(xp: []const u32, w: []const u32, kp: usize, out: *[rows][cols]i32) void { for (0..rows) |i| for (0..cols) |c| { var acc: i32 = 0; for (0..kp) |kk| { const a = xp[kk * rows + i]; const b = w[kk * cols + c]; acc += low16(a) * low16(b) + high16(a) * high16(b); } out[i][c] = acc; };}
/// VPMADDWD microkernel with the same contract as `emulate`. ymm0..11 hold/// the 6×16 int32 accumulators (row i in ymm(2i), ymm(2i+1)).pub fn kernel(xp: [*]const u32, w: [*]const u32, kp: usize, out: *[rows][cols]i32) void { if (comptime !compiled) unreachable else kernelX86(xp, w, kp, out);}
fn kernelX86(xp: [*]const u32, w: [*]const u32, kp: usize, out: *[rows][cols]i32) void { std.debug.assert(kp > 0); asm volatile ( \\ mov %[xp], %%r8 \\ mov %[w], %%r9 \\ mov %[kp], %%r10 \\ vpxor %%ymm0, %%ymm0, %%ymm0 \\ vpxor %%ymm1, %%ymm1, %%ymm1 \\ vpxor %%ymm2, %%ymm2, %%ymm2 \\ vpxor %%ymm3, %%ymm3, %%ymm3 \\ vpxor %%ymm4, %%ymm4, %%ymm4 \\ vpxor %%ymm5, %%ymm5, %%ymm5 \\ vpxor %%ymm6, %%ymm6, %%ymm6 \\ vpxor %%ymm7, %%ymm7, %%ymm7 \\ vpxor %%ymm8, %%ymm8, %%ymm8 \\ vpxor %%ymm9, %%ymm9, %%ymm9 \\ vpxor %%ymm10, %%ymm10, %%ymm10 \\ vpxor %%ymm11, %%ymm11, %%ymm11 \\ 1: \\ vmovdqu (%%r9), %%ymm12 \\ vmovdqu 32(%%r9), %%ymm13 \\ vpbroadcastd 0(%%r8), %%ymm14 \\ vpmaddwd %%ymm12, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm0, %%ymm0 \\ vpmaddwd %%ymm13, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm1, %%ymm1 \\ vpbroadcastd 4(%%r8), %%ymm14 \\ vpmaddwd %%ymm12, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm2, %%ymm2 \\ vpmaddwd %%ymm13, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm3, %%ymm3 \\ vpbroadcastd 8(%%r8), %%ymm14 \\ vpmaddwd %%ymm12, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm4, %%ymm4 \\ vpmaddwd %%ymm13, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm5, %%ymm5 \\ vpbroadcastd 12(%%r8), %%ymm14 \\ vpmaddwd %%ymm12, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm6, %%ymm6 \\ vpmaddwd %%ymm13, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm7, %%ymm7 \\ vpbroadcastd 16(%%r8), %%ymm14 \\ vpmaddwd %%ymm12, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm8, %%ymm8 \\ vpmaddwd %%ymm13, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm9, %%ymm9 \\ vpbroadcastd 20(%%r8), %%ymm14 \\ vpmaddwd %%ymm12, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm10, %%ymm10 \\ vpmaddwd %%ymm13, %%ymm14, %%ymm15 \\ vpaddd %%ymm15, %%ymm11, %%ymm11 \\ add $64, %%r9 \\ add $24, %%r8 \\ dec %%r10 \\ jnz 1b \\ mov %[out], %%r8 \\ vmovdqu %%ymm0, 0(%%r8) \\ vmovdqu %%ymm1, 32(%%r8) \\ vmovdqu %%ymm2, 64(%%r8) \\ vmovdqu %%ymm3, 96(%%r8) \\ vmovdqu %%ymm4, 128(%%r8) \\ vmovdqu %%ymm5, 160(%%r8) \\ vmovdqu %%ymm6, 192(%%r8) \\ vmovdqu %%ymm7, 224(%%r8) \\ vmovdqu %%ymm8, 256(%%r8) \\ vmovdqu %%ymm9, 288(%%r8) \\ vmovdqu %%ymm10, 320(%%r8) \\ vmovdqu %%ymm11, 352(%%r8) \\ vzeroupper : : [xp] "r" (xp), [w] "r" (w), [kp] "r" (kp), [out] "r" (out), : .{ .memory = true, .cc = true, .r8 = true, .r9 = true, .r10 = true, .ymm0 = true, .ymm1 = true, .ymm2 = true, .ymm3 = true, .ymm4 = true, .ymm5 = true, .ymm6 = true, .ymm7 = true, .ymm8 = true, .ymm9 = true, .ymm10 = true, .ymm11 = true, .ymm12 = true, .ymm13 = true, .ymm14 = true, .ymm15 = true, });}
/// Run the kernel against `emulate` on a synthetic block; integer results must match exactly.pub fn selfTest() bool { if (comptime !compiled) return false; const in = 90; const kp = comptime pairs(in); var w: [cols * in]f32 = undefined; var x: [rows * in]f32 = undefined; for (&w, 0..) |*v, i| v.* = @sin(@as(f32, @floatFromInt(i)) * 0.7); for (&x, 0..) |*v, i| v.* = @cos(@as(f32, @floatFromInt(i)) * 0.3) * 5; var wp: [cols * kp]u32 = undefined; var ws: [cols]f32 = undefined; packWeights(&wp, &ws, &w, in, cols); var xp: [kp * rows]u32 = undefined; var xs: [rows]f32 = undefined; packRows(&xp, &xs, &x, in, rows); var want: [rows][cols]i32 = undefined; var got: [rows][cols]i32 align(32) = undefined; emulate(&xp, &wp, kp, &want); kernel(&xp, &wp, kp, &got); return std.meta.eql(want, got);}
test "quantized block reconstructs the fp32 product within quantization error" { const in = 40; const kp = comptime pairs(in); var w: [cols * in]f32 = undefined; var x: [4 * in]f32 = undefined; for (&w, 0..) |*v, i| v.* = @sin(@as(f32, @floatFromInt(i))) * 0.3; for (&x, 0..) |*v, i| v.* = @cos(@as(f32, @floatFromInt(i)) * 0.5) * 3; var wp: [cols * kp]u32 = undefined; var ws: [cols]f32 = undefined; packWeights(&wp, &ws, &w, in, cols); var xp: [kp * rows]u32 = undefined; var xs: [rows]f32 = undefined; packRows(&xp, &xs, &x, in, 4); var acc: [rows][cols]i32 = undefined; emulate(&xp, &wp, kp, &acc); for (0..4) |r| for (0..cols) |c| { var want: f32 = 0; for (0..in) |k| want += x[r * in + k] * w[c * in + k]; const got = @as(f32, @floatFromInt(acc[r][c])) * xs[r] * ws[c]; try std.testing.expectApproxEqAbs(want, got, 0.05); }; for (4..rows) |r| for (0..cols) |c| try std.testing.expectEqual(@as(i32, 0), acc[r][c]);}
test "accumulator bound covers the widest supported layer (bge-base FFN)" { try std.testing.expect(max_in >= 3072);}
test "the kernel self test passes exactly when the kernel is compiled" { try std.testing.expectEqual(compiled, selfTest());}