Skip to content

CLIP ​

A CLIP-style dual encoder: a package of four modules whose two towers share one encoder, with three entries over one parameter set.

Commands ​

bash
linnet lint --std stdlib examples/03-clip
linnet inspect --parameters --std stdlib examples/03-clip/src/lib.linnet
linnet stablehlo --std stdlib --entry similarity --bind Height=32 --bind Width=32 --bind Channels=3 \
                 --bind Patch=8 --bind Vocab=100 --bind Context=16 --bind D=64 --bind Heads=4 \
                 --bind Inner=128 --bind Layers=2 --bind Embed=32 --bind T=f32 \
                 --bind B=1 --bind C=3 --bind S=8 examples/03-clip/src/lib.linnet

Key ideas ​

One encoder, two towers ​

src/encoder.linnet defines Encoder<D, Heads, Inner, Layers, T>. src/vision.linnet runs it over patch tokens without a mask (none), and src/text.linnet over token embeddings with some(causal_mask(...)). src/lib.linnet imports both towers with crate. paths.

Three entries, one parameter set ​

The root block Clip has three entries. embed_image and embed_text return L2-normalized embeddings. similarity<B, C, S> contracts them into [images, texts] cosine similarities, scaled by a learned temperature stored as Tensor[1; f32]. A backend exposes each entry as a separate function over the same weights.

Source ​

linnet.toml ​

toml
[package]
name = "clip"
version = "0.1.0"
language = "0.1"

src/encoder.linnet ​

linnet
// The transformer encoder both towers share. Whether attention is causal
// is a compile-time choice of the caller, passed as an optional mask.
module clip.encoder

use std.nn.activations::{gelu}
use std.nn.attention::{attention}
use std.nn.linear::{Linear}
use std.nn.norm::{layer_norm}

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 EncoderLayer<D: Dim, Heads: Dim, Inner: Dim, T: Float>
where
    Heads > 0,
    D % Heads == 0
{
    sub norm_1: LayerNorm<D, T>
    sub qkv: Linear<D, 3 * D, T>
    sub out: Linear<D, D, 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],
        mask: Tensor[N, N; bool]?,
    ) -> Tensor[B, N, D; T] {
        let projected = qkv.forward(norm_1.forward(x))
        let mixed = attention(
            heads<B, N, Heads, D / Heads, T>(projected[:, :, 0:D]),
            heads<B, N, Heads, D / Heads, T>(projected[:, :, D:2 * D]),
            heads<B, N, Heads, D / Heads, T>(projected[:, :, 2 * D:3 * D]),
            rsqrt(cast<f32>(D / Heads)),
            mask,
        )
        let attended = x + out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [B, N, D]))
        return attended + down.forward(gelu(up.forward(norm_2.forward(attended))))
    }
}

fn heads<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 Encoder<D: Dim, Heads: Dim, Inner: Dim, Layers: Dim, T: Float>
where
    Heads > 0,
    D % Heads == 0
{
    sub layers: [EncoderLayer<D, Heads, Inner, T>; Layers]

    pub fn forward<B: Dim, N: Dim>(
        x: Tensor[B, N, D; T],
        mask: Tensor[N, N; bool]?,
    ) -> Tensor[B, N, D; T] {
        var h = x
        static for layer in layers {
            h = layer.forward(h, mask)
        }
        return h
    }
}

// Rows scaled to unit length, in f32.
pub fn normalize<B: Dim, D: Dim, T: Float>(x: Tensor[B, D; T]) -> Tensor[B, D; f32] {
    let xf = cast<f32>(x)
    let squared[b] = sum[d] xf[b, d] * xf[b, d]
    let unit[b, d] = xf[b, d] * rsqrt(squared[b] + 1e-12)
    return unit
}

src/lib.linnet ​

linnet
// CLIP: two towers trained to agree. Each tower is its own module and the
// root block exposes three entries — one per embedding and the similarity
// matrix — over the same parameters.
module clip

use crate.encoder::{normalize}
use crate.text::{TextTower}
use crate.vision::{VisionTower}

pub block Clip<
    Height: Dim,
    Width: Dim,
    Channels: Dim,
    Patch: Dim,
    Vocab: Dim,
    Context: Dim,
    D: Dim,
    Heads: Dim,
    Inner: Dim,
    Layers: Dim,
    Embed: Dim,
    T: Float = f32,
