Skip to content

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 ​

bash
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.linnet

Key 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 ​

linnet
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 ​

linnet
pub enum Pooling { ClassToken, Mean }
pub const POOLING: Pooling = Pooling.ClassToken

match POOLING { ... } picks the pooled representation. POOLING is a constant, so the emitter and the exporters keep only the chosen arm.

Source ​

vit.linnet ​

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
}

Browse on GitHub

Released under the MIT License.