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:
1 parent
a2ff07234c
commit
5669c18b3e
2 files changed
+272
No files matched your search
@@ -2,6 +2,7 @@ pub const tensor = @import("tensor.zig");
|
|||||||
pub const pickle = @import("pickle.zig");
|
pub const pickle = @import("pickle.zig");
|
||||||
pub const pt = @import("pt.zig");
|
pub const pt = @import("pt.zig");
|
||||||
pub const layers = @import("layers.zig");
|
pub const layers = @import("layers.zig");
|
||||||
|
pub const transformer = @import("transformer.zig");
|
||||||
|
|
||||||
pub const Tensor = tensor.Tensor;
|
pub const Tensor = tensor.Tensor;
|
||||||
pub const DType = tensor.DType;
|
pub const DType = tensor.DType;
|
||||||
|
|||||||
@@ -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;
|
||||||
|
}
|
||||||
Reference in new issue
Block a user