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.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538//! fp32 kernels for a BERT encoder over packed (unpadded) token rows.
const std = @import("std");const Allocator = std.mem.Allocator;const sme = @import("sme.zig");const bf16 = @import("bf16.zig");const int8 = @import("int8.zig");
/// output columns per packed weight panelpub const nr = 16;const V = @Vector(nr, f32);/// token rows per microkernel tileconst mr = 4;/// rows and reduction depth per cache block: a 32×256 activation block (32 KiB)/// and a 256×16 weight panel slice (16 KiB) both stay in L1const row_block = 32;const k_block = 256;
/// A linear layer y = x·Wᵀ + b with W prepacked into K-major panels of `nr`/// output columns: panels[p][k][j] = W[p*nr + j][k].pub const Linear = struct { in: usize, out: usize, panels: []V = &.{}, bias: []f32, /// bf16 pair panels (bf16.packWeights); when non-empty, forward uses them bf16_panels: []u32 = &.{}, /// run the AVX-512 kernel rather than its Zig emulation bf16_native: bool = false, /// int8 weight pair panels (int8.packWeights) and per-column scales int8_panels: []u32 = &.{}, int8_scales: []f32 = &.{}, int8_native: bool = false,
/// `w` is row-major [out][in], as PyTorch stores nn.Linear weights. pub fn pack(gpa: Allocator, w: []const f32, bias: []const f32, in: usize, out: usize) !Linear { std.debug.assert(out % nr == 0 and w.len == in * out and bias.len == out); const panels = try gpa.alloc(V, out / nr * in); errdefer gpa.free(panels); for (0..out / nr) |p| { for (0..in) |k| { var v: V = undefined; inline for (0..nr) |j| v[j] = w[(p * nr + j) * in + k]; panels[p * in + k] = v; } } return .{ .in = in, .out = out, .panels = panels, .bias = try gpa.dupe(f32, bias) }; }
/// Pack W as bf16 pairs. `native` selects the AVX-512 kernel (the caller /// has checked bf16.selfTest); otherwise the Zig emulation runs. pub fn packBf16(gpa: Allocator, w: []const f32, bias: []const f32, in: usize, out: usize, native: bool) !Linear { std.debug.assert(out % bf16.cols == 0 and in <= max_in and w.len == in * out and bias.len == out); const panels = try gpa.alloc(u32, out * bf16.pairs(in)); errdefer gpa.free(panels); bf16.packWeights(panels, w, in, out); return .{ .in = in, .out = out, .bf16_panels = panels, .bf16_native = native, .bias = try gpa.dupe(f32, bias) }; }
/// Quantize W to int8 per output column. `native` selects the AVX2 kernel /// (the caller has checked int8.selfTest); otherwise the Zig emulation runs. pub fn packInt8(gpa: Allocator, w: []const f32, bias: []const f32, in: usize, out: usize, native: bool) !Linear { std.debug.assert(out % int8.cols == 0 and in <= max_in and w.len == in * out and bias.len == out); const panels = try gpa.alloc(u32, out * int8.pairs(in)); errdefer gpa.free(panels); const scales = try gpa.alloc(f32, out); errdefer gpa.free(scales); int8.packWeights(panels, scales, w, in, out); return .{ .in = in, .out = out, .int8_panels = panels, .int8_scales = scales, .int8_native = native, .bias = try gpa.dupe(f32, bias) }; }
pub fn deinit(self: *Linear, gpa: Allocator) void { gpa.free(self.panels); gpa.free(self.bf16_panels); gpa.free(self.int8_panels); gpa.free(self.int8_scales); gpa.free(self.bias); }
/// y[rows][out] = x[rows][in]·Wᵀ + b pub fn forward(self: *const Linear, x: []const f32, y: []f32, rows: usize) void { std.debug.assert(x.len >= rows * self.in and y.len >= rows * self.out); if (self.bf16_panels.len != 0) return self.forwardBf16(x, y, rows); if (self.int8_panels.len != 0) return self.forwardInt8(x, y, rows); if (sme.available and self.out % sme.cols == 0 and self.in <= sme_max_in and sme.svlBytes() == sme.lanes * @sizeOf(f32)) return self.forwardSme(x, y, rows); self.forwardNeon(x, y, rows); }
const sme_max_in = 2048; /// widest input the bf16 and int8 paths stack-buffer (bge-base's FFN is 3072) pub const max_in = 4096;
fn forwardInt8(self: *const Linear, x: []const f32, y: []f32, rows: usize) void { const kp = int8.pairs(self.in); const panel_len = kp * int8.cols; var xp: [int8.pairs(max_in) * int8.rows]u32 = undefined; var xs: [int8.rows]f32 = undefined; var r0: usize = 0; while (r0 < rows) : (r0 += int8.rows) { const m = @min(int8.rows, rows - r0); int8.packRows(xp[0 .. kp * int8.rows], &xs, x[r0 * self.in ..], self.in, m); for (0..self.out / int8.cols) |p| { var acc: [int8.rows][int8.cols]i32 align(32) = undefined; const w = self.int8_panels[p * panel_len ..][0..panel_len]; if (self.int8_native) int8.kernel(&xp, w.ptr, kp, &acc) else int8.emulate(xp[0 .. kp * int8.rows], w, kp, &acc); const ws: V = self.int8_scales[p * int8.cols ..][0..int8.cols].*; const b: V = self.bias[p * int8.cols ..][0..int8.cols].*; for (0..m) |i| { const a: V = @floatFromInt(@as(@Vector(int8.cols, i32), acc[i])); y[(r0 + i) * self.out + p * int8.cols ..][0..int8.cols].* = a * ws * @as(V, @splat(xs[i])) + b; } } } }
fn forwardBf16(self: *const Linear, x: []const f32, y: []f32, rows: usize) void { const kp = bf16.pairs(self.in); const panel_len = kp * bf16.panel_width; var xp: [bf16.pairs(max_in) * bf16.rows]u32 = undefined; var r0: usize = 0; while (r0 < rows) : (r0 += bf16.rows) { const m = @min(bf16.rows, rows - r0); bf16.packRows(xp[0 .. kp * bf16.rows], x[r0 * self.in ..], self.in, m); for (0..self.out / bf16.cols) |g| { var out_block: [bf16.rows][bf16.cols]f32 align(64) = undefined; const w0 = self.bf16_panels[2 * g * panel_len ..][0..panel_len]; const w1 = self.bf16_panels[(2 * g + 1) * panel_len ..][0..panel_len]; const b = self.bias[g * bf16.cols ..][0..bf16.cols]; if (self.bf16_native) bf16.kernel(&xp, w0.ptr, w1.ptr, kp, b.ptr, &out_block) else bf16.emulate(xp[0 .. kp * bf16.rows], w0, w1, kp, b, &out_block); for (0..m) |i| @memcpy(y[(r0 + i) * self.out + g * bf16.cols ..][0..bf16.cols], &out_block[i]); } } }
fn forwardSme(self: *const Linear, x: []const f32, y: []f32, rows: usize) void { const k_len = self.in; var xa: [sme_max_in * sme.lanes]f32 = undefined; var r0: usize = 0; while (r0 < rows) : (r0 += sme.lanes) { const m = @min(sme.lanes, rows - r0); for (0..k_len) |k| { for (0..sme.lanes) |i| xa[k * sme.lanes + i] = if (i < m) x[(r0 + i) * k_len + k] else 0; } for (0..self.out / sme.cols) |g| { var w: [4][*]const f32 = undefined; inline for (0..4) |t| w[t] = @ptrCast(self.panels[(4 * g + t) * k_len ..].ptr); sme.block(&xa, w, k_len, y[r0 * self.out + g * sme.cols ..].ptr, self.out, m); } } for (0..rows) |r| { const row = y[r * self.out ..][0..self.out]; var c: usize = 0; while (c < self.out) : (c += nr) row[c..][0..nr].* = @as(V, row[c..][0..nr].*) + @as(V, self.bias[c..][0..nr].*); } }
pub fn forwardNeon(self: *const Linear, x: []const f32, y: []f32, rows: usize) void { var r0: usize = 0; while (r0 < rows) : (r0 += row_block) { const r1 = @min(r0 + row_block, rows); var k0: usize = 0; while (k0 < self.in) : (k0 += k_block) { const k1 = @min(k0 + k_block, self.in); for (0..self.out / nr) |p| { var r = r0; while (r + mr <= r1) : (r += mr) self.tile(mr, x, y, r, p, k0, k1); switch (r1 - r) { 0 => {}, 1 => self.tile(1, x, y, r, p, k0, k1), 2 => self.tile(2, x, y, r, p, k0, k1), 3 => self.tile(3, x, y, r, p, k0, k1), 4 => self.tile(4, x, y, r, p, k0, k1), 5 => self.tile(5, x, y, r, p, k0, k1), else => unreachable, } } } } }
/// accumulate x[r0..r0+m][k0..k1] against panel p into y, seeding from the /// bias on the first k block fn tile(self: *const Linear, comptime m: usize, x: []const f32, y: []f32, r0: usize, p: usize, k0: usize, k1: usize) void { const k_len = self.in; const panel = self.panels[p * k_len ..][k0..k1]; var acc: [m]V = undefined; inline for (0..m) |i| acc[i] = if (k0 == 0) self.bias[p * nr ..][0..nr].* else y[(r0 + i) * self.out + p * nr ..][0..nr].*; var k = k0; while (k + 4 <= k1) : (k += 4) { var xs: [m]@Vector(4, f32) = undefined; inline for (0..m) |i| xs[i] = x[(r0 + i) * k_len + k ..][0..4].*; inline for (0..4) |j| { const w = panel[k - k0 + j]; inline for (0..m) |i| acc[i] = @mulAdd(V, @splat(xs[i][j]), w, acc[i]); } } while (k < k1) : (k += 1) { const w = panel[k - k0]; inline for (0..m) |i| acc[i] = @mulAdd(V, @splat(x[(r0 + i) * k_len + k]), w, acc[i]); } inline for (0..m) |i| y[(r0 + i) * self.out + p * nr ..][0..nr].* = acc[i]; }};
/// x = LayerNorm(x + residual) row by row, in place on x.pub fn addLayerNorm(x: []f32, residual: []const f32, gamma: []const f32, beta: []const f32, rows: usize, eps: f32) void { const h = gamma.len; for (0..rows) |r| { const row = x[r * h ..][0..h]; const res = residual[r * h ..][0..h]; for (row, res) |*a, b| a.* += b; layerNormRow(row, gamma, beta, eps); }}
pub fn layerNormRow(row: []f32, gamma: []const f32, beta: []const f32, eps: f32) void { std.debug.assert(row.len % nr == 0); const n: f32 = @floatFromInt(row.len); var acc: V = @splat(0); var i: usize = 0; while (i < row.len) : (i += nr) acc += @as(V, row[i..][0..nr].*); const mean = @reduce(.Add, acc) / n; const m: V = @splat(mean); acc = @splat(0); i = 0; while (i < row.len) : (i += nr) { const d = @as(V, row[i..][0..nr].*) - m; acc = @mulAdd(V, d, d, acc); } const inv: V = @splat(1.0 / @sqrt(@reduce(.Add, acc) / n + eps)); i = 0; while (i < row.len) : (i += nr) row[i..][0..nr].* = (@as(V, row[i..][0..nr].*) - m) * inv * @as(V, gamma[i..][0..nr].*) + @as(V, beta[i..][0..nr].*);}
/// exact (erf) GELU, in placepub fn gelu(x: []f32) void { var i: usize = 0; while (i + nr <= x.len) : (i += nr) { const v: V = x[i..][0..nr].*; x[i..][0..nr].* = v * @as(V, @splat(0.5)) * (@as(V, @splat(1)) + erf(nr, v * @as(V, @splat(std.math.sqrt1_2)))); } while (i < x.len) : (i += 1) { const v: @Vector(1, f32) = .{x[i]}; x[i] = (v * @as(@Vector(1, f32), @splat(0.5)) * (@as(@Vector(1, f32), @splat(1)) + erf(1, v * @as(@Vector(1, f32), @splat(std.math.sqrt1_2)))))[0]; }}
/// Abramowitz & Stegun 7.1.26, |error| < 1.5e-7fn erf(comptime n: usize, x: @Vector(n, f32)) @Vector(n, f32) { const T = @Vector(n, f32); const a = @abs(x); const t = @as(T, @splat(1)) / (@as(T, @splat(1)) + @as(T, @splat(0.3275911)) * a); var poly: T = @splat(1.061405429); poly = poly * t + @as(T, @splat(-1.453152027)); poly = poly * t + @as(T, @splat(1.421413741)); poly = poly * t + @as(T, @splat(-0.284496736)); poly = poly * t + @as(T, @splat(0.254829592)); poly = poly * t; const y = @as(T, @splat(1)) - poly * exp(n, -a * a); return @select(f32, x < @as(T, @splat(0)), -y, y);}
/// head widths the attention kernel is specialized for (bge-small and MiniLM/// use 32; bge-base and other 768-wide BERTs use 64)pub const head_dims = [_]usize{ 32, 64 };const max_head_dim = std.mem.max(usize, &head_dims);
/// queries scored together, so each K and V vector loaded serves this many FMAsconst query_tile = 4;
/// floats of scratch `attention` needs for a sequence of `len` rowspub fn attentionScratch(len: usize) usize { const padded = std.mem.alignForward(usize, len, nr); return max_head_dim * padded + query_tile * padded;}
/// Multi-head self-attention for one sequence of `len` rows. `qkv` rows hold/// [q | k | v] (each `hidden` wide); writes context rows for the first/// `queries` positions to `ctx`./// K is transposed per head so scores accumulate as vectors along the keys.pub fn attention(qkv: []const f32, ctx: []f32, scratch: []f32, len: usize, hidden: usize, heads: usize, queries: usize) void { inline for (head_dims) |hd| { if (hidden == heads * hd) return attentionHeads(hd, qkv, ctx, scratch, len, hidden, heads, queries); } unreachable; // spec.read rejects other head widths
}
fn attentionHeads(comptime hd: usize, qkv: []const f32, ctx: []f32, scratch: []f32, len: usize, hidden: usize, heads: usize, queries: usize) void { std.debug.assert(queries <= len); std.debug.assert(hidden == heads * hd); const padded = std.mem.alignForward(usize, len, nr); const kt = scratch[0 .. hd * padded]; const scores = scratch[hd * padded ..][0 .. query_tile * padded]; const stride = 3 * hidden; for (0..heads) |h| { const off = h * hd; for (0..hd) |d| { const col = kt[d * padded ..][0..padded]; for (0..len) |j| col[j] = qkv[j * stride + hidden + off + d]; @memset(col[len..], 0); } var i: usize = 0; while (i + query_tile <= queries) : (i += query_tile) queryTile(hd, query_tile, qkv, ctx, kt, scores, len, padded, hidden, off, i); switch (queries - i) { 0 => {}, inline 1...query_tile - 1 => |t| queryTile(hd, t, qkv, ctx, kt, scores, len, padded, hidden, off, i), else => unreachable, } }}
/// softmax(q·Kᵀ/√d)·V for queries first..first+t of one head. Query and/// probability scalars are broadcast from memory and the context accumulates/// in 16-lane slices, so a four-query tile fits AVX2's 16 vector registers.fn queryTile(comptime hd: usize, comptime t: usize, qkv: []const f32, ctx: []f32, kt: []const f32, scores: []f32, len: usize, padded: usize, hidden: usize, off: usize, first: usize) void { const stride = 3 * hidden; const scale: f32 = 1.0 / @sqrt(@as(f32, hd)); var c: usize = 0; while (c < padded) : (c += nr) { var acc: [t]V = @splat(@splat(0)); inline for (0..hd) |d| { const kv: V = kt[d * padded + c ..][0..nr].*; inline for (0..t) |a| acc[a] = @mulAdd(V, @splat(qkv[(first + a) * stride + off + d]), kv, acc[a]); } inline for (0..t) |a| scores[a * padded + c ..][0..nr].* = acc[a] * @as(V, @splat(scale)); } var inv: [t]f32 = undefined; inline for (0..t) |a| { const row = scores[a * padded ..][0..len]; var max: f32 = -std.math.inf(f32); for (row) |sc| max = @max(max, sc); inv[a] = 1.0 / expShiftedSum(row, max); } inline for (0..hd / nr) |slice| { var out: [t]V = @splat(@splat(0)); for (0..len) |j| { const v: V = qkv[j * stride + 2 * hidden + off + slice * nr ..][0..nr].*; inline for (0..t) |a| out[a] = @mulAdd(V, @splat(scores[a * padded + j]), v, out[a]); } inline for (0..t) |a| ctx[(first + a) * hidden + off + slice * nr ..][0..nr].* = out[a] * @as(V, @splat(inv[a])); }}
/// s[i] = exp(s[i] - shift), returning the sumfn expShiftedSum(s: []f32, shift: f32) f32 { var acc: V = @splat(0); var i: usize = 0; while (i + nr <= s.len) : (i += nr) { const e = exp(nr, @as(V, s[i..][0..nr].*) - @as(V, @splat(shift))); s[i..][0..nr].* = e; acc += e; } var sum = @reduce(.Add, acc); while (i < s.len) : (i += 1) { s[i] = exp(1, .{s[i] - shift})[0]; sum += s[i]; } return sum;}
/// Cephes-style expf: x = n·ln2 + r, a degree-6 polynomial for exp(r), then 2ⁿ/// by exponent bits. About 1 ulp; @exp on vectors lowers to scalar libcalls.pub fn exp(comptime n: usize, x_in: @Vector(n, f32)) @Vector(n, f32) { const T = @Vector(n, f32); const x = @min(@max(x_in, @as(T, @splat(-87.33654))), @as(T, @splat(88.72283))); const k = @round(x * @as(T, @splat(std.math.log2e))); var r = x - k * @as(T, @splat(0.693359375)); r = r - k * @as(T, @splat(-2.12194440e-4)); var p: T = @splat(1.9875691500e-4); p = p * r + @as(T, @splat(1.3981999507e-3)); p = p * r + @as(T, @splat(8.3334519073e-3)); p = p * r + @as(T, @splat(4.1665795894e-2)); p = p * r + @as(T, @splat(1.6666665459e-1)); p = p * r + @as(T, @splat(5.0000001201e-1)); const y = p * r * r + r + @as(T, @splat(1)); const ki: @Vector(n, i32) = @intFromFloat(k); const scale: T = @bitCast((ki + @as(@Vector(n, i32), @splat(127))) << @splat(23)); return y * scale;}
test "linear matches naive matmul for every row remainder" { const gpa = std.testing.allocator; const in = 5; const out = 32; var w: [out * in]f32 = undefined; var b: [out]f32 = undefined; for (&w, 0..) |*v, i| v.* = @as(f32, @floatFromInt(i % 7)) - 3; for (&b, 0..) |*v, i| v.* = @floatFromInt(i); var lin = try Linear.pack(gpa, &w, &b, in, out); defer lin.deinit(gpa);
for (1..8) |rows| { var x: [7 * in]f32 = undefined; for (&x, 0..) |*v, i| v.* = @as(f32, @floatFromInt(i % 5)) * 0.5; var y: [7 * out]f32 = undefined; lin.forward(&x, &y, rows); for (0..rows) |r| for (0..out) |o| { var want = b[o]; for (0..in) |k| want += x[r * in + k] * w[o * in + k]; try std.testing.expectApproxEqAbs(want, y[r * out + o], 1e-4); }; }}
test "linear matches naive matmul across cache blocks, sme and neon" { const gpa = std.testing.allocator; const in = 300; const out = 128; const rows = 37; const w = try gpa.alloc(f32, in * out); defer gpa.free(w); const b = try gpa.alloc(f32, out); defer gpa.free(b); const x = try gpa.alloc(f32, rows * in); defer gpa.free(x); const y = try gpa.alloc(f32, rows * out); defer gpa.free(y); for (w, 0..) |*v, i| v.* = @as(f32, @floatFromInt(i % 13)) * 0.1 - 0.6; for (b, 0..) |*v, i| v.* = @as(f32, @floatFromInt(i % 5)) - 2; for (x, 0..) |*v, i| v.* = @as(f32, @floatFromInt(i % 11)) * 0.05 - 0.25; var lin = try Linear.pack(gpa, w, b, in, out); defer lin.deinit(gpa);
for ([_]bool{ false, true }) |neon| { @memset(y, std.math.nan(f32)); if (neon) lin.forwardNeon(x, y, rows) else lin.forward(x, y, rows); for (0..rows) |r| for (0..out) |o| { var want: f64 = b[o]; for (0..in) |k| want += @as(f64, x[r * in + k]) * w[o * in + k]; try std.testing.expectApproxEqAbs(@as(f32, @floatCast(want)), y[r * out + o], 1e-3); }; }}
test "bf16 and int8 linears match fp32 within their error across row remainders" { const gpa = std.testing.allocator; const in = 70; const out = 64; const w = try gpa.alloc(f32, in * out); defer gpa.free(w); const b = try gpa.alloc(f32, out); defer gpa.free(b); for (w, 0..) |*v, i| v.* = @sin(@as(f32, @floatFromInt(i)) * 0.3) * 0.2; for (b, 0..) |*v, i| v.* = @as(f32, @floatFromInt(i % 7)) * 0.1; var exact = try Linear.pack(gpa, w, b, in, out); defer exact.deinit(gpa); var approx = try Linear.packBf16(gpa, w, b, in, out, bf16.selfTest()); defer approx.deinit(gpa); var quant = try Linear.packInt8(gpa, w, b, in, out, int8.selfTest()); defer quant.deinit(gpa); for ([_]usize{ 1, 8, 13 }) |rows| { const x = try gpa.alloc(f32, rows * in); defer gpa.free(x); for (x, 0..) |*v, i| v.* = @cos(@as(f32, @floatFromInt(i)) * 0.11); const y0 = try gpa.alloc(f32, rows * out); defer gpa.free(y0); const y1 = try gpa.alloc(f32, rows * out); defer gpa.free(y1); exact.forward(x, y0, rows); approx.forward(x, y1, rows); for (y0, y1) |a, c| try std.testing.expectApproxEqAbs(a, c, 0.03); quant.forward(x, y1, rows); for (y0, y1) |a, c| try std.testing.expectApproxEqAbs(a, c, 0.03); }}
test "gelu matches erf reference values" { var x = [_]f32{ -3, -1, -0.5, 0, 0.5, 1, 3 } ++ @as([16]f32, @splat(0.25)); gelu(&x); const want = [_]f32{ -0.0040496, -0.1586553, -0.1542686, 0, 0.3457314, 0.8413447, 2.9959504 }; for (want, x[0..want.len]) |w, g| try std.testing.expectApproxEqAbs(w, g, 1e-6); try std.testing.expectApproxEqAbs(@as(f32, 0.1496766), x[10], 1e-6);}
test "attention matches a naive softmax(qk/sqrt(d))v for each head width" { inline for (head_dims) |head_dim| try checkAttention(head_dim);}
fn checkAttention(comptime head_dim: usize) !void { const gpa = std.testing.allocator; const heads = 2; const hidden = heads * head_dim; for ([_]usize{ 1, 5, 17, 33 }) |len| { const qkv = try gpa.alloc(f32, len * 3 * hidden); defer gpa.free(qkv); const ctx = try gpa.alloc(f32, len * hidden); defer gpa.free(ctx); const scratch = try gpa.alloc(f32, attentionScratch(len)); defer gpa.free(scratch); for (qkv, 0..) |*v, i| v.* = @sin(@as(f32, @floatFromInt(i)) * 0.37); attention(qkv, ctx, scratch, len, hidden, heads, len);
for (0..heads) |h| for (0..len) |i| { var sc: [64]f64 = undefined; var max: f64 = -std.math.inf(f64); for (0..len) |j| { var dotp: f64 = 0; for (0..head_dim) |d| dotp += @as(f64, qkv[i * 3 * hidden + h * head_dim + d]) * qkv[j * 3 * hidden + hidden + h * head_dim + d]; sc[j] = dotp / @sqrt(@as(f64, head_dim)); max = @max(max, sc[j]); } var sum: f64 = 0; for (sc[0..len]) |*v| { v.* = @exp(v.* - max); sum += v.*; } for (0..head_dim) |d| { var want: f64 = 0; for (0..len) |j| want += sc[j] / sum * qkv[j * 3 * hidden + 2 * hidden + h * head_dim + d]; try std.testing.expectApproxEqAbs(@as(f32, @floatCast(want)), ctx[i * hidden + h * head_dim + d], 1e-5); } }; }}
test "vector exp stays within a few ulp of std exp" { var worst: f32 = 0; var x: f32 = -87; while (x < 88) : (x += 0.01137) { const got = exp(1, .{x})[0]; const want = @exp(x); worst = @max(worst, @abs(got - want) / want); } try std.testing.expect(worst < 4e-7);}
test "layernorm normalizes to zero mean unit variance before affine" { var row: [nr]f32 = undefined; for (&row, 0..) |*v, i| v.* = @floatFromInt(i % 4 + 1); const g: [nr]f32 = @splat(1); const b: [nr]f32 = @splat(0); layerNormRow(&row, &g, &b, 1e-12); try std.testing.expectApproxEqAbs(@as(f32, -1.3416408), row[0], 1e-5); try std.testing.expectApproxEqAbs(@as(f32, 1.3416408), row[3], 1e-5);}