Something went wrong. Try again.
Anonymize your writing style. Zig WASM engine detects authorship markers, fine-tuned LLM rewrites to remove them. Runs entirely in-browser. fantasma.qstorage.quilibrium.com
wasm privacy qwen zig
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724const std = @import("std");const q = @import("quantized.zig");const cache = @import("kv_cache.zig");const Config = cache.Config;const Tensor = q.Tensor;
// ============================================================// Model weight structure// ============================================================
/// Weights for a full attention layerconst FullAttentionWeights = struct { q_proj: Tensor, // [num_heads * head_dim * 2, hidden_size] (includes gate) k_proj: Tensor, // [num_kv_heads * head_dim, hidden_size] v_proj: Tensor, // [num_kv_heads * head_dim, hidden_size] o_proj: Tensor, // [hidden_size, num_heads * head_dim] q_norm: Tensor, // [head_dim] k_norm: Tensor, // [head_dim]};
/// Weights for a GatedDeltaNet (linear attention) layerconst DeltaNetWeights = struct { in_proj_qkv: Tensor, // [conv_dim, hidden_size] in_proj_z: Tensor, // [value_dim, hidden_size] in_proj_a: Tensor, // [num_heads, hidden_size] in_proj_b: Tensor, // [num_heads, hidden_size] conv1d: Tensor, // [conv_dim, 1, kernel_size] A_log: Tensor, // [num_heads] (F32) dt_bias: Tensor, // [num_heads] norm: Tensor, // [value_head_dim] (F32) out_proj: Tensor, // [hidden_size, value_dim]};
/// Weights for a single decoder layerconst LayerWeights = struct { input_layernorm: Tensor, // [hidden_size] post_attention_layernorm: Tensor, // [hidden_size]
// Attention (one of these is active) full_attn: ?FullAttentionWeights, delta_net: ?DeltaNetWeights,
// MLP gate_proj: Tensor, // [intermediate_size, hidden_size] up_proj: Tensor, // [intermediate_size, hidden_size] down_proj: Tensor, // [hidden_size, intermediate_size]};
/// Complete model weightspub const ModelWeights = struct { embed_tokens: Tensor, // [vocab_size, hidden_size] norm: Tensor, // [hidden_size] layers: [Config.num_layers]LayerWeights,
// RoPE precomputed cos/sin tables rope_cos: []f32, // [max_seq_len, rotary_dim] rope_sin: []f32, // [max_seq_len, rotary_dim]
allocator: std.mem.Allocator,};
// ============================================================// Weight loading from safetensors// ============================================================
pub fn loadWeights(allocator: std.mem.Allocator, model_dir: []const u8, max_seq_len: u32) !ModelWeights { // Build path to safetensors file var path_buf: [1024]u8 = undefined; const path = std.fmt.bufPrint(&path_buf, "{s}/model.safetensors-00001-of-00001.safetensors", .{model_dir}) catch return error.PathTooLong;
// Memory-map the file const file = try std.fs.cwd().openFile(path, .{}); defer file.close(); const file_size = try file.getEndPos();
// Read entire file into memory (mmap would be better but this works cross-platform) const file_data = try allocator.alloc(u8, file_size); // Note: we don't free file_data — it backs all tensor views for the model lifetime const bytes_read = try file.readAll(file_data); if (bytes_read != file_size) return error.IncompleteRead;
// Parse safetensors header var tensor_infos: [512]q.TensorInfo = undefined; const header = try q.parseSafetensorsHeader(file_data, &tensor_infos);
// Helper: find tensor by name const infos = tensor_infos[0..header.count];
var weights: ModelWeights = undefined; weights.allocator = allocator;
// Embedding weights.embed_tokens = try findAndLoad(infos, "model.language_model.embed_tokens.weight", file_data); weights.norm = try findAndLoad(infos, "model.language_model.norm.weight", file_data);
// Layers for (0..Config.num_layers) |layer_idx| { var lw: LayerWeights = undefined;
// Layer norms lw.input_layernorm = try findAndLoadLayer(infos, layer_idx, "input_layernorm.weight", file_data); lw.post_attention_layernorm = try findAndLoadLayer(infos, layer_idx, "post_attention_layernorm.weight", file_data);
// MLP lw.gate_proj = try findAndLoadLayer(infos, layer_idx, "mlp.gate_proj.weight", file_data); lw.up_proj = try findAndLoadLayer(infos, layer_idx, "mlp.up_proj.weight", file_data); lw.down_proj = try findAndLoadLayer(infos, layer_idx, "mlp.down_proj.weight", file_data);
if (Config.layer_is_full_attn[layer_idx]) { lw.full_attn = FullAttentionWeights{ .q_proj = try findAndLoadLayer(infos, layer_idx, "self_attn.q_proj.weight", file_data), .k_proj = try findAndLoadLayer(infos, layer_idx, "self_attn.k_proj.weight", file_data), .v_proj = try findAndLoadLayer(infos, layer_idx, "self_attn.v_proj.weight", file_data), .o_proj = try findAndLoadLayer(infos, layer_idx, "self_attn.o_proj.weight", file_data), .q_norm = try findAndLoadLayer(infos, layer_idx, "self_attn.q_norm.weight", file_data), .k_norm = try findAndLoadLayer(infos, layer_idx, "self_attn.k_norm.weight", file_data), }; lw.delta_net = null; } else { lw.delta_net = DeltaNetWeights{ .in_proj_qkv = try findAndLoadLayer(infos, layer_idx, "linear_attn.in_proj_qkv.weight", file_data), .in_proj_z = try findAndLoadLayer(infos, layer_idx, "linear_attn.in_proj_z.weight", file_data), .in_proj_a = try findAndLoadLayer(infos, layer_idx, "linear_attn.in_proj_a.weight", file_data), .in_proj_b = try findAndLoadLayer(infos, layer_idx, "linear_attn.in_proj_b.weight", file_data), .conv1d = try findAndLoadLayer(infos, layer_idx, "linear_attn.conv1d.weight", file_data), .A_log = try findAndLoadLayer(infos, layer_idx, "linear_attn.A_log", file_data), .dt_bias = try findAndLoadLayer(infos, layer_idx, "linear_attn.dt_bias", file_data), .norm = try findAndLoadLayer(infos, layer_idx, "linear_attn.norm.weight", file_data), .out_proj = try findAndLoadLayer(infos, layer_idx, "linear_attn.out_proj.weight", file_data), }; lw.full_attn = null; }
weights.layers[layer_idx] = lw; }
// Precompute RoPE cos/sin tables const rotary_dim = Config.rotary_dim; weights.rope_cos = try allocator.alloc(f32, max_seq_len * rotary_dim); weights.rope_sin = try allocator.alloc(f32, max_seq_len * rotary_dim);
// inv_freq = 1 / (theta ^ (2i / dim)) for i in 0..dim/2 // HuggingFace layout: emb = cat(freqs, freqs) → [f0,f1,...,f31,f0,f1,...,f31] const half_dim = rotary_dim / 2; for (0..max_seq_len) |pos| { for (0..half_dim) |i| { const freq = 1.0 / std.math.pow(f32, Config.rope_theta, @as(f32, @floatFromInt(2 * i)) / @as(f32, @floatFromInt(rotary_dim))); const angle = @as(f32, @floatFromInt(pos)) * freq; const cos_val = @cos(angle); const sin_val = @sin(angle); // First half and second half are identical (cat(freqs, freqs)) weights.rope_cos[pos * rotary_dim + i] = cos_val; weights.rope_cos[pos * rotary_dim + half_dim + i] = cos_val; weights.rope_sin[pos * rotary_dim + i] = sin_val; weights.rope_sin[pos * rotary_dim + half_dim + i] = sin_val; } }
return weights;}
fn findAndLoad(infos: []const q.TensorInfo, name: []const u8, file_data: []const u8) !Tensor { for (infos) |*info| { if (std.mem.eql(u8, info.name, name)) { return try q.tensorFromInfo(info, file_data); } } return q.SafetensorsError.TensorNotFound;}
fn findAndLoadLayer(infos: []const q.TensorInfo, layer_idx: usize, suffix: []const u8, file_data: []const u8) !Tensor { var name_buf: [256]u8 = undefined; const name = std.fmt.bufPrint(&name_buf, "model.language_model.layers.{d}.{s}", .{ layer_idx, suffix }) catch return error.PathTooLong; return findAndLoad(infos, name, file_data);}
// ============================================================// Weight loading from Q4MF format// ============================================================
/// Load model weights from Q4MF format (INT4 quantized)./// file_data must remain valid for the lifetime of the returned ModelWeights.pub fn loadWeightsQ4(allocator: std.mem.Allocator, file_data: []const u8, max_seq_len: u32) !ModelWeights { const header = try q.parseQ4MF(allocator, file_data);
var weights: ModelWeights = undefined; weights.allocator = allocator;
// Embedding weights.embed_tokens = try findAndLoadQ4(header, "model.language_model.embed_tokens.weight", allocator); weights.norm = try findAndLoadQ4(header, "model.language_model.norm.weight", allocator);
// Layers for (0..Config.num_layers) |layer_idx| { var lw: LayerWeights = undefined;
lw.input_layernorm = try findAndLoadLayerQ4(header, layer_idx, "input_layernorm.weight", allocator); lw.post_attention_layernorm = try findAndLoadLayerQ4(header, layer_idx, "post_attention_layernorm.weight", allocator);
lw.gate_proj = try findAndLoadLayerQ4(header, layer_idx, "mlp.gate_proj.weight", allocator); lw.up_proj = try findAndLoadLayerQ4(header, layer_idx, "mlp.up_proj.weight", allocator); lw.down_proj = try findAndLoadLayerQ4(header, layer_idx, "mlp.down_proj.weight", allocator);
if (Config.layer_is_full_attn[layer_idx]) { lw.full_attn = FullAttentionWeights{ .q_proj = try findAndLoadLayerQ4(header, layer_idx, "self_attn.q_proj.weight", allocator), .k_proj = try findAndLoadLayerQ4(header, layer_idx, "self_attn.k_proj.weight", allocator), .v_proj = try findAndLoadLayerQ4(header, layer_idx, "self_attn.v_proj.weight", allocator), .o_proj = try findAndLoadLayerQ4(header, layer_idx, "self_attn.o_proj.weight", allocator), .q_norm = try findAndLoadLayerQ4(header, layer_idx, "self_attn.q_norm.weight", allocator), .k_norm = try findAndLoadLayerQ4(header, layer_idx, "self_attn.k_norm.weight", allocator), }; lw.delta_net = null; } else { lw.delta_net = DeltaNetWeights{ .in_proj_qkv = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.in_proj_qkv.weight", allocator), .in_proj_z = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.in_proj_z.weight", allocator), .in_proj_a = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.in_proj_a.weight", allocator), .in_proj_b = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.in_proj_b.weight", allocator), .conv1d = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.conv1d.weight", allocator), .A_log = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.A_log", allocator), .dt_bias = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.dt_bias", allocator), .norm = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.norm.weight", allocator), .out_proj = try findAndLoadLayerQ4(header, layer_idx, "linear_attn.out_proj.weight", allocator), }; lw.full_attn = null; }
weights.layers[layer_idx] = lw; }
// Precompute RoPE cos/sin tables const rotary_dim = Config.rotary_dim; weights.rope_cos = try allocator.alloc(f32, max_seq_len * rotary_dim); weights.rope_sin = try allocator.alloc(f32, max_seq_len * rotary_dim);
const half_dim = rotary_dim / 2; for (0..max_seq_len) |pos| { for (0..half_dim) |i| { const freq = 1.0 / std.math.pow(f32, Config.rope_theta, @as(f32, @floatFromInt(2 * i)) / @as(f32, @floatFromInt(rotary_dim))); const angle = @as(f32, @floatFromInt(pos)) * freq; const cos_val = @cos(angle); const sin_val = @sin(angle); weights.rope_cos[pos * rotary_dim + i] = cos_val; weights.rope_cos[pos * rotary_dim + half_dim + i] = cos_val; weights.rope_sin[pos * rotary_dim + i] = sin_val; weights.rope_sin[pos * rotary_dim + half_dim + i] = sin_val; } }
return weights;}
fn findAndLoadQ4(header: q.Q4MFHeader, name: []const u8, allocator: std.mem.Allocator) !Tensor { for (header.tensors) |*info| { if (std.mem.eql(u8, info.name, name)) { return try q.tensorFromQ4MF(info, header.group_size, allocator); } } return q.SafetensorsError.TensorNotFound;}
fn findAndLoadLayerQ4(header: q.Q4MFHeader, layer_idx: usize, suffix: []const u8, allocator: std.mem.Allocator) !Tensor { var name_buf: [256]u8 = undefined; const name = std.fmt.bufPrint(&name_buf, "model.language_model.layers.{d}.{s}", .{ layer_idx, suffix }) catch return error.PathTooLong; return findAndLoadQ4(header, name, allocator);}
// ============================================================// Forward pass - single token// ============================================================
/// Scratch buffers for a single forward passpub const ForwardState = struct { hidden: [Config.hidden_size]f32, residual: [Config.hidden_size]f32, normed: [Config.hidden_size]f32,
// Attention temporaries q_buf: [Config.num_attention_heads * Config.head_dim * 2]f32, // includes gate k_buf: [Config.num_kv_heads * Config.head_dim]f32, v_buf: [Config.num_kv_heads * Config.head_dim]f32, attn_out: [Config.attn_output_dim]f32,
// DeltaNet temporaries dn_qkv: [Config.linear_conv_dim]f32, dn_z: [Config.linear_value_dim]f32, dn_a: [Config.linear_num_heads]f32, dn_b: [Config.linear_num_heads]f32, dn_out: [Config.linear_value_dim]f32,
// MLP temporaries mlp_gate: [Config.intermediate_size]f32, mlp_up: [Config.intermediate_size]f32, mlp_down: [Config.hidden_size]f32,
// Logits (reuses space since it's only needed at the end) logits: []f32,
allocator: std.mem.Allocator,
pub fn init(allocator: std.mem.Allocator) !ForwardState { var state: ForwardState = undefined; state.allocator = allocator; state.logits = try allocator.alloc(f32, Config.vocab_size); return state; }
pub fn deinit(self: *ForwardState) void { self.allocator.free(self.logits); }};
/// Run a single forward pass for one token./// Returns logits over the vocabulary.pub fn forward( weights: *const ModelWeights, model_cache: *cache.ModelCache, state: *ForwardState, token: u32, pos: u32,) []f32 { // 1. Embedding lookup const embed_offset = @as(usize, token) * Config.hidden_size; for (0..Config.hidden_size) |i| { state.hidden[i] = weights.embed_tokens.get(embed_offset + i); }
// 2. Process each layer for (0..Config.num_layers) |layer_idx| { const lw = &weights.layers[layer_idx];
// Save residual @memcpy(&state.residual, &state.hidden);
// Pre-attention RMSNorm q.rmsNorm(&state.normed, &state.hidden, &lw.input_layernorm, Config.rms_norm_eps);
// Attention (full or delta) if (Config.layer_is_full_attn[layer_idx]) { fullAttentionForward(weights, lw, &model_cache.kv_caches[layer_idx].?, state, pos); } else { deltaNetForward(lw, &model_cache.delta_states[layer_idx].?, state); }
// Residual connection for (0..Config.hidden_size) |i| { state.hidden[i] = state.residual[i] + state.hidden[i]; }
// Save residual for MLP @memcpy(&state.residual, &state.hidden);
// Pre-MLP RMSNorm q.rmsNorm(&state.normed, &state.hidden, &lw.post_attention_layernorm, Config.rms_norm_eps);
// MLP (SwiGLU) mlpForward(lw, state);
// Residual connection for (0..Config.hidden_size) |i| { state.hidden[i] = state.residual[i] + state.hidden[i]; } }
// 3. Final RMSNorm q.rmsNorm(&state.normed, &state.hidden, &weights.norm, Config.rms_norm_eps);
// 4. LM head (tied weights: logits = embed_tokens^T @ normed) // embed_tokens is [vocab_size, hidden_size], so each row is a token embedding for (0..Config.vocab_size) |i| { state.logits[i] = weights.embed_tokens.dot(i * Config.hidden_size, &state.normed); }
return state.logits;}
// ============================================================// Full Attention (GQA with output gate + QK-norm + partial RoPE)// ============================================================
fn fullAttentionForward( weights: *const ModelWeights, lw: *const LayerWeights, kv: *cache.KVCache, state: *ForwardState, pos: u32,) void { const attn = &lw.full_attn.?; const num_heads = Config.num_attention_heads; const num_kv_heads = Config.num_kv_heads; const head_dim = Config.head_dim; const kv_groups = num_heads / num_kv_heads; const rotary_dim = Config.rotary_dim;
// Project Q (includes gate), K, V q.matVec(&attn.q_proj, &state.normed, &state.q_buf); q.matVec(&attn.k_proj, &state.normed, &state.k_buf); q.matVec(&attn.v_proj, &state.normed, &state.v_buf);
// Split Q into query and gate: q_buf has [num_heads * head_dim * 2] // Organized as [num_heads][head_dim * 2], split each head's output into q and gate var gate_buf: [Config.attn_output_dim]f32 = undefined; for (0..num_heads) |h| { const src_offset = h * head_dim * 2; const dst_offset = h * head_dim; // First head_dim elements are query, second head_dim are gate for (0..head_dim) |d| { state.attn_out[dst_offset + d] = state.q_buf[src_offset + d]; // query gate_buf[dst_offset + d] = state.q_buf[src_offset + head_dim + d]; // gate } } // Now state.attn_out[0..num_heads*head_dim] = query, gate_buf = gate // Copy query back (we'll use attn_out for the final output later) var query: [Config.attn_output_dim]f32 = undefined; @memcpy(&query, &state.attn_out);
// QK-norm: normalize each head's Q and K with RMSNorm for (0..num_heads) |h| { const offset = h * head_dim; q.rmsNorm( query[offset .. offset + head_dim], query[offset .. offset + head_dim], &attn.q_norm, Config.rms_norm_eps, ); } for (0..num_kv_heads) |h| { const offset = h * head_dim; q.rmsNorm( state.k_buf[offset .. offset + head_dim], state.k_buf[offset .. offset + head_dim], &attn.k_norm, Config.rms_norm_eps, ); }
// Apply partial RoPE to query and key const cos = weights.rope_cos[pos * rotary_dim .. (pos + 1) * rotary_dim]; const sin = weights.rope_sin[pos * rotary_dim .. (pos + 1) * rotary_dim];
for (0..num_heads) |h| { applyRope(query[h * head_dim ..][0..head_dim], cos, sin, rotary_dim); } for (0..num_kv_heads) |h| { applyRope(state.k_buf[h * head_dim ..][0..head_dim], cos, sin, rotary_dim); }
// Store K, V in cache kv.append(&state.k_buf, &state.v_buf);
// Compute attention for each query head const scale = 1.0 / @sqrt(@as(f32, @floatFromInt(head_dim))); const seq_len = kv.len;
for (0..num_heads) |h| { const kv_h = h / kv_groups; // which KV head this query head uses const q_offset = h * head_dim;
// Dot product with all cached keys var max_score: f32 = -std.math.inf(f32);
// We need a temporary for attention scores // Use mlp_gate as scratch (it's big enough for seq_len <= intermediate_size) const scores = state.mlp_gate[0..seq_len];
for (0..seq_len) |t| { const cached_k = kv.getKey(@intCast(t)); const k_head = cached_k[kv_h * head_dim .. (kv_h + 1) * head_dim];
var score: f32 = 0.0; for (0..head_dim) |d| { score += query[q_offset + d] * k_head[d]; } score *= scale; scores[t] = score; if (score > max_score) max_score = score; }
// Softmax var sum_exp: f32 = 0.0; for (scores) |*s| { s.* = @exp(s.* - max_score); sum_exp += s.*; } if (sum_exp > 0.0) { const inv_sum = 1.0 / sum_exp; for (scores) |*s| { s.* *= inv_sum; } }
// Weighted sum of values const out_offset = h * head_dim; @memset(state.attn_out[out_offset .. out_offset + head_dim], 0.0); for (0..seq_len) |t| { const cached_v = kv.getValue(@intCast(t)); const v_head = cached_v[kv_h * head_dim .. (kv_h + 1) * head_dim]; const w = scores[t]; for (0..head_dim) |d| { state.attn_out[out_offset + d] += w * v_head[d]; } } }
// Apply output gate: attn_out *= sigmoid(gate) for (0..Config.attn_output_dim) |i| { state.attn_out[i] *= q.sigmoid(gate_buf[i]); }
// Flatten gate to [batch, seq, num_heads * head_dim] and reshape as needed // Already in the right shape. The gate_buf reshape in Python is: // gate = gate.reshape(*input_shape, -1) # flatten across heads // Our gate_buf is already [num_heads * head_dim] = [2048]
// Output projection: hidden = o_proj @ attn_out q.matVec(&attn.o_proj, &state.attn_out, &state.hidden);}
/// Apply rotary position embedding to the first rotary_dim elements of x./// x is [head_dim], only first rotary_dim elements are rotated.fn applyRope(x: []f32, cos: []const f32, sin: []const f32, rotary_dim: u32) void { // RoPE: for pairs (x[2i], x[2i+1]): // x'[2i] = x[2i] * cos[2i] - x[2i+1] * sin[2i] // x'[2i+1] = x[2i] * sin[2i+1] + x[2i+1] * cos[2i+1] // Wait — HuggingFace uses rotate_half style: split in half, not interleaved. // rotate_half: x1 = x[..dim//2], x2 = x[dim//2..], return cat(-x2, x1) // So: x_embed = x * cos + rotate_half(x) * sin const half = rotary_dim / 2; var i: usize = 0; while (i < half) : (i += 1) { const x0 = x[i]; const x1 = x[i + half]; // cos/sin are [rotary_dim], duplicated as [cos0,cos0,cos1,cos1,...] // But with rotate_half convention: // x_embed[i] = x[i] * cos[i] + (-x[i+half]) * sin[i] // x_embed[i+half] = x[i+half] * cos[i+half] + x[i] * sin[i+half] x[i] = x0 * cos[i] - x1 * sin[i]; x[i + half] = x1 * cos[i + half] + x0 * sin[i + half]; }}
// ============================================================// GatedDeltaNet (linear attention) - recurrent mode// ============================================================
fn deltaNetForward( lw: *const LayerWeights, delta_state: *cache.DeltaNetState, state: *ForwardState,) void { const dn = &lw.delta_net.?; const num_heads = Config.linear_num_heads; const key_dim = Config.linear_key_head_dim; const value_dim = Config.linear_value_head_dim; const conv_dim = Config.linear_conv_dim; const kernel_size = Config.linear_conv_kernel;
// 1. Project QKV q.matVec(&dn.in_proj_qkv, &state.normed, &state.dn_qkv);
// 2. Project Z (gate), A, B q.matVec(&dn.in_proj_z, &state.normed, &state.dn_z); q.matVec(&dn.in_proj_a, &state.normed, &state.dn_a); q.matVec(&dn.in_proj_b, &state.normed, &state.dn_b);
// 3. Causal conv1d update (single step) // conv_state is [conv_dim, kernel_size - 1], stored as [conv_dim][state_cols] // Full window = [conv_state..., new_input] = kernel_size values // Must compute conv BEFORE updating state const state_cols = kernel_size - 1; for (0..conv_dim) |ch| { const cs_base = ch * state_cols;
// Compute depthwise convolution: dot(window, weight) // window = [conv_state[0], conv_state[1], conv_state[2], new_input] var conv_out: f32 = 0.0; for (0..state_cols) |k| { conv_out += delta_state.conv_state[cs_base + k] * dn.conv1d.get(ch * kernel_size + k); } conv_out += state.dn_qkv[ch] * dn.conv1d.get(ch * kernel_size + state_cols);
// Update state: shift left and append new input var j: usize = 0; while (j + 1 < state_cols) : (j += 1) { delta_state.conv_state[cs_base + j] = delta_state.conv_state[cs_base + j + 1]; } delta_state.conv_state[cs_base + state_cols - 1] = state.dn_qkv[ch];
// Apply SiLU activation state.dn_qkv[ch] = q.silu(conv_out); }
// 4. Split QKV: [key_dim, key_dim, value_dim] const key_total = Config.linear_key_dim; // 2048 const q_start: usize = 0; const k_start = key_total; const v_start = key_total * 2;
// 5. Compute decay g and write gate beta per head // g = -exp(A_log) * softplus(a + dt_bias) // beta = sigmoid(b) var g_vals: [Config.linear_num_heads]f32 = undefined; var beta_vals: [Config.linear_num_heads]f32 = undefined; for (0..num_heads) |h| { const a_log = dn.A_log.get(h); const a_val = state.dn_a[h]; const dt_b = dn.dt_bias.get(h); g_vals[h] = -@exp(a_log) * q.softplus(a_val + dt_b); beta_vals[h] = q.sigmoid(state.dn_b[h]); }
// 6. Recurrent delta rule (single step per head) // For each head h: // q_t = l2norm(qkv[q_start + h*key_dim .. + key_dim]) // k_t = l2norm(qkv[k_start + h*key_dim .. + key_dim]) // v_t = qkv[v_start + h*value_dim .. + value_dim] // decay = exp(g[h]) // S = S * decay + k_t * (v_t - S^T @ k_t)^T * beta[h] // output[h] = S^T @ q_t const scale = 1.0 / @sqrt(@as(f32, @floatFromInt(key_dim)));
for (0..num_heads) |h| { const state_matrix = delta_state.getHeadState(@intCast(h));
// Extract and L2-normalize Q and K for this head var q_head: [Config.linear_key_head_dim]f32 = undefined; var k_head: [Config.linear_key_head_dim]f32 = undefined;
@memcpy(&q_head, state.dn_qkv[q_start + h * key_dim ..][0..key_dim]); @memcpy(&k_head, state.dn_qkv[k_start + h * key_dim ..][0..key_dim]);
q.l2Norm(&q_head); q.l2Norm(&k_head);
// Scale query for (&q_head) |*qv| { qv.* *= scale; }
const v_head = state.dn_qkv[v_start + h * value_dim ..][0..value_dim];
// Decay the state: S *= exp(g) const decay = @exp(g_vals[h]); for (state_matrix) |*s| { s.* *= decay; }
// Compute kv_mem = S @ k_t (matrix-vector product, S is [key_dim x value_dim]) // S is stored as [key_dim][value_dim], so S @ k means: for each value dim v, // sum over key_dim d: S[d][v] * k[d] // Actually, looking at the Python code: // kv_mem = (last_recurrent_state * k_t.unsqueeze(-1)).sum(dim=-2) // This means: kv_mem[v] = sum_d(S[d,v] * k[d]) var kv_mem: [Config.linear_value_head_dim]f32 = undefined; for (0..value_dim) |v| { var sum: f32 = 0.0; for (0..key_dim) |d| { sum += state_matrix[d * value_dim + v] * k_head[d]; } kv_mem[v] = sum; }
// delta = (v_t - kv_mem) * beta var delta: [Config.linear_value_head_dim]f32 = undefined; for (0..value_dim) |v| { delta[v] = (v_head[v] - kv_mem[v]) * beta_vals[h]; }
// Update state: S += k_t * delta^T // S[d,v] += k[d] * delta[v] for (0..key_dim) |d| { for (0..value_dim) |v| { state_matrix[d * value_dim + v] += k_head[d] * delta[v]; } }
// Output: out[v] = sum_d(S[d,v] * q[d]) for (0..value_dim) |v| { var sum: f32 = 0.0; for (0..key_dim) |d| { sum += state_matrix[d * value_dim + v] * q_head[d]; } state.dn_out[h * value_dim + v] = sum; } }
// 7. Gated RMS norm: out = rmsNorm(dn_out) * silu(z) // Applied per-head (norm weight is [value_head_dim]) for (0..num_heads) |h| { const offset = h * value_dim; q.rmsNormGated( state.dn_out[offset .. offset + value_dim], state.dn_out[offset .. offset + value_dim], state.dn_z[offset .. offset + value_dim], &dn.norm, Config.rms_norm_eps, ); }
// 8. Output projection: hidden = out_proj @ dn_out q.matVec(&dn.out_proj, &state.dn_out, &state.hidden);}
// ============================================================// SwiGLU MLP// ============================================================
fn mlpForward(lw: *const LayerWeights, state: *ForwardState) void { // gate = silu(gate_proj @ normed) // up = up_proj @ normed // hidden = down_proj @ (gate * up) q.matVec(&lw.gate_proj, &state.normed, &state.mlp_gate); q.matVec(&lw.up_proj, &state.normed, &state.mlp_up);
// gate = silu(gate) * up for (0..Config.intermediate_size) |i| { state.mlp_gate[i] = q.silu(state.mlp_gate[i]) * state.mlp_up[i]; }
// hidden = down_proj @ gate q.matVec(&lw.down_proj, state.mlp_gate[0..Config.intermediate_size], &state.hidden);}