Vision Transformer
A Vision Transformer that maps images to class logits, with unmasked attention (none). It shows reshapes the checker proves, shapes sized by expressions, and pooling chosen at compile time.
Commands
linnet check --std stdlib examples/02-vit/vit.linnet
linnet stablehlo --std stdlib --bind Height=32 --bind Width=32 --bind Channels=3 --bind Patch=8 \
--bind D=64 --bind Heads=4 --bind Inner=128 --bind Layers=2 --bind Classes=10 \
--bind T=f32 --bind B=1 examples/02-vit/vit.linnetKey ideas
Patchify by reshape
patchify<B, C, Height, Width, P, T> reshapes [B, C, Height, Width] to [B, (Height / P) * (Width / P), C * P * P] through a permute. With Height % Patch == 0 and Width % Patch == 0 declared, the solver proves the element counts match. Without them, it refuses.
Sizes as expressions
param positions: Tensor[(Height / Patch) * (Width / Patch) + 1, D; T]Model.forward expands the class token to [B, 1, D] with broadcast_to and concatenates it in front of the patches. The position table has one row per patch plus one for the class token, and the checker matches it against that concatenation.
Compile-time choice
pub enum Pooling { ClassToken, Mean }
pub const POOLING: Pooling = Pooling.ClassTokenmatch POOLING { ... } picks the pooled representation. POOLING is a constant, so the emitter and the exporters keep only the chosen arm.
Source
vit.linnet
// A Vision Transformer. Patches are cut with `reshape` and `permute`, whose
// element counts the checker proves from the `where` clause; the class token
// and position table are parameters sized by compile-time arithmetic; the
// pooling strategy is a compile-time constant matched at build time.
module examples.vit
use std.nn.activations::{gelu}
use std.nn.attention::{attention}
use std.nn.linear::{Linear}
use std.nn.norm::{layer_norm}
pub enum Pooling {
ClassToken,
Mean,
}
pub const POOLING: Pooling = Pooling.ClassToken
// `[B, C, Height, Width]` -> `[B, Patches, C * P * P]`, row-major over patches.
pub fn patchify<B: Dim, C: Dim, Height: Dim, Width: Dim, P: Dim, T: Float>(
image: Tensor[B, C, Height, Width; T],
) -> Tensor[B, (Height / P) * (Width / P), C * P * P; T]
where
P > 0,
Height % P == 0,
Width % P == 0
{
let grid = reshape(image, [B, C, Height / P, P, Width / P, P])
let ordered = permute(grid, [0, 2, 4, 1, 3, 5])
return reshape(ordered, [B, (Height / P) * (Width / P), C * P * P])
}
pub block LayerNorm<D: Dim, T: Float> {
param weight: Tensor[D; T]
param bias: Tensor[D; T]
pub fn forward<*S: Shape>(x: Tensor[*S, D; T]) -> Tensor[*S, D; T] {
return layer_norm(x, weight, some(bias))
}
}
pub block Attention<D: Dim, Heads: Dim, T: Float>
where
Heads > 0,
D % Heads == 0
{
sub q: Linear<D, D, T>
sub k: Linear<D, D, T>
sub v: Linear<D, D, T>
sub out: Linear<D, D, T>
// Every token attends to every other: no mask.
pub fn forward<B: Dim, N: Dim>(x: Tensor[B, N, D; T]) -> Tensor[B, N, D; T] {
let mixed = attention(
split<B, N, Heads, D / Heads, T>(q.forward(x)),
split<B, N, Heads, D / Heads, T>(k.forward(x)),
split<B, N, Heads, D / Heads, T>(v.forward(x)),
rsqrt(cast<f32>(D / Heads)),
none,
)
return out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [B, N, D]))
}
}
fn split<B: Dim, N: Dim, Heads: Dim, Dh: Dim, T: Float>(
x: Tensor[B, N, Heads * Dh; T],
) -> Tensor[B, Heads, N, Dh; T] {
return permute(reshape(x, [B, N, Heads, Dh]), [0, 2, 1, 3])
}
pub block EncoderLayer<D: Dim, Heads: Dim, Inner: Dim, T: Float>
where
Heads > 0,
D % Heads == 0
{
sub norm_1: LayerNorm<D, T>
sub attention: Attention<D, Heads, T>
sub norm_2: LayerNorm<D, T>
sub up: Linear<D, Inner, T>
sub down: Linear<Inner, D, T>
pub fn forward<B: Dim, N: Dim>(x: Tensor[B, N, D; T]) -> Tensor[B, N, D; T] {
let attended = x + attention.forward(norm_1.forward(x))
return attended + down.forward(gelu(up.forward(norm_2.forward(attended))))
}
}
pub block Model<
Height: Dim,
Width: Dim,
Channels: Dim,
Patch: Dim,
D: Dim,
Heads: Dim,
Inner: Dim,
Layers: Dim,
Classes: Dim,
T: Float = f32,
>
where
Patch > 0,
Heads > 0,
Height % Patch == 0,
Width % Patch == 0,
D % Heads == 0
{
sub patch_embed: Linear<Channels * Patch * Patch, D, T>
param class_token: Tensor[1, 1, D; T]
param positions: Tensor[(Height / Patch) * (Width / Patch) + 1, D; T]
sub layers: [EncoderLayer<D, Heads, Inner, T>; Layers]
sub norm: LayerNorm<D, T>
sub head: Linear<D, Classes, T>
pub entry forward<B: Dim>(
image: Tensor[B, Channels, Height, Width; T],
) -> Tensor[B, Classes; T] {
let patches = patch_embed.forward(
patchify<B, Channels, Height, Width, Patch, T>(image),
)
let tokens = concat(broadcast_to(class_token, [B, 1, D]), patches, axis = 1) + positions
var x = tokens
static for layer in layers {
x = layer.forward(x)
}
let normalized = norm.forward(x)
let pooled = match POOLING {
ClassToken => normalized[:, 0, :]
Mean => mean_tokens(normalized)
}
return head.forward(pooled)
}
}
fn mean_tokens<B: Dim, N: Dim, D: Dim, T: Float>(x: Tensor[B, N, D; T]) -> Tensor[B, D; T] {
let pooled[b, d] = sum[n] x[b, n, d] / cast<T>(N)
return pooled
}