diff --git a/src/crucible.zig b/src/crucible.zig index 6b326e5..a4d8736 100644 --- a/src/crucible.zig +++ b/src/crucible.zig @@ -2,6 +2,7 @@ pub const tensor = @import("tensor.zig"); pub const pickle = @import("pickle.zig"); pub const pt = @import("pt.zig"); pub const layers = @import("layers.zig"); +pub const transformer = @import("transformer.zig"); pub const Tensor = tensor.Tensor; pub const DType = tensor.DType; diff --git a/src/transformer.zig b/src/transformer.zig new file mode 100644 index 0000000..b2a35dc --- /dev/null +++ b/src/transformer.zig @@ -0,0 +1,271 @@ +//! Transformer layer primitives — the missing set for smol transformer + +//! graph-augmented inference. Built on the same f32x8 FMA philosophy as the +//! conv/linear primitives. + +const std = @import("std"); +const V = @Vector(8, f32); +const softmax = @import("layers.zig").softmax; + +// --------------------------------------------------------------------------- +// Normalization +// --------------------------------------------------------------------------- + +/// LayerNorm: (x - mean) / sqrt(var + eps) * weight + bias over the last dim. +/// weight/bias may be null (no affine). +pub fn layerNorm(input: []const f32, weight: ?[]const f32, bias: ?[]const f32, eps: f32, output: []f32) void { + const n: f32 = @floatFromInt(input.len); + var mean: f32 = 0; + for (input) |v| mean += v; + mean /= n; + + var variance: f32 = 0; + for (input) |v| { + const d = v - mean; + variance += d * d; + } + variance /= n; + + const inv_std = 1.0 / @sqrt(variance + eps); + for (input, 0..) |v, i| { + const normalized = (v - mean) * inv_std; + output[i] = if (weight) |w| normalized * w[i] + bias.?[i] else normalized; + } +} + +/// RMSNorm: x / sqrt(mean(x^2) + eps) * weight. Used in LLaMA/Mistral-family. +pub fn rmsNorm(input: []const f32, weight: []const f32, eps: f32, output: []f32) void { + var sum_squares: f32 = 0; + for (input) |v| sum_squares += v * v; + const inv_rms = 1.0 / @sqrt(sum_squares / @as(f32, @floatFromInt(input.len)) + eps); + for (input, 0..) |v, i| { + output[i] = v * inv_rms * weight[i]; + } +} + +// --------------------------------------------------------------------------- +// Activations (transformer-relevant) +// --------------------------------------------------------------------------- + +/// GELU (tanh approximation, matching PyTorch's default). +pub fn gelu(input: []const f32, output: []f32) void { + for (input, 0..) |x, i| { + // tanh approximation: 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3))) + const inner = 0.7978845608028654 * (x + 0.044715 * x * x * x); + output[i] = 0.5 * x * (1.0 + std.math.tanh(inner)); + } +} + +/// SiLU (Swish): x * sigmoid(x). Used in LLaMA/Mistral-family. +pub fn silu(input: []const f32, output: []f32) void { + for (input, 0..) |x, i| { + output[i] = x / (1.0 + @exp(-x)); + } +} + +// --------------------------------------------------------------------------- +// Embedding +// --------------------------------------------------------------------------- + +/// Embedding lookup: weight is (vocab_size, embed_dim), returns the row. +pub fn embedding(weight: []const f32, embed_dim: usize, token_ids: []const usize, output: []f32) void { + for (token_ids, 0..) |tid, i| { + const src = weight[tid * embed_dim ..][0..embed_dim]; + @memcpy(output[i * embed_dim ..][0..embed_dim], src); + } +} + +// --------------------------------------------------------------------------- +// Matmul (GEMM) +// --------------------------------------------------------------------------- + +/// 2D matrix multiply: A (M, K) x B (K, N) -> C (M, N). f32x8 FMA over K. +pub fn matmul(a: []const f32, b: []const f32, c: []f32, m: usize, k: usize, n: usize) void { + @memset(c, 0); + for (0..m) |i| { + const a_row = a[i * k ..][0..k]; + const c_row = c[i * n ..][0..n]; + for (0..k) |kk| { + const av: V = @splat(a_row[kk]); + const b_row = b[kk * n ..][0..n]; + var x: usize = 0; + while (x + 8 <= n) : (x += 8) { + const bv: V = b_row[x..][0..8].*; + const cv: V = c_row[x..][0..8].*; + c_row[x..][0..8].* = @mulAdd(V, av, bv, cv); + } + while (x < n) : (x += 1) { + c_row[x] += a_row[kk] * b_row[x]; + } + } + } +} + +// --------------------------------------------------------------------------- +// Attention +// --------------------------------------------------------------------------- + +/// Scaled dot-product attention for a single head: +/// scores = softmax(Q @ K^T / sqrt(d_k)) @ V +/// q: (seq_q, d_k), k: (seq_k, d_k), v: (seq_k, d_v) +/// output: (seq_q, d_v) +/// Optional mask: (seq_q, seq_k), added to scores before softmax. +pub fn attention( + alloc: std.mem.Allocator, + q: []const f32, + k: []const f32, + v: []const f32, + seq_q: usize, + seq_k: usize, + d_k: usize, + d_v: usize, + mask: ?[]const f32, + output: []f32, +) !void { + // scores = Q @ K^T / sqrt(d_k) → (seq_q, seq_k) + const scores = try alloc.alloc(f32, seq_q * seq_k); + defer alloc.free(scores); + const scale = 1.0 / @sqrt(@as(f32, @floatFromInt(d_k))); + + // Q @ K^T + matmul(q, k, scores, seq_q, d_k, seq_k); + for (scores) |*s| s.* *= scale; + + // apply mask (additive) + if (mask) |m| { + for (scores, 0..) |*s, i| s.* += m[i]; + } + + // row-wise softmax + for (0..seq_q) |i| { + softmax(scores[i * seq_k ..][0..seq_k]); + } + + // output = softmax(Q @ K^T / sqrt(d_k)) @ V → (seq_q, d_v) + matmul(scores, v, output, seq_q, seq_k, d_v); +} + +/// Multi-head attention: splits q/k/v into `n_heads` heads along the embed +/// dimension, runs attention per head, concatenates outputs. +/// q: (seq_q, embed), k: (seq_k, embed), v: (seq_k, embed) +/// output: (seq_q, embed) +pub fn multiHeadAttention( + alloc: std.mem.Allocator, + q: []const f32, + k: []const f32, + v: []const f32, + seq_q: usize, + seq_k: usize, + embed: usize, + n_heads: usize, + mask: ?[]const f32, // optional additive mask (e.g. padding) applied after causal + output: []f32, +) !void { + const head_dim = embed / n_heads; + var concat = try alloc.alloc(f32, seq_q * embed); + defer alloc.free(concat); + + for (0..n_heads) |h| { + const q_off = h * head_dim; + const k_off = h * head_dim; + const v_off = h * head_dim; + + // extract per-head slices (stride = embed) + const q_head = try alloc.alloc(f32, seq_q * head_dim); + defer alloc.free(q_head); + const k_head = try alloc.alloc(f32, seq_k * head_dim); + defer alloc.free(k_head); + const v_head = try alloc.alloc(f32, seq_k * head_dim); + defer alloc.free(v_head); + + for (0..seq_q) |i| { + for (0..head_dim) |d| { + q_head[i * head_dim + d] = q[i * embed + q_off + d]; + } + } + for (0..seq_k) |i| { + for (0..head_dim) |d| { + k_head[i * head_dim + d] = k[i * embed + k_off + d]; + } + for (0..head_dim) |d| { + v_head[i * head_dim + d] = v[i * embed + v_off + d]; + } + } + + // attention per head + const head_out = try alloc.alloc(f32, seq_q * head_dim); + defer alloc.free(head_out); + + // scores + const scores = try alloc.alloc(f32, seq_q * seq_k); + defer alloc.free(scores); + const scale = 1.0 / @sqrt(@as(f32, @floatFromInt(head_dim))); + matmul(q_head, k_head, scores, seq_q, head_dim, seq_k); + for (scores) |*s| s.* *= scale; + + // causal mask (lower triangular) + for (0..seq_q) |i| { + for (0..seq_k) |j| { + if (j > i) scores[i * seq_k + j] = -std.math.floatMax(f32) / 2.0; + } + } + // apply external mask if provided + if (mask) |m| { + for (scores, 0..) |*s, i| s.* += m[i]; + } + + // row-wise softmax + for (0..seq_q) |i| { + softmax(scores[i * seq_k ..][0..seq_k]); + } + + // output = scores @ v_head + matmul(scores, v_head, head_out, seq_q, seq_k, head_dim); + + // scatter into concat + for (0..seq_q) |i| { + for (0..head_dim) |d| { + concat[i * embed + q_off + d] = head_out[i * head_dim + d]; + } + } + } + + @memcpy(output[0..seq_q * embed], concat); +} + +// --------------------------------------------------------------------------- +// RoPE (Rotary Position Embeddings) +// --------------------------------------------------------------------------- + +/// Apply RoPE in place on a (seq_len, embed_dim) tensor. +/// Requires embed_dim to be even. cos/sin are precomputed or computed inline. +pub fn rope(input: []f32, seq_len: usize, embed_dim: usize, base: f32) void { + const half = embed_dim / 2; + for (0..seq_len) |pos| { + for (0..half) |d| { + const freq = @exp(-@log(base) * @as(f32, @floatFromInt(2 * d)) / @as(f32, @floatFromInt(embed_dim))); + const angle = @as(f32, @floatFromInt(pos)) * freq; + const cos_a = @cos(angle); + const sin_a = @sin(angle); + const idx = pos * embed_dim; + const x0 = input[idx + d]; + const x1 = input[idx + d + half]; + input[idx + d] = x0 * cos_a - x1 * sin_a; + input[idx + d + half] = x0 * sin_a + x1 * cos_a; + } + } +} + +// --------------------------------------------------------------------------- +// Causal mask generation +// --------------------------------------------------------------------------- + +/// Generate a causal mask (lower triangular = 0, upper triangular = -inf). +pub fn causalMask(alloc: std.mem.Allocator, seq_q: usize, seq_k: usize) ![]f32 { + const mask = try alloc.alloc(f32, seq_q * seq_k); + for (0..seq_q) |i| { + for (0..seq_k) |j| { + mask[i * seq_k + j] = if (j > i) -std.math.floatMax(f32) / 2.0 else 0; + } + } + return mask; +}