//! 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 16 pub const rows = 6; /// output columns per call and per packed panel pub const cols = 16; pub const weight_max = 127; pub const act_max = 2047; /// K beyond this could overflow the int32 accumulators pub 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 accepts pub 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 int32 pub 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()); }