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.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315//! bf16 matmul for AVX-512 BF16 (e.g. AMD Zen 4), dispatched at runtime.//!//! Weights and activations are rounded to bf16 and packed as pairs along K, the//! layout VDPBF16PS consumes: each 32-bit lane holds (k even, k odd), and one//! instruction accumulates both products of 16 lanes into fp32. The binary is//! built for the x86_64_v3 baseline; the assembler accepts AVX-512 anyway, and//! `available()` checks CPUID and XCR0 before anything runs it.//!//! `emulate` computes the same rounded products in plain Zig, so the packing and//! the numerics are testable anywhere, and `selfTest` checks the kernel against//! it on the machine that will run it.
const std = @import("std");const builtin = @import("builtin");
/// The kernel and CPU probes need LLVM's assembler: Zig's self-hosted x86/// backend (the Debug default on x86_64 Linux) rejects them, so there the/// kernel is simply absent and callers fall back to f32.pub const compiled = builtin.cpu.arch == .x86_64 and builtin.zig_backend == .stage2_llvm;
/// rows per microkernel callpub const rows = 8;/// output columns per call: two 16-lane panelspub const cols = 32;/// output columns per packed panelpub const panel_width = 16;
pub fn fromF32(f: f32) u16 { const bits: u32 = @bitCast(f); if (std.math.isNan(f)) return @intCast((bits >> 16) | 0x40); const round = 0x7FFF + ((bits >> 16) & 1); return @intCast((bits +% round) >> 16);}
pub fn toF32(h: u16) f32 { return @bitCast(@as(u32, h) << 16);}
fn pair(lo: f32, hi: f32) u32 { return @as(u32, fromF32(lo)) | (@as(u32, fromF32(hi)) << 16);}
pub fn pairs(in: usize) usize { return (in + 1) / 2;}
/// Pack row-major W[out][in] into panels of 16 output columns, K-pair-major:/// out[(p * pairs + kk) * 16 + j] = (W[p*16+j][2kk], W[p*16+j][2kk+1]).pub fn packWeights(dst: []u32, w: []const f32, in: usize, out: usize) void { const kp = pairs(in); std.debug.assert(out % panel_width == 0 and dst.len == out / panel_width * kp * panel_width); for (0..out / panel_width) |p| for (0..kp) |kk| for (0..panel_width) |j| { const row = w[(p * panel_width + j) * in ..][0..in]; const k = 2 * kk; dst[(p * kp + kk) * panel_width + j] = pair(row[k], if (k + 1 < in) row[k + 1] else 0); };}
/// Pack up to 8 activation rows, K-pair-major: dst[kk * 8 + i] = (x[i][2kk], x[i][2kk+1])./// Missing rows are zero.pub fn packRows(dst: []u32, x: []const f32, in: usize, m: usize) void { const kp = pairs(in); std.debug.assert(m <= rows and dst.len >= kp * rows); for (0..kp) |kk| for (0..rows) |i| { const k = 2 * kk; dst[kk * rows + i] = if (i >= m) 0 else pair(x[i * in + k], if (k + 1 < in) x[i * in + k + 1] else 0); };}
/// out[i][c] = bias[c] + Σ_kk dot(xp pair, weight pair), for the two panels at w0, w1.pub fn emulate(xp: []const u32, w0: []const u32, w1: []const u32, kp: usize, bias: []const f32, out: *[rows][cols]f32) void { for (0..rows) |i| for (0..cols) |c| { const w = if (c < panel_width) w0 else w1; const j = c % panel_width; var acc: f32 = bias[c]; for (0..kp) |kk| { const a = xp[kk * rows + i]; const b = w[kk * panel_width + j]; acc += toF32(@truncate(a)) * toF32(@truncate(b)) + toF32(@truncate(a >> 16)) * toF32(@truncate(b >> 16)); } out[i][c] = acc; };}
pub fn available() bool { if (!compiled) return false; const leaf1 = cpuid(1, 0); const osxsave = leaf1.ecx & (1 << 27) != 0; if (!osxsave) return false; // XCR0: SSE, AVX, opmask, ZMM_Hi256 and Hi16_ZMM state enabled by the OS const xcr0 = xgetbv(); if (xcr0 & 0xE6 != 0xE6) return false; const leaf7 = cpuid(7, 0); const avx512f = leaf7.ebx & (1 << 16) != 0; const avx512bw = leaf7.ebx & (1 << 30) != 0; const avx512vl = leaf7.ebx & (1 << 31) != 0; const bf16 = cpuid(7, 1).eax & (1 << 5) != 0; return avx512f and avx512bw and avx512vl and bf16;}
const Regs = struct { eax: u32, ebx: u32, ecx: u32, edx: u32 };
fn cpuid(leaf: u32, sub: u32) Regs { if (comptime !compiled) return .{ .eax = 0, .ebx = 0, .ecx = 0, .edx = 0 } else return cpuidX86(leaf, sub);}
fn cpuidX86(leaf: u32, sub: u32) Regs { var eax: u32 = undefined; var ebx: u32 = undefined; var ecx: u32 = undefined; var edx: u32 = undefined; asm volatile ("cpuid" : [eax] "={eax}" (eax), [ebx] "={ebx}" (ebx), [ecx] "={ecx}" (ecx), [edx] "={edx}" (edx), : [leaf] "{eax}" (leaf), [sub] "{ecx}" (sub), ); return .{ .eax = eax, .ebx = ebx, .ecx = ecx, .edx = edx };}
fn xgetbv() u64 { if (comptime !compiled) return 0 else return xgetbvX86();}
fn xgetbvX86() u64 { var lo: u32 = undefined; var hi: u32 = undefined; asm volatile ("xgetbv" : [lo] "={eax}" (lo), [hi] "={edx}" (hi), : [idx] "{ecx}" (@as(u32, 0)), ); return @as(u64, hi) << 32 | lo;}
/// The AVX-512 BF16 microkernel: same contract as `emulate`. zmm0..15 hold the/// 8×32 fp32 accumulators (row i in zmm(2i), zmm(2i+1)), seeded from the bias.pub fn kernel(xp: [*]const u32, w0: [*]const u32, w1: [*]const u32, kp: usize, bias: [*]const f32, out: *[rows][cols]f32) void { if (comptime !compiled) unreachable else kernelX86(xp, w0, w1, kp, bias, out);}
fn kernelX86(xp: [*]const u32, w0: [*]const u32, w1: [*]const u32, kp: usize, bias: [*]const f32, out: *[rows][cols]f32) void { std.debug.assert(kp > 0); asm volatile ( \\ mov %[xp], %%r8 \\ mov %[w0], %%r9 \\ mov %[w1], %%r10 \\ mov %[kp], %%r11 \\ vmovups (%[bias]), %%zmm0 \\ vmovups 64(%[bias]), %%zmm1 \\ vmovaps %%zmm0, %%zmm2 \\ vmovaps %%zmm1, %%zmm3 \\ vmovaps %%zmm0, %%zmm4 \\ vmovaps %%zmm1, %%zmm5 \\ vmovaps %%zmm0, %%zmm6 \\ vmovaps %%zmm1, %%zmm7 \\ vmovaps %%zmm0, %%zmm8 \\ vmovaps %%zmm1, %%zmm9 \\ vmovaps %%zmm0, %%zmm10 \\ vmovaps %%zmm1, %%zmm11 \\ vmovaps %%zmm0, %%zmm12 \\ vmovaps %%zmm1, %%zmm13 \\ vmovaps %%zmm0, %%zmm14 \\ vmovaps %%zmm1, %%zmm15 \\ 1: \\ vmovups (%%r9), %%zmm16 \\ vmovups (%%r10), %%zmm17 \\ vpbroadcastd 0(%%r8), %%zmm18 \\ vpbroadcastd 4(%%r8), %%zmm19 \\ vpbroadcastd 8(%%r8), %%zmm20 \\ vpbroadcastd 12(%%r8), %%zmm21 \\ vdpbf16ps %%zmm16, %%zmm18, %%zmm0 \\ vdpbf16ps %%zmm17, %%zmm18, %%zmm1 \\ vdpbf16ps %%zmm16, %%zmm19, %%zmm2 \\ vdpbf16ps %%zmm17, %%zmm19, %%zmm3 \\ vdpbf16ps %%zmm16, %%zmm20, %%zmm4 \\ vdpbf16ps %%zmm17, %%zmm20, %%zmm5 \\ vdpbf16ps %%zmm16, %%zmm21, %%zmm6 \\ vdpbf16ps %%zmm17, %%zmm21, %%zmm7 \\ vpbroadcastd 16(%%r8), %%zmm18 \\ vpbroadcastd 20(%%r8), %%zmm19 \\ vpbroadcastd 24(%%r8), %%zmm20 \\ vpbroadcastd 28(%%r8), %%zmm21 \\ vdpbf16ps %%zmm16, %%zmm18, %%zmm8 \\ vdpbf16ps %%zmm17, %%zmm18, %%zmm9 \\ vdpbf16ps %%zmm16, %%zmm19, %%zmm10 \\ vdpbf16ps %%zmm17, %%zmm19, %%zmm11 \\ vdpbf16ps %%zmm16, %%zmm20, %%zmm12 \\ vdpbf16ps %%zmm17, %%zmm20, %%zmm13 \\ vdpbf16ps %%zmm16, %%zmm21, %%zmm14 \\ vdpbf16ps %%zmm17, %%zmm21, %%zmm15 \\ add $64, %%r9 \\ add $64, %%r10 \\ add $32, %%r8 \\ dec %%r11 \\ jnz 1b \\ vmovups %%zmm0, 0(%[out]) \\ vmovups %%zmm1, 64(%[out]) \\ vmovups %%zmm2, 128(%[out]) \\ vmovups %%zmm3, 192(%[out]) \\ vmovups %%zmm4, 256(%[out]) \\ vmovups %%zmm5, 320(%[out]) \\ vmovups %%zmm6, 384(%[out]) \\ vmovups %%zmm7, 448(%[out]) \\ vmovups %%zmm8, 512(%[out]) \\ vmovups %%zmm9, 576(%[out]) \\ vmovups %%zmm10, 640(%[out]) \\ vmovups %%zmm11, 704(%[out]) \\ vmovups %%zmm12, 768(%[out]) \\ vmovups %%zmm13, 832(%[out]) \\ vmovups %%zmm14, 896(%[out]) \\ vmovups %%zmm15, 960(%[out]) \\ vzeroupper : : [xp] "r" (xp), [w0] "r" (w0), [w1] "r" (w1), [kp] "r" (kp), [bias] "r" (bias), [out] "r" (out), : .{ .memory = true, .cc = true, .r8 = true, .r9 = true, .r10 = true, .r11 = true, .zmm0 = true, .zmm1 = true, .zmm2 = true, .zmm3 = true, .zmm4 = true, .zmm5 = true, .zmm6 = true, .zmm7 = true, .zmm8 = true, .zmm9 = true, .zmm10 = true, .zmm11 = true, .zmm12 = true, .zmm13 = true, .zmm14 = true, .zmm15 = true, .zmm16 = true, .zmm17 = true, .zmm18 = true, .zmm19 = true, .zmm20 = true, .zmm21 = true, });}
/// Run the kernel against `emulate` on a synthetic block. The caller only/// enables the kernel when this returns true.pub fn selfTest() bool { if (!available()) return false; const in = 96; const kp = comptime pairs(in); var w: [cols * in]f32 = undefined; var x: [rows * in]f32 = undefined; var bias: [cols]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); for (&bias, 0..) |*v, i| v.* = @as(f32, @floatFromInt(i)) * 0.01; var wp: [cols * kp]u32 = undefined; packWeights(&wp, &w, in, cols); var xp: [kp * rows]u32 = undefined; packRows(&xp, &x, in, rows); var want: [rows][cols]f32 = undefined; var got: [rows][cols]f32 align(64) = undefined; emulate(&xp, wp[0 .. kp * panel_width], wp[kp * panel_width ..], kp, &bias, &want); kernel(&xp, wp[0..].ptr, wp[kp * panel_width ..].ptr, kp, &bias, &got); for (want, got) |wr, gr| for (wr, gr) |a, b| { if (!(@abs(a - b) <= 1e-3 * @max(1, @abs(a)))) return false; }; return true;}
test "bf16 rounding is round-to-nearest-even" { try std.testing.expectEqual(@as(u16, 0x3F80), fromF32(1.0)); try std.testing.expectEqual(@as(u16, 0x3F80), fromF32(@bitCast(@as(u32, 0x3F808000)))); try std.testing.expectEqual(@as(u16, 0x3F82), fromF32(@bitCast(@as(u32, 0x3F818000)))); try std.testing.expectEqual(@as(u16, 0x3F81), fromF32(@bitCast(@as(u32, 0x3F808001)))); try std.testing.expectEqual(@as(f32, -2.0), toF32(fromF32(-2.0)));}
test "emulated bf16 block tracks the fp32 product within bf16 error" { const in = 10; const kp = comptime pairs(in); var w: [cols * in]f32 = undefined; var x: [3 * in]f32 = undefined; var bias: [cols]f32 = undefined; for (&w, 0..) |*v, i| v.* = @sin(@as(f32, @floatFromInt(i))); for (&x, 0..) |*v, i| v.* = @cos(@as(f32, @floatFromInt(i)) * 0.5); for (&bias, 0..) |*v, i| v.* = @floatFromInt(i); var wp: [cols * kp]u32 = undefined; packWeights(&wp, &w, in, cols); var xp: [kp * rows]u32 = undefined; packRows(&xp, &x, in, 3); var got: [rows][cols]f32 = undefined; emulate(&xp, wp[0 .. kp * panel_width], wp[kp * panel_width ..], kp, &bias, &got); for (0..3) |r| for (0..cols) |c| { var want: f32 = bias[c]; for (0..in) |k| want += x[r * in + k] * w[c * in + k]; try std.testing.expectApproxEqAbs(want, got[r][c], 0.05); }; for (3..rows) |r| for (0..cols) |c| try std.testing.expectEqual(bias[c], got[r][c]);}
test "the kernel self test passes exactly when AVX-512 BF16 is available" { try std.testing.expectEqual(available(), selfTest());}