Skip to content

Coming from PyTorch ​

Three things change: shapes are checked before anything runs, source never carries weights, and one file runs in PyTorch, JAX, XLA and ONNX Runtime. Operations, parameter names and weights on disk stay the same.

The same block, twice ​

python
class Block(nn.Module):
    def __init__(self, h, heads):
        super().__init__()
        self.norm = RMSNorm(h)
        self.qkv = nn.Linear(h, 3 * h, bias=False)
        self.out = nn.Linear(h, h, bias=False)
        self.heads = heads

    def forward(self, x):
        b, s, h = x.shape
        q, k, v = self.qkv(self.norm(x)).split(h, dim=-1)
        q = q.view(b, s, self.heads, -1).transpose(1, 2)
        k = k.view(b, s, self.heads, -1).transpose(1, 2)
        v = v.view(b, s, self.heads, -1).transpose(1, 2)
        mixed = F.scaled_dot_product_attention(q, k, v, is_causal=True)
        return x + self.out(mixed.transpose(1, 2).reshape(b, s, h))
linnet
pub block Block<H: Dim, Heads: Dim, T: Float = bf16>
where
    Heads > 0,
    H % Heads == 0
{
    sub norm: RmsNorm<H, T>
    sub qkv: Linear<H, 3 * H, T>
    sub out: Linear<H, H, T>

    pub entry forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let projected = qkv.forward(norm.forward(x))
        let q = heads<B, S, Heads, H / Heads, T>(projected[:, :, 0:H])
        let k = heads<B, S, Heads, H / Heads, T>(projected[:, :, H:2 * H])
        let v = heads<B, S, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
        let mixed = attention(q, k, v, rsqrt(cast<f32>(H / Heads)), some(causal_mask<S, S>()))
        return x + out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [B, S, H]))
    }
}
PyTorchLinnet
b, s, h = x.shape at runtimeB, S, H are generics; every shape derives from them
.view(..., -1): a wrong divisor reshapes silentlyreshape to a shape proved from H % Heads == 0, or rejected
__init__ builds submodulessub norm: RmsNorm<H, T> declares one; nothing allocates
state_dict() keysthe same names: norm.weight, qkv.weight, out.weight
F.scaled_dot_product_attentionattention from stdlib/, written in Linnet

Running it ​

python
from linnet.torch import load

model = load("src/block.linnet",
             generics={"H": 1024, "Heads": 16, "T": "bf16"},
             weights="checkpoints/block/",     # SafeTensors named by parameter path
             device="cuda",
             numerics="equivalent",            # PyTorch kernels for library operations
             compile="inductor")               # generated source under torch.compile

numerics="fast" runs softmax, normalization and attention in the input dtype, as PyTorch reference models do on bf16. trainable=True lets any torch.optim optimizer train the model. See PyTorch and JAX and Flax.

Importing a model ​

python
from linnet.torch import export_linnet

export_linnet(module, (example_input,), output="src/model.linnet", weights="weights/")

export_linnet traces with torch.export and writes checked Linnet that keeps the parameter names, or fails naming the operation it cannot express. For ONNX and JAX, use linnet.onnx.import_onnx and linnet.jax.export_linnet.

What you get ​

nn.ModuleLinnet
Shape errorsat runtime, on the first batch that reaches themat check time, both shapes named
Loading a modelruns the model's Python; pickle in .pt fileslinnet check runs nothing; weights are SafeTensors by path
Performancenative kernelsthe same kernels from generated source, replayable as CUDA graphs, or XLA (Llama 3.1 8B decodes one request at 167 tok/s on an H100, transformers compiled at 110); see Benchmarks

The trade: no Python inside the model and no data-dependent shapes. Loops are while over scalars or compile-time static for.

Where things live ​

PyTorchLinnet
nn.Module subclassblock
nn.Parameterparam name: Tensor[...; T]
register_bufferbuffer, or state for caches the block updates
forwardpub entry forward<...>(...); any number of entries
nn.ModuleListsub layers: [Layer<...>; N] and static for layer in layers
x @ w.T, einsummatmul, or index notation sum[k] x[i, k] * w[j, k]
F.softmax, F.layer_norm, ...std.nn.softmax::softmax, std.nn.norm::layer_norm, ...
F.cross_entropystd.nn.loss::cross_entropy, a weighted sum; linear_cross_entropy with the output head
torch.export, ONNX exportlinnet stablehlo, linnet onnx, linnet torch, linnet plan

Next: the Quickstart, then the Language tour.

Released under the MIT License.