models/whisper-tinyREADME.md6.4 kB
md
# Whisper tiny
OpenAI's smallest multilingual speech recognition model, and the first
encoder-decoder model in Nest. The encoder reads a 30-second window of audio
as an 80-channel log-mel spectrogram of 3000 frames, halves the frame rate
with two convolutions, and runs four pre-norm transformer layers of width 384
over the 1500 positions that remain. The decoder is an ordinary causal
language model over the 51865-token multilingual vocabulary with a
cross-attention over those encoder states inserted between its self-attention
and its MLP, and an output head tied to its token embedding.
## The convolutional stem
The stem is two 1-D convolutions of kernel three over the mel channels, the
first at stride 1 and the second at stride 2, each followed by GELU. The
standard library has no `conv1d` -- `std.nn.conv::conv2d` takes a square
kernel -- so the source carries its own, built the way `conv2d` is built: an
index expression may not do arithmetic, so the tap positions are a tensor
made from `iota` and the padded signal is gathered through it.
```linnet
let taps[o, k] =
iota<i32>((L + 2 * Pad - K) / Stride + 1)[o] * cast<i32>(Stride) + iota<i32>(K)[k]
let windows[b, o, ci, k] = padded[b, taps[o, k], ci]
```
Signals are `[B, Positions, Channels]` throughout the stem, which is the
layout the encoder layers want anyway, so the gathered window flattens over
(channel, tap) and the convolution finishes as one `linear` against the
checkpoint's `[Cout, Cin, K]` kernel flattened the same way. That is a
deliberate choice rather than a shortcut: there is no `conv1d` op for a
backend to recognize, so a contraction written out in index notation would
stay a five-axis product, while `linear` selects `F.linear`, a `dot_general`,
or an ONNX `Gemm` on every backend.
The output length is the form the shape solver derives rather than one
written by hand. For kernel three at stride two with one frame of padding,
`(Frames + 2 * 1 - 3) / 2 + 1` reduces to `(Frames - 1) / 2 + 1`, and that
expression is the length of `encoder_positions`: the position table's size is
checked arithmetic, not a constant that happens to be 1500.
Whisper's positions are sinusoidal in origin, but the checkpoint stores them
as an ordinary tensor, so the source simply reads them.
## The key projection has no bias
In both stacks, queries, values and the output projection carry a bias and
keys do not. That asymmetry is in the reference implementation and it is
real: a key bias shifts every key by the same vector, which a softmax over
dot products does not ignore. `k_proj.bias` is therefore left out of
`bindings.json`, `Linear`'s bias is an optional parameter, and the absent
optional makes the `none` branch of `std.nn.linear::linear` run.
## Entries
Transcription encodes once and decodes many times, so the two stacks are
separate entries rather than one forward pass. The audio length the decoder
cross-attends over is its own generic, so encoder states shorter than a full
30-second window are accepted as they are.
```python
from linnet import nest
model = nest.load("whisper-tiny", backend="torch")
states = model.run_entry("encode", [mel]) # [B, 80, 3000] -> [B, 1500, 384]
logits = model.run_entry("decode", [tokens, states]) # -> [B, S, 51865]
```
`mel` is what `WhisperFeatureExtractor` produces (`input_features`): a
log-mel spectrogram of 80 bins and 3000 frames, padded to 30 seconds.
`tokens` begins with Whisper's prompt of special tokens, for example
`<|startoftranscript|> <|en|> <|transcribe|> <|notimestamps|>`, which the
Hub repository's tokenizer supplies.
| Entry | |
| --- | --- |
| `encode<B>(mel)` | encoder states for a 30-second window |
| `decode<B, S, A>(tokens, audio)` | logits for every token position |
`decode` recomputes the whole prefix each step; the cached entries below
do not. The encoder attends over all 1500 positions, including the silence
a shorter clip was padded with, exactly as the reference encoder does;
`std.nn.attention::attention` takes an unbatched `Tensor[Q, K; bool]` mask,
so a per-sequence padding mask could not be expressed even if the reference
used one.
## Transcribing with caches
`listen`, `prefill`, and `step` transcribe without recomputing: each
decoder layer keeps the keys and values of every token so far (`MaxTokens`
positions) and those of the encoder's states (1500 positions), sized by the
`Batch` generic (1 by default).
```python
model.run_entry("listen", [mel]) # encode; fill the cross-attention caches
logits = model.run_entry("prefill", [prompt]) # the prompt's tokens; logits after it
logits = model.run_entry("step", [token, position]) # one more token at `position`
```
| Entry | |
| --- | --- |
| `listen(mel)` | encoder states, and every decoder layer's cross-attention keys and values |
| `prefill<S>(tokens)` | the prompt from position 0 into the self-attention caches; logits after its last token |
| `step(token, pos)` | one token at `pos` over the caches; its logits |
Greedy transcription of the benchmark's clip gives the same tokens through
these entries as through `decode` over the whole prefix, on PyTorch (CUDA
graphs and generated source), JAX, and ONNX Runtime alike.
## Numerics
Whisper's `config.json` asks for plain `gelu`, the error-function form, and
its encoder applies that activation twice in the convolutional stem before
any transformer layer runs. The source uses `std.nn.activations::gelu_erf`.
The tanh approximation `gelu` would not do here: it perturbs the stem by
2.0e-3 per activation, and the four encoder layers amplify that into 5.7e-2
on the encoder states.
Against `WhisperForConditionalGeneration` exactly as published, in f32 on CPU
with a fixed random mel and the five-token prefix above:
| | max abs diff |
| --- | --- |
| encoder states, mel from `N(0, 1)` | 1.3e-4 |
| encoder states, mel from `U(-1, 1)`, a real log-mel's range | 3.7e-4 |
| decoder logits, end to end | 2.8e-5 |
| decoder logits, on the reference's own encoder states | 2.4e-5 |
The argmax agrees at every position. What remains is f32 rounding, most of it
the difference between `scaled_dot_product_attention` over 1500 positions and
the reference's eager attention; feeding the decoder the reference's own
encoder states isolates it and leaves 2e-5.
## Provenance
- Weights: [openai/whisper-tiny](https://huggingface.co/openai/whisper-tiny), Apache-2.0.
- Code: [openai/whisper](https://github.com/openai/whisper).
- Paper: [Robust Speech Recognition via Large-Scale Weak Supervision](https://arxiv.org/abs/2212.04356).bench.json7.9 kB
json
{
"date": "2026-09-28T19:33:54+00:00",
"environment": {
"python": "3.12.3",
"machine": "x86_64",
"torch": "2.14.0",
"jax": "0.11.2",
"transformers": "5.17.0",
"diffusers": "0.40.0",
"onnxruntime": "1.30.0",
"linnet": "0.1.0",
"gpu": "NVIDIA H100 80GB HBM3",
"gpus": "1",
"driver": "580.126.09"
},
"workload": {
"device": "cuda",
"prompt_tokens": 512,
"new_tokens": 128,
"batch": 1,
"warmup": 3,
"iters": 10,
"seed": 0,
"offload_gib": 8.0,
"serve_requests": 256,
"serve_concurrency": 64,
"serve_prompt_min": 128,
"serve_prompt_max": 512,
"serve_new": 128
},
"methods": [
{
"method": "transformers (eager)",
"kind": "reference",
"metrics": {
"encode_ms": 2.5331340730190277,
"transcribe_ms": 63.134714029729366,
"load_s": 1.9117679838091135,
"peak_vram_mib": 934.0
},
"max_abs_diff": 0.0,
"notes": "",
"error": null,
"key": "transformers-eager"
},
{
"method": "transformers (torch.compile, static cache)",
"kind": "reference",
"metrics": {
"encode_ms": 2.177765592932701,
"transcribe_ms": 51.773007959127426,
"load_s": 0.6697663590312004,
"peak_vram_mib": 1120.0
},
"max_abs_diff": 0.0008020401000976562,
"notes": "",
"error": null,
"key": "transformers-compile"
},
{
"method": "Linnet torch (generated source)",
"kind": "linnet",
"metrics": {
"encode_ms": 2.359773963689804,
"transcribe_ms": 58.74755047261715,
"load_s": 0.3051711730659008,
"peak_vram_mib": 1064.0
},
"max_abs_diff": 0.07889080047607422,
"notes": "transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-torch"
},
{
"method": "Linnet torch (CUDA graphs)",
"kind": "linnet",
"metrics": {
"encode_ms": 1.8138191662728786,
"transcribe_ms": 22.346744779497385,
"load_s": 0.3012003432959318,
"peak_vram_mib": 1184.0
},
"max_abs_diff": 0.07913780212402344,
"notes": "transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-cudagraphs"
},
{
"method": "Linnet torch (torch.compile, inductor)",
"kind": "linnet",
"metrics": {
"encode_ms": 1.9942112267017365,
"transcribe_ms": 34.278989769518375,
"load_s": 0.2968304455280304,
"peak_vram_mib": 1062.0
},
"max_abs_diff": 0.07913780212402344,
"notes": "transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-inductor"
},
{
"method": "Linnet JAX (XLA, StableHLO)",
"kind": "linnet",
"metrics": {
"encode_ms": 1.0177195072174072,
"transcribe_ms": 16.396986320614815,
"load_s": 0.2223910056054592,
"peak_vram_mib": 432.3,
"driver_vram_mib": 1206.9
},
"max_abs_diff": 0.6507987976074219,
"notes": "transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
"error": null,
"key": "linnet-jax"
},
{
"method": "Linnet JAX (XLA, generated source)",
"kind": "linnet",
"metrics": {
"encode_ms": 1.0147206485271454,
"transcribe_ms": 16.603680327534676,
"load_s": 0.21561535447835922,
"peak_vram_mib": 432.3,
"driver_vram_mib": 1206.9
},
"max_abs_diff": 0.6820688247680664,
"notes": "transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
"error": null,
"key": "linnet-jax-source"
},
{
"method": "Linnet ONNX f32 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"encode_ms": 1.7017573118209839,
"transcribe_ms": 16.175041440874338,
"load_s": 0.23641015402972698,
"peak_vram_mib": 1846.0
},
"max_abs_diff": 0.7507038116455078,
"notes": "f32; ONNX Runtime's CUDA execution provider; first calls build the sessions; transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-onnx-f32"
},
{
"method": "Linnet ONNX f32 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"encode_ms": 0.7919538766145706,
"transcribe_ms": 24.299739859998226,
"load_s": 0.15421363804489374,
"peak_vram_mib": 3418.0
},
"max_abs_diff": 0.7840394973754883,
"notes": "f32; ONNX Runtime's TensorRT execution provider; first calls build the sessions; transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-onnx-f32-trt"
},
{
"method": "Linnet ONNX f16 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"encode_ms": 1.6398252919316292,
"transcribe_ms": 16.539379954338074,
"load_s": 0.2333100289106369,
"peak_vram_mib": 1674.0
},
"max_abs_diff": 0.7178497314453125,
"notes": "f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-onnx-f16"
},
{
"method": "Linnet ONNX f16 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"encode_ms": 0.4132934845983982,
"transcribe_ms": 20.40411438792944,
"load_s": 0.24507786612957716,
"peak_vram_mib": 3142.0
},
"max_abs_diff": 1.3493905067443848,
"notes": "f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-onnx-f16-trt"
},
{
"method": "Linnet ONNX bf16 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"encode_ms": 1.257970929145813,
"transcribe_ms": 17.268605530261993,
"load_s": 0.24303966015577316,
"peak_vram_mib": 1552.0
},
"max_abs_diff": 8.531454086303711,
"notes": "bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-onnx-bf16"
},
{
"method": "Linnet ONNX bf16 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"encode_ms": 0.45439181849360466,
"transcribe_ms": 20.365712698549032,
"load_s": 0.14582019858062267,
"peak_vram_mib": 3080.0
},
"max_abs_diff": 5.794703006744385,
"notes": "bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; transcribed with the card's KV caches: `listen` encodes and fills the cross-attention caches, `prefill` feeds the prompt, `step` one token each",
"error": null,
"key": "linnet-onnx-bf16-trt"
}
],
"reference": "transformers-eager"
}bindings.json14.7 kB
json
{
"conv1_weight": "model.encoder.conv1.weight",
"conv1_bias": "model.encoder.conv1.bias",
"conv2_weight": "model.encoder.conv2.weight",
"conv2_bias": "model.encoder.conv2.bias",
"encoder_positions": "model.encoder.embed_positions.weight",
"encoder_layers.0.self_attn_layer_norm.weight": "model.encoder.layers.0.self_attn_layer_norm.weight",
"encoder_layers.0.self_attn_layer_norm.bias": "model.encoder.layers.0.self_attn_layer_norm.bias",
"encoder_layers.0.self_attn.q_proj.weight": "model.encoder.layers.0.self_attn.q_proj.weight",
"encoder_layers.0.self_attn.q_proj.bias": "model.encoder.layers.0.self_attn.q_proj.bias",
"encoder_layers.0.self_attn.k_proj.weight": "model.encoder.layers.0.self_attn.k_proj.weight",
"encoder_layers.0.self_attn.v_proj.weight": "model.encoder.layers.0.self_attn.v_proj.weight",
"encoder_layers.0.self_attn.v_proj.bias": "model.encoder.layers.0.self_attn.v_proj.bias",
"encoder_layers.0.self_attn.out_proj.weight": "model.encoder.layers.0.self_attn.out_proj.weight",
"encoder_layers.0.self_attn.out_proj.bias": "model.encoder.layers.0.self_attn.out_proj.bias",
"encoder_layers.0.final_layer_norm.weight": "model.encoder.layers.0.final_layer_norm.weight",
"encoder_layers.0.final_layer_norm.bias": "model.encoder.layers.0.final_layer_norm.bias",
"encoder_layers.0.mlp.fc1.weight": "model.encoder.layers.0.fc1.weight",
"encoder_layers.0.mlp.fc1.bias": "model.encoder.layers.0.fc1.bias",
"encoder_layers.0.mlp.fc2.weight": "model.encoder.layers.0.fc2.weight",
"encoder_layers.0.mlp.fc2.bias": "model.encoder.layers.0.fc2.bias",
"encoder_layers.1.self_attn_layer_norm.weight": "model.encoder.layers.1.self_attn_layer_norm.weight",
"encoder_layers.1.self_attn_layer_norm.bias": "model.encoder.layers.1.self_attn_layer_norm.bias",
"encoder_layers.1.self_attn.q_proj.weight": "model.encoder.layers.1.self_attn.q_proj.weight",
"encoder_layers.1.self_attn.q_proj.bias": "model.encoder.layers.1.self_attn.q_proj.bias",
"encoder_layers.1.self_attn.k_proj.weight": "model.encoder.layers.1.self_attn.k_proj.weight",
"encoder_layers.1.self_attn.v_proj.weight": "model.encoder.layers.1.self_attn.v_proj.weight",
"encoder_layers.1.self_attn.v_proj.bias": "model.encoder.layers.1.self_attn.v_proj.bias",
"encoder_layers.1.self_attn.out_proj.weight": "model.encoder.layers.1.self_attn.out_proj.weight",
"encoder_layers.1.self_attn.out_proj.bias": "model.encoder.layers.1.self_attn.out_proj.bias",
"encoder_layers.1.final_layer_norm.weight": "model.encoder.layers.1.final_layer_norm.weight",
"encoder_layers.1.final_layer_norm.bias": "model.encoder.layers.1.final_layer_norm.bias",
"encoder_layers.1.mlp.fc1.weight": "model.encoder.layers.1.fc1.weight",
"encoder_layers.1.mlp.fc1.bias": "model.encoder.layers.1.fc1.bias",
"encoder_layers.1.mlp.fc2.weight": "model.encoder.layers.1.fc2.weight",
"encoder_layers.1.mlp.fc2.bias": "model.encoder.layers.1.fc2.bias",
"encoder_layers.2.self_attn_layer_norm.weight": "model.encoder.layers.2.self_attn_layer_norm.weight",
"encoder_layers.2.self_attn_layer_norm.bias": "model.encoder.layers.2.self_attn_layer_norm.bias",
"encoder_layers.2.self_attn.q_proj.weight": "model.encoder.layers.2.self_attn.q_proj.weight",
"encoder_layers.2.self_attn.q_proj.bias": "model.encoder.layers.2.self_attn.q_proj.bias",
"encoder_layers.2.self_attn.k_proj.weight": "model.encoder.layers.2.self_attn.k_proj.weight",
"encoder_layers.2.self_attn.v_proj.weight": "model.encoder.layers.2.self_attn.v_proj.weight",
"encoder_layers.2.self_attn.v_proj.bias": "model.encoder.layers.2.self_attn.v_proj.bias",
"encoder_layers.2.self_attn.out_proj.weight": "model.encoder.layers.2.self_attn.out_proj.weight",
"encoder_layers.2.self_attn.out_proj.bias": "model.encoder.layers.2.self_attn.out_proj.bias",
"encoder_layers.2.final_layer_norm.weight": "model.encoder.layers.2.final_layer_norm.weight",
"encoder_layers.2.final_layer_norm.bias": "model.encoder.layers.2.final_layer_norm.bias",
"encoder_layers.2.mlp.fc1.weight": "model.encoder.layers.2.fc1.weight",
"encoder_layers.2.mlp.fc1.bias": "model.encoder.layers.2.fc1.bias",
"encoder_layers.2.mlp.fc2.weight": "model.encoder.layers.2.fc2.weight",
"encoder_layers.2.mlp.fc2.bias": "model.encoder.layers.2.fc2.bias",
"encoder_layers.3.self_attn_layer_norm.weight": "model.encoder.layers.3.self_attn_layer_norm.weight",
"encoder_layers.3.self_attn_layer_norm.bias": "model.encoder.layers.3.self_attn_layer_norm.bias",
"encoder_layers.3.self_attn.q_proj.weight": "model.encoder.layers.3.self_attn.q_proj.weight",
"encoder_layers.3.self_attn.q_proj.bias": "model.encoder.layers.3.self_attn.q_proj.bias",
"encoder_layers.3.self_attn.k_proj.weight": "model.encoder.layers.3.self_attn.k_proj.weight",
"encoder_layers.3.self_attn.v_proj.weight": "model.encoder.layers.3.self_attn.v_proj.weight",
"encoder_layers.3.self_attn.v_proj.bias": "model.encoder.layers.3.self_attn.v_proj.bias",
"encoder_layers.3.self_attn.out_proj.weight": "model.encoder.layers.3.self_attn.out_proj.weight",
"encoder_layers.3.self_attn.out_proj.bias": "model.encoder.layers.3.self_attn.out_proj.bias",
"encoder_layers.3.final_layer_norm.weight": "model.encoder.layers.3.final_layer_norm.weight",
"encoder_layers.3.final_layer_norm.bias": "model.encoder.layers.3.final_layer_norm.bias",
"encoder_layers.3.mlp.fc1.weight": "model.encoder.layers.3.fc1.weight",
"encoder_layers.3.mlp.fc1.bias": "model.encoder.layers.3.fc1.bias",
"encoder_layers.3.mlp.fc2.weight": "model.encoder.layers.3.fc2.weight",
"encoder_layers.3.mlp.fc2.bias": "model.encoder.layers.3.fc2.bias",
"encoder_norm.weight": "model.encoder.layer_norm.weight",
"encoder_norm.bias": "model.encoder.layer_norm.bias",
"token_embedding": "model.decoder.embed_tokens.weight",
"decoder_positions": "model.decoder.embed_positions.weight",
"decoder_layers.0.self_attn_layer_norm.weight": "model.decoder.layers.0.self_attn_layer_norm.weight",
"decoder_layers.0.self_attn_layer_norm.bias": "model.decoder.layers.0.self_attn_layer_norm.bias",
"decoder_layers.0.self_attn.q_proj.weight": "model.decoder.layers.0.self_attn.q_proj.weight",
"decoder_layers.0.self_attn.q_proj.bias": "model.decoder.layers.0.self_attn.q_proj.bias",
"decoder_layers.0.self_attn.k_proj.weight": "model.decoder.layers.0.self_attn.k_proj.weight",
"decoder_layers.0.self_attn.v_proj.weight": "model.decoder.layers.0.self_attn.v_proj.weight",
"decoder_layers.0.self_attn.v_proj.bias": "model.decoder.layers.0.self_attn.v_proj.bias",
"decoder_layers.0.self_attn.out_proj.weight": "model.decoder.layers.0.self_attn.out_proj.weight",
"decoder_layers.0.self_attn.out_proj.bias": "model.decoder.layers.0.self_attn.out_proj.bias",
"decoder_layers.0.encoder_attn_layer_norm.weight": "model.decoder.layers.0.encoder_attn_layer_norm.weight",
"decoder_layers.0.encoder_attn_layer_norm.bias": "model.decoder.layers.0.encoder_attn_layer_norm.bias",
"decoder_layers.0.encoder_attn.q_proj.weight": "model.decoder.layers.0.encoder_attn.q_proj.weight",
"decoder_layers.0.encoder_attn.q_proj.bias": "model.decoder.layers.0.encoder_attn.q_proj.bias",
"decoder_layers.0.encoder_attn.k_proj.weight": "model.decoder.layers.0.encoder_attn.k_proj.weight",
"decoder_layers.0.encoder_attn.v_proj.weight": "model.decoder.layers.0.encoder_attn.v_proj.weight",
"decoder_layers.0.encoder_attn.v_proj.bias": "model.decoder.layers.0.encoder_attn.v_proj.bias",
"decoder_layers.0.encoder_attn.out_proj.weight": "model.decoder.layers.0.encoder_attn.out_proj.weight",
"decoder_layers.0.encoder_attn.out_proj.bias": "model.decoder.layers.0.encoder_attn.out_proj.bias",
"decoder_layers.0.final_layer_norm.weight": "model.decoder.layers.0.final_layer_norm.weight",
"decoder_layers.0.final_layer_norm.bias": "model.decoder.layers.0.final_layer_norm.bias",
"decoder_layers.0.mlp.fc1.weight": "model.decoder.layers.0.fc1.weight",
"decoder_layers.0.mlp.fc1.bias": "model.decoder.layers.0.fc1.bias",
"decoder_layers.0.mlp.fc2.weight": "model.decoder.layers.0.fc2.weight",
"decoder_layers.0.mlp.fc2.bias": "model.decoder.layers.0.fc2.bias",
"decoder_layers.1.self_attn_layer_norm.weight": "model.decoder.layers.1.self_attn_layer_norm.weight",
"decoder_layers.1.self_attn_layer_norm.bias": "model.decoder.layers.1.self_attn_layer_norm.bias",
"decoder_layers.1.self_attn.q_proj.weight": "model.decoder.layers.1.self_attn.q_proj.weight",
"decoder_layers.1.self_attn.q_proj.bias": "model.decoder.layers.1.self_attn.q_proj.bias",
"decoder_layers.1.self_attn.k_proj.weight": "model.decoder.layers.1.self_attn.k_proj.weight",
"decoder_layers.1.self_attn.v_proj.weight": "model.decoder.layers.1.self_attn.v_proj.weight",
"decoder_layers.1.self_attn.v_proj.bias": "model.decoder.layers.1.self_attn.v_proj.bias",
"decoder_layers.1.self_attn.out_proj.weight": "model.decoder.layers.1.self_attn.out_proj.weight",
"decoder_layers.1.self_attn.out_proj.bias": "model.decoder.layers.1.self_attn.out_proj.bias",
"decoder_layers.1.encoder_attn_layer_norm.weight": "model.decoder.layers.1.encoder_attn_layer_norm.weight",
"decoder_layers.1.encoder_attn_layer_norm.bias": "model.decoder.layers.1.encoder_attn_layer_norm.bias",
"decoder_layers.1.encoder_attn.q_proj.weight": "model.decoder.layers.1.encoder_attn.q_proj.weight",
"decoder_layers.1.encoder_attn.q_proj.bias": "model.decoder.layers.1.encoder_attn.q_proj.bias",
"decoder_layers.1.encoder_attn.k_proj.weight": "model.decoder.layers.1.encoder_attn.k_proj.weight",
"decoder_layers.1.encoder_attn.v_proj.weight": "model.decoder.layers.1.encoder_attn.v_proj.weight",
"decoder_layers.1.encoder_attn.v_proj.bias": "model.decoder.layers.1.encoder_attn.v_proj.bias",
"decoder_layers.1.encoder_attn.out_proj.weight": "model.decoder.layers.1.encoder_attn.out_proj.weight",
"decoder_layers.1.encoder_attn.out_proj.bias": "model.decoder.layers.1.encoder_attn.out_proj.bias",
"decoder_layers.1.final_layer_norm.weight": "model.decoder.layers.1.final_layer_norm.weight",
"decoder_layers.1.final_layer_norm.bias": "model.decoder.layers.1.final_layer_norm.bias",
"decoder_layers.1.mlp.fc1.weight": "model.decoder.layers.1.fc1.weight",
"decoder_layers.1.mlp.fc1.bias": "model.decoder.layers.1.fc1.bias",
"decoder_layers.1.mlp.fc2.weight": "model.decoder.layers.1.fc2.weight",
"decoder_layers.1.mlp.fc2.bias": "model.decoder.layers.1.fc2.bias",
"decoder_layers.2.self_attn_layer_norm.weight": "model.decoder.layers.2.self_attn_layer_norm.weight",
"decoder_layers.2.self_attn_layer_norm.bias": "model.decoder.layers.2.self_attn_layer_norm.bias",
"decoder_layers.2.self_attn.q_proj.weight": "model.decoder.layers.2.self_attn.q_proj.weight",
"decoder_layers.2.self_attn.q_proj.bias": "model.decoder.layers.2.self_attn.q_proj.bias",
"decoder_layers.2.self_attn.k_proj.weight": "model.decoder.layers.2.self_attn.k_proj.weight",
"decoder_layers.2.self_attn.v_proj.weight": "model.decoder.layers.2.self_attn.v_proj.weight",
"decoder_layers.2.self_attn.v_proj.bias": "model.decoder.layers.2.self_attn.v_proj.bias",
"decoder_layers.2.self_attn.out_proj.weight": "model.decoder.layers.2.self_attn.out_proj.weight",
"decoder_layers.2.self_attn.out_proj.bias": "model.decoder.layers.2.self_attn.out_proj.bias",
"decoder_layers.2.encoder_attn_layer_norm.weight": "model.decoder.layers.2.encoder_attn_layer_norm.weight",
"decoder_layers.2.encoder_attn_layer_norm.bias": "model.decoder.layers.2.encoder_attn_layer_norm.bias",
"decoder_layers.2.encoder_attn.q_proj.weight": "model.decoder.layers.2.encoder_attn.q_proj.weight",
"decoder_layers.2.encoder_attn.q_proj.bias": "model.decoder.layers.2.encoder_attn.q_proj.bias",
"decoder_layers.2.encoder_attn.k_proj.weight": "model.decoder.layers.2.encoder_attn.k_proj.weight",
"decoder_layers.2.encoder_attn.v_proj.weight": "model.decoder.layers.2.encoder_attn.v_proj.weight",
"decoder_layers.2.encoder_attn.v_proj.bias": "model.decoder.layers.2.encoder_attn.v_proj.bias",
"decoder_layers.2.encoder_attn.out_proj.weight": "model.decoder.layers.2.encoder_attn.out_proj.weight",
"decoder_layers.2.encoder_attn.out_proj.bias": "model.decoder.layers.2.encoder_attn.out_proj.bias",
"decoder_layers.2.final_layer_norm.weight": "model.decoder.layers.2.final_layer_norm.weight",
"decoder_layers.2.final_layer_norm.bias": "model.decoder.layers.2.final_layer_norm.bias",
"decoder_layers.2.mlp.fc1.weight": "model.decoder.layers.2.fc1.weight",
"decoder_layers.2.mlp.fc1.bias": "model.decoder.layers.2.fc1.bias",
"decoder_layers.2.mlp.fc2.weight": "model.decoder.layers.2.fc2.weight",
"decoder_layers.2.mlp.fc2.bias": "model.decoder.layers.2.fc2.bias",
"decoder_layers.3.self_attn_layer_norm.weight": "model.decoder.layers.3.self_attn_layer_norm.weight",
"decoder_layers.3.self_attn_layer_norm.bias": "model.decoder.layers.3.self_attn_layer_norm.bias",
"decoder_layers.3.self_attn.q_proj.weight": "model.decoder.layers.3.self_attn.q_proj.weight",
"decoder_layers.3.self_attn.q_proj.bias": "model.decoder.layers.3.self_attn.q_proj.bias",
"decoder_layers.3.self_attn.k_proj.weight": "model.decoder.layers.3.self_attn.k_proj.weight",
"decoder_layers.3.self_attn.v_proj.weight": "model.decoder.layers.3.self_attn.v_proj.weight",
"decoder_layers.3.self_attn.v_proj.bias": "model.decoder.layers.3.self_attn.v_proj.bias",
"decoder_layers.3.self_attn.out_proj.weight": "model.decoder.layers.3.self_attn.out_proj.weight",
"decoder_layers.3.self_attn.out_proj.bias": "model.decoder.layers.3.self_attn.out_proj.bias",
"decoder_layers.3.encoder_attn_layer_norm.weight": "model.decoder.layers.3.encoder_attn_layer_norm.weight",
"decoder_layers.3.encoder_attn_layer_norm.bias": "model.decoder.layers.3.encoder_attn_layer_norm.bias",
"decoder_layers.3.encoder_attn.q_proj.weight": "model.decoder.layers.3.encoder_attn.q_proj.weight",
"decoder_layers.3.encoder_attn.q_proj.bias": "model.decoder.layers.3.encoder_attn.q_proj.bias",
"decoder_layers.3.encoder_attn.k_proj.weight": "model.decoder.layers.3.encoder_attn.k_proj.weight",
"decoder_layers.3.encoder_attn.v_proj.weight": "model.decoder.layers.3.encoder_attn.v_proj.weight",
"decoder_layers.3.encoder_attn.v_proj.bias": "model.decoder.layers.3.encoder_attn.v_proj.bias",
"decoder_layers.3.encoder_attn.out_proj.weight": "model.decoder.layers.3.encoder_attn.out_proj.weight",
"decoder_layers.3.encoder_attn.out_proj.bias": "model.decoder.layers.3.encoder_attn.out_proj.bias",
"decoder_layers.3.final_layer_norm.weight": "model.decoder.layers.3.final_layer_norm.weight",
"decoder_layers.3.final_layer_norm.bias": "model.decoder.layers.3.final_layer_norm.bias",
"decoder_layers.3.mlp.fc1.weight": "model.decoder.layers.3.fc1.weight",
"decoder_layers.3.mlp.fc1.bias": "model.decoder.layers.3.fc1.bias",
"decoder_layers.3.mlp.fc2.weight": "model.decoder.layers.3.fc2.weight",
"decoder_layers.3.mlp.fc2.bias": "model.decoder.layers.3.fc2.bias",
"decoder_norm.weight": "model.decoder.layer_norm.weight",
"decoder_norm.bias": "model.decoder.layer_norm.bias"
}linnet.toml62 B
toml
[package]
name = "whisper"
version = "0.1.0"
language = "0.1"nest.toml921 B
toml
[model]
name = "whisper-tiny"
title = "Whisper tiny"
summary = "OpenAI's smallest multilingual speech recognition model: a convolutional stem over a log-mel spectrogram, four encoder layers, and a four-layer decoder that cross-attends to them."
license = "Apache-2.0"
family = "whisper"
tags = ["automatic-speech-recognition", "audio", "encoder-decoder", "cross-attention", "multilingual"]
[links]
huggingface = "https://huggingface.co/openai/whisper-tiny"
github = "https://github.com/openai/whisper"
arxiv = "https://arxiv.org/abs/2212.04356"
[source]
path = "src/lib.linnet"
root = "Model"
entry = "encode"
[generics]
Mels = 80
Frames = 3000
D = 384
Heads = 6
Inner = 1536
EncoderLayers = 4
DecoderLayers = 4
Vocab = 51865
MaxTokens = 448
T = "f32"
[check]
B = 1
[weights]
repo = "openai/whisper-tiny"
revision = "169d4a4341b33bc18d8881c4b69c2e104e1cc0af"
files = ["model.safetensors"]
bindings = "bindings.json"samples.json319 B
json
{
"kind": "transcription",
"audio": "hf-internal-testing/librispeech_asr_dummy, clean/validation[0]",
"text": "Mr. Quilter is the apostle of the middle classes and we are glad to welcome his gospel.",
"reference_text": "Mr. Quilter is the apostle of the middle classes and we are glad to welcome his gospel."
}src
lib.linnet14.6 kB
linnet
// Whisper: speech recognition as an encoder-decoder transformer. The encoder
// reads a log-mel spectrogram through two 1-D convolutions that halve the
// frame rate, adds the position table the checkpoint stores, and runs pre-norm
// encoder layers over the result. The decoder is an ordinary causal language
// model with a cross-attention over those encoder states inserted between its
// self-attention and its MLP, and an output head tied to the token embedding.
//
// The two are separate entries because transcription uses them separately:
// the encoder runs once over the 30 seconds of audio, the decoder runs once
// per token against the states it produced.
module whisper
use std.nn.activations::{gelu_erf}
use std.nn.attention::{attention, causal_mask}
use std.nn.cache::{write_at, write_span}
use std.nn.embedding::{embedding}
use std.nn.linear::{Linear, linear}
use std.nn.norm::{layer_norm}
// Whisper normalizes with `torch.nn.LayerNorm`'s default epsilon.
pub const EPS: f32 = 1e-5
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), EPS)
}
}
// Zero padding on the position axis of a `[B, Positions, Channels]` signal.
// `P = 0` concatenates empty tensors, which every backend folds away.
fn pad1d<B: Dim, L: Dim, C: Dim, P: Dim, T: Float>(
x: Tensor[B, L, C; T],
) -> Tensor[B, L + 2 * P, C; T] {
let edge = fill<T>([B, P, C], 0.0)
return concat(concat(edge, x, axis = 1), edge, axis = 1)
}
// A 1-D convolution, which the standard library does not have: `conv2d` takes
// a square kernel, and the stem's kernel is a line of three frames. The
// window geometry is built the way `std.nn.conv::conv2d` builds it, because an
// index expression may not do arithmetic: the tap positions are a tensor made
// from `iota` and the padded signal is gathered through it.
//
// Signals are `[B, Positions, Channels]` here, the layout the rest of the
// encoder wants, so the gathered window can be flattened over (channel, tap)
// and the convolution finished as one `linear` against the kernel flattened
// the same way. That matters: there is no `conv1d` op for a backend to
// recognize, so a five-axis product written out in index notation would stay
// a five-axis product, while `linear` selects a matrix multiply everywhere.
fn conv1d<
B: Dim,
L: Dim,
Cin: Dim,
Cout: Dim,
K: Dim,
Stride: Dim,
Pad: Dim,
T: Float,
>(
x: Tensor[B, L, Cin; T],
weight: Tensor[Cout, Cin, K; T],
bias: Tensor[Cout; T],
) -> Tensor[B, (L + 2 * Pad - K) / Stride + 1, Cout; T]
where
Stride > 0,
K > 0
{
let padded = pad1d<B, L, Cin, Pad, T>(x)
// Window `o` starts at `o * Stride`, so tap `k` of it reads position
// `o * Stride + k` of the padded signal.
let taps[o, k] =
iota<i32>((L + 2 * Pad - K) / Stride + 1)[o] * cast<i32>(Stride) + iota<i32>(K)[k]
// The gather transposes as it reads, so that (channel, tap) flattens in
// the order the checkpoint's `[Cout, Cin, K]` kernel flattens in.
let windows[b, o, ci, k] = padded[b, taps[o, k], ci]
return linear(
reshape(windows, [B, (L + 2 * Pad - K) / Stride + 1, Cin * K]),
reshape(weight, [Cout, Cin * K]),
some(bias),
)
}
// Attention over an explicit key/value source, which is what lets one block
// serve both self-attention (the source is the sequence itself) and
// cross-attention (the source is the encoder's states).
//
// Whisper projects queries, values, and the output with a bias but keys
// without one. The asymmetry is in the reference implementation and it is
// real: a key bias shifts every key by the same vector, which a softmax over
// dot products does not ignore, so leaving it out is a choice the checkpoint
// was trained with rather than a redundancy.
pub block Attention<D: Dim, Heads: Dim, T: Float>
where
Heads > 0,
D % Heads == 0
{
sub q_proj: Linear<D, D, T>
sub k_proj: Linear<D, D, T>
sub v_proj: Linear<D, D, T>
sub out_proj: Linear<D, D, T>
pub fn forward<B: Dim, Q: Dim, Src: Dim>(
x: Tensor[B, Q, D; T],
source: Tensor[B, Src, D; T],
mask: Tensor[Q, Src; bool]?,
) -> Tensor[B, Q, D; T] {
return attend(x, keys(source), values(source), mask)
}
// The keys and the values of `source`, split into heads: what a cache
// holds so that later tokens do not project them again.
pub fn keys<B: Dim, N: Dim>(source: Tensor[B, N, D; T]) -> Tensor[B, Heads, N, D / Heads; T] {
return heads<B, N, Heads, D / Heads, T>(k_proj.forward(source))
}
pub fn values<B: Dim, N: Dim>(
source: Tensor[B, N, D; T],
) -> Tensor[B, Heads, N, D / Heads; T] {
return heads<B, N, Heads, D / Heads, T>(v_proj.forward(source))
}
// `x`'s queries over keys and values already split into heads.
pub fn attend<B: Dim, Q: Dim, K: Dim>(
x: Tensor[B, Q, D; T],
keys: Tensor[B, Heads, K, D / Heads; T],
values: Tensor[B, Heads, K, D / Heads; T],
mask: Tensor[Q, K; bool]?,
) -> Tensor[B, Q, D; T] {
let mixed = attention(
heads<B, Q, Heads, D / Heads, T>(q_proj.forward(x)),
keys,
values,
rsqrt(cast<f32>(D / Heads)),
mask,
)
return out_proj.forward(reshape(permute(mixed, [0, 2, 1, 3]), [B, Q, D]))
}
}
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])
}
// The MLP of both stacks: one widening projection, GELU, one narrowing
// projection.
pub block Mlp<D: Dim, Inner: Dim, T: Float> {
sub fc1: Linear<D, Inner, T>
sub fc2: Linear<Inner, D, T>
pub fn forward<*S: Shape>(x: Tensor[*S, D; T]) -> Tensor[*S, D; T] {
return fc2.forward(gelu_erf(fc1.forward(x)))
}
}
// Every audio position attends to every other: the spectrogram is padded to a
// fixed 30 seconds, and the reference encoder attends over the padding too.
pub block EncoderLayer<D: Dim, Heads: Dim, Inner: Dim, T: Float>
where
Heads > 0,
D % Heads == 0
{
sub self_attn_layer_norm: LayerNorm<D, T>
sub self_attn: Attention<D, Heads, T>
sub final_layer_norm: LayerNorm<D, T>
sub mlp: Mlp<D, Inner, T>
pub fn forward<B: Dim, N: Dim>(x: Tensor[B, N, D; T]) -> Tensor[B, N, D; T] {
let normed = self_attn_layer_norm.forward(x)
let attended = x + self_attn.forward(normed, normed, none)
return attended + mlp.forward(final_layer_norm.forward(attended))
}
}
// Three residual steps: causal self-attention over the tokens decoded so far,
// then attention over the encoder's states, then the MLP. The cross-attention
// needs no mask, since every token may read the whole utterance.
//
// For decoding one token at a time the layer keeps what it would otherwise
// recompute: the keys and values of every token so far (`self_k`, `self_v`,
// `MaxTokens` positions) and those of the encoder's states (`cross_k`,
// `cross_v`, `Audio` positions), which `listen` fills once per utterance.
pub block DecoderLayer<
D: Dim,
Heads: Dim,
Inner: Dim,
Batch: Dim,
MaxTokens: Dim,
Audio: Dim,
T: Float,
>
where
Heads > 0,
D % Heads == 0,
MaxTokens > 0
{
sub self_attn_layer_norm: LayerNorm<D, T>
sub self_attn: Attention<D, Heads, T>
sub encoder_attn_layer_norm: LayerNorm<D, T>
sub encoder_attn: Attention<D, Heads, T>
sub final_layer_norm: LayerNorm<D, T>
sub mlp: Mlp<D, Inner, T>
state self_k: Tensor[Batch, Heads, MaxTokens, D / Heads; T]
state self_v: Tensor[Batch, Heads, MaxTokens, D / Heads; T]
state cross_k: Tensor[Batch, Heads, Audio, D / Heads; T]
state cross_v: Tensor[Batch, Heads, Audio, D / Heads; T]
pub fn forward<B: Dim, S: Dim, A: Dim>(
x: Tensor[B, S, D; T],
audio: Tensor[B, A, D; T],
) -> Tensor[B, S, D; T] {
let normed = self_attn_layer_norm.forward(x)
let attended = x + self_attn.forward(normed, normed, some(causal_mask<S, S>()))
let crossed =
attended + encoder_attn.forward(encoder_attn_layer_norm.forward(attended), audio, none)
return crossed + mlp.forward(final_layer_norm.forward(crossed))
}
// The encoder's states as this layer's cross-attention reads them.
pub fn listen(audio: Tensor[Batch, Audio, D; T]) -> Tensor[Batch, Audio, D; T] {
cross_k = encoder_attn.keys(audio)
cross_v = encoder_attn.values(audio)
return audio
}
// A prompt of `S` tokens from the first position: their keys and values
// go into the caches, and they attend causally among themselves.
pub fn prefill<S: Dim>(x: Tensor[Batch, S, D; T]) -> Tensor[Batch, S, D; T]
where S > 0 {
let normed = self_attn_layer_norm.forward(x)
let k = self_attn.keys(normed)
let v = self_attn.values(normed)
self_k = write_span(self_k, k, 0)
self_v = write_span(self_v, v, 0)
let attended = x + self_attn.attend(normed, k, v, some(causal_mask<S, S>()))
return cross(attended)
}
// One token at position `pos`, over the tokens before it in the cache.
pub fn step(x: Tensor[Batch, 1, D; T], pos: i32) -> Tensor[Batch, 1, D; T] {
let normed = self_attn_layer_norm.forward(x)
self_k = write_at(self_k, self_attn.keys(normed), pos)
self_v = write_at(self_v, self_attn.values(normed), pos)
let slots = iota<i32>(MaxTokens)
let seen[t] = slots[t] <= pos
let attended =
x + self_attn.attend(normed, self_k, self_v, some(reshape(seen, [1, MaxTokens])))
return cross(attended)
}
// Cross-attention over the cached encoder states, then the MLP.
fn cross<S: Dim>(attended: Tensor[Batch, S, D; T]) -> Tensor[Batch, S, D; T] {
let crossed =
attended +
encoder_attn.attend(
encoder_attn_layer_norm.forward(attended),
cross_k,
cross_v,
none,
)
return crossed + mlp.forward(final_layer_norm.forward(crossed))
}
}
pub block Model<
Mels: Dim,
Frames: Dim,
D: Dim,
Heads: Dim,
Inner: Dim,
EncoderLayers: Dim,
DecoderLayers: Dim,
Vocab: Dim,
MaxTokens: Dim,
T: Float = f32,
Batch: Dim = 1,
>
where
Heads > 0,
D % Heads == 0,
Frames > 0,
MaxTokens > 0
{
// The stem, with the `[Cout, Cin, K]` kernels the checkpoint stores.
param conv1_weight: Tensor[D, Mels, 3; T]
param conv1_bias: Tensor[D; T]
param conv2_weight: Tensor[D, D, 3; T]
param conv2_bias: Tensor[D; T]
// Sinusoidal in origin, but the checkpoint stores it as a tensor like any
// other, and its length follows from the second convolution's stride.
param encoder_positions: Tensor[(Frames - 1) / 2 + 1, D; T]
sub encoder_layers: [EncoderLayer<D, Heads, Inner, T>; EncoderLayers]
sub encoder_norm: LayerNorm<D, T>
param token_embedding: Tensor[Vocab, D; T]
param decoder_positions: Tensor[MaxTokens, D; T]
sub decoder_layers: [DecoderLayer<D, Heads, Inner, Batch, MaxTokens, (Frames - 1) / 2 +
1, T>; DecoderLayers]
sub decoder_norm: LayerNorm<D, T>
// The spectrogram to the encoder's states: `Frames` mel frames become
// `(Frames - 1) / 2 + 1` positions, the form the shape solver derives
// from a stride-2 kernel of three with one frame of padding on each side.
// The mel bins arrive as the channel axis, the layout the feature
// extractor and the checkpoint's kernels use, and the permutation to
// positions-first happens once here.
pub entry encode<B: Dim>(
mel: Tensor[B, Mels, Frames; T],
) -> Tensor[B, (Frames - 1) / 2 + 1, D; T] {
return encoded(mel)
}
fn encoded<B: Dim>(mel: Tensor[B, Mels, Frames; T]) -> Tensor[B, (Frames - 1) / 2 + 1, D; T] {
let signal = permute(mel, [0, 2, 1])
let first = gelu_erf(
conv1d<B, Frames, Mels, D, 3, 1, 1, T>(signal, conv1_weight, conv1_bias),
)
let second = gelu_erf(conv1d<B, Frames, D, D, 3, 2, 1, T>(first, conv2_weight, conv2_bias))
var x = second + encoder_positions
static for layer in encoder_layers {
x = layer.forward(x)
}
return encoder_norm.forward(x)
}
// Tokens and encoder states to logits for every token position. The audio
// length is its own generic rather than the encoder's output length, so a
// shorter set of states is accepted as it is.
pub entry decode<B: Dim, S: Dim, A: Dim>(
tokens: Tensor[B, S; i32],
audio: Tensor[B, A, D; T],
) -> Tensor[B, S, Vocab; T]
where S <= MaxTokens {
// Positions 0..S: the decoder is fed from the start of the transcript.
var x = embedding(tokens, token_embedding) + embedding(iota<i32>(S), decoder_positions)
static for layer in decoder_layers {
x = layer.forward(x, audio)
}
// The head is the token embedding read as a projection; Whisper ties
// the two, so the checkpoint has no separate output matrix.
return linear(decoder_norm.forward(x), token_embedding)
}
// Transcribing with caches: `listen` encodes the spectrogram and gives
// every decoder layer the keys and values of its states, `prefill`
// feeds the prompt (start of transcript, language, task) and returns
// the logits after it, and `step` feeds one more token at `pos`,
// reading the caches instead of recomputing every token before it.
pub entry listen(mel: Tensor[Batch, Mels, Frames; T]) -> Tensor[Batch, (Frames - 1) / 2 +
1, D; T] {
let audio = encoded(mel)
static for layer in decoder_layers {
let _ = layer.listen(audio)
}
return audio
}
pub entry prefill<S: Dim>(tokens: Tensor[Batch, S; i32]) -> Tensor[Batch, Vocab; T]
where
S > 0,
S <= MaxTokens
{
var x = embedding(tokens, token_embedding) + embedding(iota<i32>(S), decoder_positions)
static for layer in decoder_layers {
x = layer.prefill(x)
}
return linear(decoder_norm.forward(x[:, S - 1, :]), token_embedding)
}
pub entry step(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] {
let position[d] = decoder_positions[cast<i64>(pos), d]
var x = embedding(token, token_embedding) + reshape(position, [1, 1, D])
static for layer in decoder_layers {
x = layer.step(x, pos)
}
return linear(decoder_norm.forward(x[:, 0, :]), token_embedding)
}
}