Models / whisper-tiny

Whisper tiny

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.

37.8M parameterswhisperApache-2.0automatic-speech-recognitionaudioencoder-decodercross-attentionmultilingual
models/whisper-tiny8 files · 44.9 kB in the registry · GitHub
README.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).

Open on GitHub

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"
}

Open on GitHub

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"
}

Open on GitHub

linnet.toml62 B
toml
[package]
name = "whisper"
version = "0.1.0"
language = "0.1"

Open on GitHub

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"

Open on GitHub

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."
}

Open on GitHub

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)
    }
}

Open on GitHub

model.safetensorsHugging Face ↗