>
where
    Patch > 0,
    Heads > 0,
    Height % Patch == 0,
    Width % Patch == 0,
    D % Heads == 0
{
    sub vision: VisionTower<Height, Width, Channels, Patch, D, Heads, Inner, Layers, Embed, T>
    sub text: TextTower<Vocab, Context, D, Heads, Inner, Layers, Embed, T>
    // The learned temperature, stored as its logarithm.
    param logit_scale: Tensor[1; f32]

    pub entry embed_image<B: Dim>(
        image: Tensor[B, Channels, Height, Width; T],
    ) -> Tensor[B, Embed; f32] {
        return normalize(vision.forward(image))
    }

    pub entry embed_text<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, Embed; f32]
    where
        S > 0,
        S <= Context
    {
        return normalize(text.forward(tokens))
    }

    // `[images, texts]` cosine similarities scaled by the temperature.
    pub entry similarity<B: Dim, C: Dim, S: Dim>(
        image: Tensor[B, Channels, Height, Width; T],
        tokens: Tensor[C, S; i32],
    ) -> Tensor[B, C; f32]
    where
        S > 0,
        S <= Context
    {
        let images = normalize(vision.forward(image))
        let texts = normalize(text.forward(tokens))
        let logits[b, c] = sum[d] images[b, d] * texts[c, d]
        return logits * exp(logit_scale)
    }
}

src/text.linnet ​

linnet
// The text tower: token and position embeddings, a causal encoder, and the
// last position projected into the joint embedding space.
module clip.text

use crate.encoder::{Encoder, LayerNorm}
use std.nn.attention::{causal_mask}
use std.nn.embedding::{embedding}
use std.nn.linear::{Linear}

pub block TextTower<
    Vocab: Dim,
    Context: Dim,
    D: Dim,
    Heads: Dim,
    Inner: Dim,
    Layers: Dim,
    Embed: Dim,
    T: Float,
>
where
    Heads > 0,
    D % Heads == 0
{
    param token_embedding: Tensor[Vocab, D; T]
    param positions: Tensor[Context, D; T]
    sub encoder: Encoder<D, Heads, Inner, Layers, T>
    sub final_norm: LayerNorm<D, T>
    sub projection: Linear<D, Embed, T>

    pub fn forward<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, Embed; T]
    where
        S > 0,
        S <= Context
    {
        let x = embedding(tokens, token_embedding) + embedding(iota<i32>(S), positions)
        let encoded = encoder.forward(x, some(causal_mask<S, S>()))
        return projection.forward(final_norm.forward(encoded[:, S - 1, :]))
    }
}

src/vision.linnet ​

linnet
// The image tower: patches, a class token, positions, the shared encoder,
// and a projection of the class token into the joint embedding space.
module clip.vision

use crate.encoder::{Encoder, LayerNorm}
use std.nn.linear::{Linear}

pub block VisionTower<
    Height: Dim,
    Width: Dim,
    Channels: Dim,
    Patch: Dim,
    D: Dim,
    Heads: Dim,
    Inner: Dim,
    Layers: Dim,
    Embed: Dim,
    T: Float,
>
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 pre_norm: LayerNorm<D, T>
    sub encoder: Encoder<D, Heads, Inner, Layers, T>
    sub post_norm: LayerNorm<D, T>
    sub projection: Linear<D, Embed, T>

    pub fn forward<B: Dim>(image: Tensor[B, Channels, Height, Width; T]) -> Tensor[B, Embed; T] {
        let grid = reshape(image, [B, Channels, Height / Patch, Patch, Width / Patch, Patch])
        let patches = reshape(
            permute(grid, [0, 2, 4, 1, 3, 5]),
            [B, (Height / Patch) * (Width / Patch), Channels * Patch * Patch],
        )
        let tokens = concat(
            broadcast_to(class_token, [B, 1, D]),
            patch_embed.forward(patches),
            axis = 1,
        )
        let encoded = encoder.forward(pre_norm.forward(tokens + positions), none)
        return projection.forward(post_norm.forward(encoded[:, 0, :]))
    }
}

Browse on GitHub

Released under the MIT License.