const std = @import("std"); /// BF16 to F32 conversion and tensor utilities for weight loading. /// Weights are stored as BF16 (u16) and converted to F32 on the fly. /// Convert a single BF16 value (u16) to F32. pub inline fn bf16ToF32(v: u16) f32 { // BF16 is the upper 16 bits of F32 const bits: u32 = @as(u32, v) << 16; return @bitCast(bits); } /// Convert a single F16 value (u16 bits) to F32. pub inline fn f16ToF32(bits: u16) f32 { return @floatCast(@as(f16, @bitCast(bits))); } /// Convert a single F32 to BF16 (truncation, no rounding). pub inline fn f32ToBf16(v: f32) u16 { const bits: u32 = @bitCast(v); return @intCast(bits >> 16); } /// A view into a weight tensor stored as BF16. /// Does not own the data — points into the mmap'd safetensors buffer. pub const Tensor = struct { /// Raw BF16 data (u16 elements) data_bf16: ?[]const u16 = null, /// Raw F32 data (some tensors like A_log are stored as F32) data_f32: ?[]const f32 = null, /// INT4 group-quantized data (Q4MF format) data_q4: ?Q4Data = null, /// Shape dimensions shape: [4]u32 = .{ 0, 0, 0, 0 }, /// Number of dimensions ndim: u8 = 0, pub fn numel(self: *const Tensor) usize { if (self.ndim == 0) return 0; var n: usize = 1; for (self.shape[0..self.ndim]) |d| { n *= d; } return n; } /// Get a single F32 value from the tensor by flat index. pub inline fn get(self: *const Tensor, idx: usize) f32 { if (self.data_f32) |f| return f[idx]; if (self.data_bf16) |b| return bf16ToF32(b[idx]); if (self.data_q4) |q| return q.get(idx); return 0.0; } /// Perform dot product: sum(self[offset..offset+len] * other[0..len]) /// where self is the weight tensor and other is an f32 slice. /// Uses SIMD acceleration for all data formats. pub fn dot(self: *const Tensor, offset: usize, other: []const f32) f32 { const n = other.len; if (self.data_f32) |f| { return simdDotF32(f[offset .. offset + n], other); } if (self.data_bf16) |b| { return simdDotBf16(b[offset .. offset + n], other); } if (self.data_q4) |q| { return q.dot(offset, other); } return 0.0; } }; /// INT4 group-quantized data (Q4MF format). /// Each group of `group_size` elements is stored as: /// scale: f16 (2 bytes) + zero: f16 (2 bytes) + packed nibbles (group_size/2 bytes) /// Dequantization: value = nibble * scale + zero pub const Q4Data = struct { raw: []const u8, num_elements: u32, group_size: u32, const GROUP_HEADER: usize = 4; // 2 bytes scale + 2 bytes zero const NIBBLES_PER_BYTE: usize = 2; fn groupBytes(self: *const Q4Data) usize { return GROUP_HEADER + self.group_size / NIBBLES_PER_BYTE; } pub inline fn get(self: *const Q4Data, idx: usize) f32 { const gb = self.groupBytes(); const group_idx = idx / self.group_size; const within = idx % self.group_size; const group_offset = group_idx * gb; const scale = f16ToF32(std.mem.readInt(u16, self.raw[group_offset..][0..2], .little)); const zero = f16ToF32(std.mem.readInt(u16, self.raw[group_offset + 2 ..][0..2], .little)); const byte_idx = within / 2; const packed_byte = self.raw[group_offset + GROUP_HEADER + byte_idx]; const nibble: u8 = if (within % 2 == 0) packed_byte & 0x0F else packed_byte >> 4; return @as(f32, @floatFromInt(nibble)) * scale + zero; } pub fn dot(self: *const Q4Data, offset: usize, other: []const f32) f32 { const n = other.len; const gs = self.group_size; const gb = self.groupBytes(); var sum: f32 = 0.0; var elem_idx = offset; var other_idx: usize = 0; while (other_idx < n) { const group_idx = elem_idx / gs; const within_start = elem_idx % gs; const group_offset = group_idx * gb; const remaining_in_group = gs - within_start; const chunk = @min(remaining_in_group, n - other_idx); const scale = f16ToF32(std.mem.readInt(u16, self.raw[group_offset..][0..2], .little)); const zero = f16ToF32(std.mem.readInt(u16, self.raw[group_offset + 2 ..][0..2], .little)); const nibble_data = self.raw[group_offset + GROUP_HEADER ..][0 .. gs / 2]; // Dequantize chunk into stack buffer and SIMD dot var deq_buf: [256]f32 = undefined; // max group_size=256 const deq = deq_buf[0..chunk]; for (0..chunk) |i| { const wi = within_start + i; const byte_idx = wi / 2; const b = nibble_data[byte_idx]; const nibble: u8 = if (wi % 2 == 0) b & 0x0F else b >> 4; deq[i] = @as(f32, @floatFromInt(nibble)) * scale + zero; } sum += simdDotF32(deq, other[other_idx .. other_idx + chunk]); elem_idx += chunk; other_idx += chunk; } return sum; } }; /// SIMD-accelerated dot product for two f32 slices. /// Uses @Vector(4, f32) which compiles to WASM SIMD128 or SSE/AVX. pub fn simdDotF32(a: []const f32, b: []const f32) f32 { const n = a.len; const VF32x4 = @Vector(4, f32); var acc: VF32x4 = @splat(0.0); var i: usize = 0; while (i + 4 <= n) : (i += 4) { const va: VF32x4 = a[i..][0..4].*; const vb: VF32x4 = b[i..][0..4].*; acc += va * vb; } var sum = @reduce(.Add, acc); // Scalar tail while (i < n) : (i += 1) { sum += a[i] * b[i]; } return sum; } /// SIMD-accelerated dot product for BF16 weights against f32 input. /// Dequantizes BF16 in chunks of 4, then SIMD multiplies. pub fn simdDotBf16(a_bf16: []const u16, b: []const f32) f32 { const n = a_bf16.len; const VF32x4 = @Vector(4, f32); var acc: VF32x4 = @splat(0.0); var i: usize = 0; while (i + 4 <= n) : (i += 4) { const va = VF32x4{ bf16ToF32(a_bf16[i]), bf16ToF32(a_bf16[i + 1]), bf16ToF32(a_bf16[i + 2]), bf16ToF32(a_bf16[i + 3]), }; const vb: VF32x4 = b[i..][0..4].*; acc += va * vb; } var sum = @reduce(.Add, acc); while (i < n) : (i += 1) { sum += bf16ToF32(a_bf16[i]) * b[i]; } return sum; } /// Matrix-vector multiply: out = weight @ input /// weight shape: [out_dim, in_dim], input shape: [in_dim], out shape: [out_dim] pub fn matVec(weight: *const Tensor, input: []const f32, out: []f32) void { const out_dim = out.len; for (0..out_dim) |i| { out[i] = weight.dot(i * input.len, input); } } /// Matrix-vector multiply and accumulate: out += weight @ input pub fn matVecAdd(weight: *const Tensor, input: []const f32, out: []f32) void { const out_dim = out.len; for (0..out_dim) |i| { out[i] += weight.dot(i * input.len, input); } } /// RMS normalization: out = x * rsqrt(mean(x^2) + eps) * (1 + weight) /// Qwen3.5 uses (1+weight) scaling since weights are initialized to 0. pub fn rmsNorm(out: []f32, x: []const f32, weight: *const Tensor, eps: f32) void { const n = x.len; var sum_sq: f32 = 0.0; for (x) |v| { sum_sq += v * v; } const rms = 1.0 / @sqrt(sum_sq / @as(f32, @floatFromInt(n)) + eps); for (0..n) |i| { out[i] = x[i] * rms * (1.0 + weight.get(i)); } } /// RMS normalization with gating: out = rmsNorm(x) * silu(gate) /// Used in GatedDeltaNet for the output normalization. pub fn rmsNormGated(out: []f32, x: []const f32, gate: []const f32, weight: *const Tensor, eps: f32) void { const n = x.len; var sum_sq: f32 = 0.0; for (x) |v| { sum_sq += v * v; } const rms = 1.0 / @sqrt(sum_sq / @as(f32, @floatFromInt(n)) + eps); for (0..n) |i| { const normed = x[i] * rms * weight.get(i); const g = gate[i]; const silu_g = g * sigmoid(g); out[i] = normed * silu_g; } } pub inline fn silu(x: f32) f32 { return x * sigmoid(x); } pub inline fn sigmoid(x: f32) f32 { return 1.0 / (1.0 + @exp(-x)); } pub inline fn softplus(x: f32) f32 { return @log(1.0 + @exp(x)); } /// L2 normalization of a vector (in-place). Returns the same slice. pub fn l2Norm(x: []f32) void { const eps: f32 = 1e-6; var sum_sq: f32 = 0.0; for (x) |v| { sum_sq += v * v; } const inv_norm = 1.0 / @sqrt(sum_sq + eps); for (x) |*v| { v.* *= inv_norm; } } // ============================================================ // Safetensors parser // ============================================================ pub const SafetensorsError = error{ InvalidFormat, UnsupportedDtype, TensorNotFound, }; pub const TensorInfo = struct { name: []const u8, dtype: []const u8, shape: [4]u32, ndim: u8, data_offset: usize, data_len: usize, }; /// Parse safetensors header and return tensor metadata. /// The file format is: [8 bytes: header_len LE u64] [header_len bytes: JSON] [raw tensor data] pub fn parseSafetensorsHeader( file_data: []const u8, out_tensors: []TensorInfo, ) SafetensorsError!struct { count: usize, data_start: usize } { if (file_data.len < 8) return SafetensorsError.InvalidFormat; const header_len = std.mem.readInt(u64, file_data[0..8], .little); if (8 + header_len > file_data.len) return SafetensorsError.InvalidFormat; const data_start = 8 + @as(usize, @intCast(header_len)); const header_json = file_data[8..data_start]; // Parse the JSON header manually (avoid std.json for WASM compat) var count: usize = 0; var pos: usize = 0; while (pos < header_json.len and count < out_tensors.len) { // Find next key (tensor name) - skip __metadata__ const key_start = findChar(header_json, pos, '"') orelse break; const key_end = findChar(header_json, key_start + 1, '"') orelse break; const key = header_json[key_start + 1 .. key_end]; pos = key_end + 1; if (std.mem.eql(u8, key, "__metadata__")) { // Skip metadata object pos = skipJsonValue(header_json, pos); continue; } // Parse tensor info object const obj_start = findChar(header_json, pos, '{') orelse break; pos = obj_start + 1; var info = TensorInfo{ .name = key, .dtype = "", .shape = .{ 0, 0, 0, 0 }, .ndim = 0, .data_offset = 0, .data_len = 0, }; // Parse fields within the tensor info object var brace_depth: u32 = 1; while (pos < header_json.len and brace_depth > 0) { if (header_json[pos] == '}') { brace_depth -= 1; pos += 1; continue; } const fk_start = findChar(header_json, pos, '"') orelse break; const fk_end = findChar(header_json, fk_start + 1, '"') orelse break; const field_key = header_json[fk_start + 1 .. fk_end]; pos = fk_end + 1; // Skip colon while (pos < header_json.len and header_json[pos] != ':') pos += 1; pos += 1; while (pos < header_json.len and header_json[pos] == ' ') pos += 1; if (std.mem.eql(u8, field_key, "dtype")) { const ds = findChar(header_json, pos, '"') orelse break; const de = findChar(header_json, ds + 1, '"') orelse break; info.dtype = header_json[ds + 1 .. de]; pos = de + 1; } else if (std.mem.eql(u8, field_key, "shape")) { // Parse array of ints const arr_start = findChar(header_json, pos, '[') orelse break; const arr_end = findChar(header_json, arr_start, ']') orelse break; var dim_pos = arr_start + 1; while (dim_pos < arr_end and info.ndim < 4) { while (dim_pos < arr_end and !isDigit(header_json[dim_pos])) dim_pos += 1; if (dim_pos >= arr_end) break; var num: u32 = 0; while (dim_pos < arr_end and isDigit(header_json[dim_pos])) { num = num * 10 + @as(u32, header_json[dim_pos] - '0'); dim_pos += 1; } info.shape[info.ndim] = num; info.ndim += 1; } pos = arr_end + 1; } else if (std.mem.eql(u8, field_key, "data_offsets")) { // Parse [start, end] const arr_start = findChar(header_json, pos, '[') orelse break; const arr_end = findChar(header_json, arr_start, ']') orelse break; var dim_pos = arr_start + 1; var offsets: [2]usize = .{ 0, 0 }; var oi: usize = 0; while (dim_pos < arr_end and oi < 2) { while (dim_pos < arr_end and !isDigit(header_json[dim_pos])) dim_pos += 1; if (dim_pos >= arr_end) break; var num: usize = 0; while (dim_pos < arr_end and isDigit(header_json[dim_pos])) { num = num * 10 + @as(usize, header_json[dim_pos] - '0'); dim_pos += 1; } offsets[oi] = num; oi += 1; } info.data_offset = offsets[0]; info.data_len = offsets[1] - offsets[0]; pos = arr_end + 1; } else { pos = skipJsonValue(header_json, pos); } } // Adjust offset relative to data section info.data_offset += data_start; out_tensors[count] = info; count += 1; } return .{ .count = count, .data_start = data_start }; } /// Create a Tensor view from a TensorInfo and the raw file data. pub fn tensorFromInfo(info: *const TensorInfo, file_data: []const u8) SafetensorsError!Tensor { var t = Tensor{ .shape = info.shape, .ndim = info.ndim, }; const raw = file_data[info.data_offset .. info.data_offset + info.data_len]; if (std.mem.eql(u8, info.dtype, "BF16") or std.mem.eql(u8, info.dtype, "bf16")) { const aligned: []align(@alignOf(u16)) const u8 = @alignCast(raw); t.data_bf16 = std.mem.bytesAsSlice(u16, aligned); } else if (std.mem.eql(u8, info.dtype, "F32") or std.mem.eql(u8, info.dtype, "f32")) { const aligned: []align(@alignOf(f32)) const u8 = @alignCast(raw); t.data_f32 = std.mem.bytesAsSlice(f32, aligned); } else { return SafetensorsError.UnsupportedDtype; } return t; } fn findChar(buf: []const u8, start: usize, c: u8) ?usize { var i = start; while (i < buf.len) : (i += 1) { if (buf[i] == c and (i == 0 or buf[i - 1] != '\\')) return i; } return null; } fn isDigit(c: u8) bool { return c >= '0' and c <= '9'; } fn skipJsonValue(buf: []const u8, start: usize) usize { var pos = start; while (pos < buf.len and (buf[pos] == ' ' or buf[pos] == ':' or buf[pos] == ',')) pos += 1; if (pos >= buf.len) return pos; return switch (buf[pos]) { '"' => { // Skip string pos += 1; while (pos < buf.len) : (pos += 1) { if (buf[pos] == '"' and buf[pos - 1] != '\\') return pos + 1; } return pos; }, '{' => { var depth: u32 = 1; pos += 1; while (pos < buf.len and depth > 0) : (pos += 1) { if (buf[pos] == '{') depth += 1; if (buf[pos] == '}') depth -= 1; } return pos; }, '[' => { var depth: u32 = 1; pos += 1; while (pos < buf.len and depth > 0) : (pos += 1) { if (buf[pos] == '[') depth += 1; if (buf[pos] == ']') depth -= 1; } return pos; }, else => { // Number, bool, null while (pos < buf.len and buf[pos] != ',' and buf[pos] != '}' and buf[pos] != ']') pos += 1; return pos; }, }; } // ============================================================ // Q4MF format parser // ============================================================ pub const Q4MF_MAGIC: u32 = 0x46344D51; pub const Q4MF_VERSION: u32 = 1; pub const Q4MFTensorInfo = struct { name: []const u8, shape: [4]u32, ndim: u8, dtype: u8, // 0=Q4, 1=F32, 2=BF16 data: []const u8, }; pub const Q4MFHeader = struct { num_tensors: u32, group_size: u32, tensors: []Q4MFTensorInfo, }; /// Parse a Q4MF file buffer and return tensor metadata. /// The returned tensor infos point into file_data (zero-copy for data). pub fn parseQ4MF( allocator: std.mem.Allocator, file_data: []const u8, ) !Q4MFHeader { if (file_data.len < 16) return SafetensorsError.InvalidFormat; const magic = std.mem.readInt(u32, file_data[0..4], .little); if (magic != Q4MF_MAGIC) return SafetensorsError.InvalidFormat; const version = std.mem.readInt(u32, file_data[4..8], .little); if (version != Q4MF_VERSION) return SafetensorsError.InvalidFormat; const num_tensors = std.mem.readInt(u32, file_data[8..12], .little); const group_size = std.mem.readInt(u32, file_data[12..16], .little); const tensors = try allocator.alloc(Q4MFTensorInfo, num_tensors); var pos: usize = 16; for (0..num_tensors) |ti| { if (pos + 2 > file_data.len) return SafetensorsError.InvalidFormat; const name_len = std.mem.readInt(u16, file_data[pos..][0..2], .little); pos += 2; const name = file_data[pos .. pos + name_len]; pos += name_len; const ndim = file_data[pos]; pos += 1; var shape: [4]u32 = .{ 0, 0, 0, 0 }; for (0..ndim) |d| { shape[d] = std.mem.readInt(u32, file_data[pos..][0..4], .little); pos += 4; } const dtype = file_data[pos]; pos += 1; const data_len = std.mem.readInt(u32, file_data[pos..][0..4], .little); pos += 4; const data = file_data[pos .. pos + data_len]; pos += data_len; tensors[ti] = .{ .name = name, .shape = shape, .ndim = ndim, .dtype = dtype, .data = data, }; } return .{ .num_tensors = num_tensors, .group_size = group_size, .tensors = tensors, }; } /// Create a Tensor from Q4MF tensor info. /// For Q4 data, creates a Q4Data view. For F32/BF16, copies to aligned allocation. pub fn tensorFromQ4MF(info: *const Q4MFTensorInfo, group_size: u32, allocator: std.mem.Allocator) !Tensor { var t = Tensor{ .shape = info.shape, .ndim = info.ndim, }; var numel: u32 = 1; for (info.shape[0..info.ndim]) |s| numel *= s; switch (info.dtype) { 0 => { // Q4 — zero-copy view into file data t.data_q4 = .{ .raw = info.data, .num_elements = numel, .group_size = group_size, }; }, 1 => { // F32 — copy to aligned allocation const f32_data = try allocator.alloc(f32, numel); for (0..numel) |i| { const bits = std.mem.readInt(u32, info.data[i * 4 ..][0..4], .little); f32_data[i] = @bitCast(bits); } t.data_f32 = f32_data; }, 2 => { // BF16 — copy to aligned allocation const bf16_data = try allocator.alloc(u16, numel); for (0..numel) |i| { bf16_data[i] = std.mem.readInt(u16, info.data[i * 2 ..][0..2], .little); } t.data_bf16 = bf16_data; }, else => return SafetensorsError.UnsupportedDtype, } return t; } test "bf16 conversion" { // 1.0 in f32 is 0x3F800000, in bf16 is 0x3F80 const one = bf16ToF32(0x3F80); try std.testing.expectEqual(@as(f32, 1.0), one); const back = f32ToBf16(1.0); try std.testing.expectEqual(@as(u16, 0x3F80), back); // 0.0 const zero = bf16ToF32(0x0000); try std.testing.expectEqual(@as(f32, 0.0), zero); } test "rms norm" { var x = [_]f32{ 1.0, 2.0, 3.0, 4.0 }; var out: [4]f32 = undefined; // Weight of all zeros means (1+0)=1 scaling (identity weight) var w_data = [_]f32{ 0.0, 0.0, 0.0, 0.0 }; const w = Tensor{ .data_f32 = &w_data, .shape = .{ 4, 0, 0, 0 }, .ndim = 1 }; rmsNorm(&out, &x, &w, 1e-6); // RMS of [1,2,3,4] = sqrt((1+4+9+16)/4) = sqrt(30/4) = sqrt(7.5) const rms_val = @sqrt(30.0 / 4.0); const expected_0 = 1.0 / rms_val; try std.testing.expectApproxEqAbs(expected_0, out[0], 1e-5); } test "q4 dequantization" { // Manually construct a single group of 128 elements: // scale = 1.0 (f16), zero = 0.0 (f16), nibbles all encode their index mod 16 var group_data: [68]u8 = undefined; // scale = 1.0 as f16 = 0x3C00 group_data[0] = 0x00; group_data[1] = 0x3C; // zero = 0.0 as f16 = 0x0000 group_data[2] = 0x00; group_data[3] = 0x00; // Pack nibbles: element 2i = i%16, element 2i+1 = (i+1)%16 // Actually, simpler: set all nibbles to known values // byte[j] = low_nibble | (high_nibble << 4) // elements 0,1 -> byte 0; elements 2,3 -> byte 1; etc. for (0..64) |j| { const low: u8 = @intCast((j * 2) % 16); const high: u8 = @intCast((j * 2 + 1) % 16); group_data[4 + j] = low | (high << 4); } const q4 = Q4Data{ .raw = &group_data, .num_elements = 128, .group_size = 128, }; // Element 0: nibble=0, dequant = 0 * 1.0 + 0.0 = 0.0 try std.testing.expectApproxEqAbs(@as(f32, 0.0), q4.get(0), 1e-3); // Element 1: nibble=1, dequant = 1 * 1.0 + 0.0 = 1.0 try std.testing.expectApproxEqAbs(@as(f32, 1.0), q4.get(1), 1e-3); // Element 14: nibble=14, dequant = 14.0 try std.testing.expectApproxEqAbs(@as(f32, 14.0), q4.get(14), 1e-3); // Element 15: nibble=15, dequant = 15.0 try std.testing.expectApproxEqAbs(@as(f32, 15.0), q4.get(15), 1e-3); // Element 16: nibble=0 (wraps mod 16), dequant = 0.0 try std.testing.expectApproxEqAbs(@as(f32, 0.0), q4.get(16), 1e-3); } test "q4 dequantization with offset" { // Two groups: second group has scale=2.0, zero=10.0 var data: [68 * 2]u8 = undefined; // Group 0: scale=1.0, zero=0.0, all nibbles=5 data[0] = 0x00; data[1] = 0x3C; // f16 1.0 data[2] = 0x00; data[3] = 0x00; // f16 0.0 @memset(data[4..68], 0x55); // nibble 5 in both low and high // Group 1: scale=2.0 (0x4000), zero=10.0 (0x4900) data[68] = 0x00; data[69] = 0x40; // f16 2.0 data[70] = 0x00; data[71] = 0x49; // f16 10.0 @memset(data[72..136], 0x33); // nibble 3 in both low and high const q4 = Q4Data{ .raw = &data, .num_elements = 256, .group_size = 128, }; // Group 0, element 0: 5 * 1.0 + 0.0 = 5.0 try std.testing.expectApproxEqAbs(@as(f32, 5.0), q4.get(0), 1e-2); // Group 1, element 128: 3 * 2.0 + 10.0 = 16.0 try std.testing.expectApproxEqAbs(@as(f32, 16.0), q4.get(128), 1e-1); } test "simd dot f32" { const a = [_]f32{ 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0 }; const b = [_]f32{ 9.0, 8.0, 7.0, 6.0, 5.0, 4.0, 3.0, 2.0, 1.0 }; const result = simdDotF32(&a, &b); // 9+16+21+24+25+24+21+16+9 = 165 try std.testing.expectApproxEqAbs(@as(f32, 165.0), result, 1e-4); } test "simd dot bf16" { // 1.0 in bf16 = 0x3F80, 2.0 = 0x4000 const a = [_]u16{ 0x3F80, 0x4000, 0x3F80, 0x4000 }; const b = [_]f32{ 1.0, 1.0, 1.0, 1.0 }; const result = simdDotBf16(&a, &b); // 1*1 + 2*1 + 1*1 + 2*1 = 6.0 try std.testing.expectApproxEqAbs(@as(f32, 6.0), result, 1e-4); }