transformer: LayerNorm, RMSNorm, GELU, SiLU, Embedding, GEMM, attention (self + cross), multi-head, RoPE, causal mask

The transformer layer set for smol transformer + graph-augmented inference.
Attention supports self-attention (causal) and cross-attention (no causal
mask, K/V from graph embeddings). RoPE included for positional encoding.

👾 Generated with [Letta Code](https://letta.com)

Co-Authored-By: Letta Code <noreply@letta.com>
This commit is contained in:
pierreandLetta Code committed 2026-09-30 10:16:21 +03:00
1 parent a2ff07234c
commit 5669c18b3e
2 files changed
+272

No files matched your search

+1
View File
@@ -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;
+271
View File
@@ -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;
}