Developer reference for model.py. For export/usage instructions see
README.md.
The model is written directly with ExecuTorch custom ops rather than using
source transformations. The patterns come from examples/models/llama/,
extension/llm/export/builder.py, and optimum-executorch.
Audio waveform @ 16kHz
-> Mel spectrogram (128 bins, hop=160, window=400)
(B, 128, T_mel)
-> CausalConv1d (128 -> 1280, k=3, s=1) + GELU
-> CausalConv1d (1280 -> 1280, k=3, s=2) + GELU
(B, 1280, T_mel//2) -> transpose -> (B, T_mel//2, 1280)
-> 32x CausalEncoderLayer (RMSNorm -> RoPE attention -> RMSNorm -> SwiGLU)
-> RMSNorm
(B, T_mel//2, 1280)
-> Reshape: concat downsample_factor=4 consecutive frames
(B, T_mel//8, 5120)
-> AudioLanguageAdapter: Linear(5120, 3072) -> GELU -> Linear(3072, 3072)
(B, T_audio, 3072) = audio_embeds
audio_embeds + token_embedding(prev_token) = combined_embeds
combined_embeds
-> 26x MistralDecoderLayer (RMSNorm -> GQA attention -> adaptive RMSNorm(t_cond) -> SwiGLU)
-> RMSNorm -> Linear(3072, 131072)
(B, seq_len, 131072) = logits
The model exports three methods (offline mode):
| Method | Input | Output |
|---|---|---|
audio_encoder |
mel spectrogram (1, 128, T_mel) |
audio embeddings (1, T_mel//8, 3072) |
text_decoder |
embeddings (1, seq_len, 3072) + positions (seq_len,) |
logits (1, seq_len, 131072) |
token_embedding |
token IDs (1, seq_len) |
embeddings (1, seq_len, 3072) |
With --streaming, audio_encoder is replaced by encode_audio_chunk
which takes a mel chunk (1, 128, 8) + encoder positions (4,) and
returns audio embeddings (1, 1, 3072). Conv states are maintained as
internal buffers.
Audio and text embeddings are summed at each position (not concatenated or masked-scatter like the original non-realtime Voxtral).
| Parameter | Encoder | LM Decoder |
|---|---|---|
| dim | 1280 | 3072 |
| layers | 32 | 26 |
| heads | 32 | 32 (8 KV, GQA 4:1) |
| head_dim | 64 | 128 |
| hidden_dim (FFN) | 5120 | 9216 |
| rope_theta | 1,000,000 | 1,000,000 |
| biases (attn) | wq, wv, wo yes; wk no | none |
| biases (FFN) | w2 yes; w1, w3 no | none |
| vocab_size | — | 131,072 |
| total params | ~1B | ~3.4B |
| Audio Parameter | Value |
|---|---|
| sample_rate | 16,000 Hz |
| num_mel_bins | 128 |
| hop_length | 160 |
| window_size | 400 |
| downsample_factor | 4 |
| frame_rate | 12.5 fps |
Decoder KV cache depends on mode:
- Offline: flat buffer sized by
max_seq_len(default 4096). 26 layers × 2 × 4096 × 8 × 128 × bytes_per_elem. fp32: ≈ 832 MB, bf16: ≈ 416 MB. - Streaming: ring buffer sized to 2×
sliding_window(default 8192 → 16384 slots;--sliding-window 2048→ 4096 slots). 26 layers × 2 × 2×sliding_window × 8 × 128 × bytes_per_elem. sw=8192 fp32: ≈ 3.3 GB, bf16: ≈ 1.7 GB. sw=2048 fp32: ≈ 832 MB, bf16: ≈ 416 MB.
Encoder KV caches (streaming only): 32 layers × 2 × 1500 × 32 × 64 × bytes_per_elem. fp32: ≈ 786 MB, bf16: ≈ 393 MB.
Runtime memory = model weights (from .pte) + KV caches + working
memory. Weight sizes depend on quantization: ~16 GB (fp32), ~8 GB
(bf16), ~4 GB (8w), ~2 GB (4w/8da4w). Metal and CUDA backends are recommended to use
bf16 (--dtype bf16) when quantization is enabled.
VoxtralRealtimeModel
encoder: CausalWhisperEncoder
conv_layers: [CausalConv1d, CausalConv1d]
layers: 32x CausalEncoderLayer
attention_norm: RMSNorm
attention: EncoderAttention (wq/wk/wv/wo, F.scaled_dot_product_attention)
ffn_norm: RMSNorm
feed_forward: EncoderSwiGLU (w1/w2/w3)
norm: RMSNorm
adapter: AudioLanguageAdapter (w_in/w_out)
decoder: MistralDecoder
tok_embeddings: Embedding
layers: 26x MistralDecoderLayer
attention_norm: RMSNorm
attention: LMAttention
wq/wk/wv/wo: Linear (no bias)
kv_cache: streaming: RingKVCache/StandardRingKVCache; offline: KVCache/StaticKVCache
sdpa: SDPA (XNNPACK) or MetalSDPA (Metal) or StandardSDPA (CUDA)
ffn_norm: RMSNorm
ada_rms_norm_t_cond: Sequential(Linear, GELU, Linear)
feed_forward: LMMLP (w1/w2/w3)
norm: RMSNorm
output: Linear (tied to tok_embeddings)
StreamingAudioEncoderExport
conv1: nn.Conv1d (shared from encoder.conv_layers[0].conv)
conv2: nn.Conv1d (shared from encoder.conv_layers[1].conv)
layers: 32x CausalEncoderLayer (shared from encoder.layers)
enc_norm: RMSNorm (shared from encoder.norm)
adapter: AudioLanguageAdapter (shared from model.adapter)
kv_caches: 32x RingKVCache (XNNPACK) or StandardRingKVCache (Metal/CUDA)
sdpa: SDPA (XNNPACK) or MetalSDPA (Metal) or StandardSDPA (CUDA)
inv_freq: RoPE inverse frequencies (owned, on-the-fly computation)
The offline encoder (CausalWhisperEncoder) processes the full mel
spectrogram at once. No KV cache, no GQA (n_heads == n_kv_heads).
EncoderAttention uses F.scaled_dot_product_attention with
is_causal=True, transposing to [B, H, T, D] internally. No custom
ops needed — works on all backends (XNNPACK, Metal, CUDA, Portable).
The offline encoder uses full causal attention (no sliding window).
The model's params.json specifies sliding_window: 750 but this is
only enforced in the streaming encoder (via KV cache). For audio shorter
than 750 encoder frames (~15s), full causal is equivalent.
The text decoder (MistralDecoder) is a 26-layer Mistral decoder with
GQA (32 query heads, 8 KV heads). Backend selection is controlled by the
backend config field, passed through from the export script's --backend
flag (e.g., "xnnpack", "metal", "cuda", "portable").
The decoder KV cache depends on the export mode:
Streaming (--streaming): Ring buffer KV cache for unlimited
duration. The model's params.json specifies sliding_window: 8192
for the decoder (overridable via --sliding-window). Each query attends
to only the last sliding_window positions; old entries are overwritten
when the buffer wraps. Position tracking is analytic (no mutable state).
Sliding window masks are computed each step via create_causal_mask.
- XNNPACK/Portable:
RingKVCachewith[B, S, H, D]layout, usingtorch.ops.llama.update_cache_with_indicesfor scatter writes. - Metal/CUDA:
StandardRingKVCachewith[B, H, S, D]layout, usingindex_copy_on dim=2 with wrapped indices.
Offline (default): Flat KV cache bounded by max_seq_len (default
4096). Full causal attention — each query attends to all prior positions.
- XNNPACK/Portable:
KVCachewith[B, S, H, D]layout, usingtorch.ops.llama.update_cache. - Metal/CUDA:
StaticKVCachewith[B, H, S, D]layout, usingindex_copy_.
SDPA is its own module (not inline code), making it swappable for
backend-specific implementations.
XNNPACK/Portable: SDPA uses torch.ops.llama.custom_sdpa.
In streaming mode, receives a sliding window mask from the ring cache.
In offline mode, uses is_causal=True with no explicit mask.
Handles GQA expansion internally and upcasts to float32.
Metal: MetalSDPA uses torch.ops.aten._scaled_dot_product_attention_math_for_mps
which handles GQA natively (the kernel infers the group ratio from differing
Q vs K/V head counts), avoiding the memory bandwidth overhead of
repeat_interleave. Uses explicit additive attention masks
that must match the Q/K/V dtype (the kernel reads masks as device T*).
Both streaming and offline use [B, H, S, D] KV layout
(StandardRingKVCache and StaticKVCache share this layout), so
transpose_kv=False in all cases.
CUDA: StandardSDPA uses F.scaled_dot_product_attention with
enable_gqa=True. Uses boolean attention masks (True=attend,
False=masked) as required by the Triton SDPA kernel. Same
[B, H, S, D] KV layout as Metal.
Q/K/V projections produce [B, T, H, D] via .view(). RoPE operates
on [B, T, H, D].
XNNPACK/Portable: Both KVCache (offline) and RingKVCache
(streaming) use [B, S, H, D]. SDPA (custom_sdpa) receives Q and
KV cache in this layout — no transpose(1, 2) in the attention hot path.
Metal/CUDA: Both StandardRingKVCache (streaming) and StaticKVCache
(offline) use [B, H, S, D] layout. MetalSDPA/StandardSDPA only
transpose Q from [B, T, H, D] to [B, H, T, D] — KV is already in
the expected layout.
RoPE frequencies are computed on-the-fly using stored inv_freq
(same pattern as the streaming encoder), enabling unlimited position
indices without a precomputed table bound.
Each decoder layer has a time-conditioned FFN norm unique to this model. After the standard RMSNorm on the FFN input, a learned scale is applied:
scale = 1 + Sequential(Linear(3072->32), GELU, Linear(32->3072))(t_cond)
ffn_input = rms_norm(x) * scale
The t_cond is a sinusoidal embedding of n_delay_tokens (default 6 = 480ms),
precomputed once and passed to each decoder layer as a constant.
The ada_rms_norm_t_cond modules add ~5.1M parameters across 26
layers (26 × (3072×32 + 32×3072) = 26 × 196,608), quantized by
--qlinear.
For streaming/live transcription, StreamingAudioEncoderExport processes
audio incrementally (8 mel frames = 80ms per step) instead of the full
mel at once. It shares all weights with the offline encoder but uses a
different forward path:
mel_chunk (1, 128, 8) + enc_input_pos (4,)
conv1_state (1, 128, 2) and conv2_state (1, 1280, 2) are internal buffers
-> cat(state, chunk) -> raw Conv1d (no CausalConv1d padding) -> GELU
-> cat(state, conv1_out) -> raw Conv1d -> GELU
(1, 1280, 4) -> transpose -> (1, 4, 1280)
-> 32x streaming encoder layer (ring KV cache + SDPA)
-> RMSNorm
(1, 4, 1280)
-> Reshape downsample (1, 1, 5120) -> Adapter (1, 1, 3072)
-> audio_embeds (1, 1, 3072)
XNNPACK/Portable: Uses RingKVCache (update_cache_with_indices
custom op) and SDPA (custom_sdpa).
Metal: Uses StandardRingKVCache (index_copy_-based ring
buffer) and MetalSDPA (native MPS SDPA kernel).
Masks are created in the model dtype to match the kernel's device T* expectation.
CUDA: Uses StandardRingKVCache and StandardSDPA
(F.scaled_dot_product_attention with explicit
sliding window masks).
Each 80ms step produces one audio embedding (1, 1, 3072). The
runner (StreamingSession::decode_step) then:
- Looks up the embedding for the previous token via
token_embedding - Sums audio + token embeddings element-wise (same as offline mode)
- Feeds the combined embedding to
text_decoderat the current position - Samples one token from the output logits
After audio ends, flush() pads the unfinished tail with silence and keeps
running the same audio-conditioned streaming path until the final partial
step and transcription delay are drained.
The causal convolutions need left context across chunk boundaries. Instead of zero-padding (offline) or recompute-with-overlap (vLLM), explicit conv state carries the tail of the previous chunk:
- Conv1 (kernel=3, stride=1): state = last 2 mel frames from previous
chunk.
cat(state, chunk)→ (1, 128, 10) → Conv1d → (1, 1280, 8). - Conv2 (kernel=3, stride=2): state = last 2 conv1 GELU output frames.
cat(state, conv1_out)→ (1, 1280, 10) → Conv1d → (1, 1280, 4).
The raw nn.Conv1d is called directly (bypassing CausalConv1d.forward
which would zero-pad). This produces identical results to the offline
encoder — verified to within fp32 precision (max diff < 2e-5).
Each of the 32 encoder transformer layers gets its own ring buffer KV
cache (RingKVCache for XNNPACK/Portable, StandardRingKVCache
for Metal/CUDA) that overwrites old entries when the window is exceeded,
enabling streaming of arbitrary length audio.
- Cache shape:
(1, 2*max_enc_len, 32, 64)per layer. The buffer is 2x the window size because writes happen before attention. With a 1x buffer (size = window), writingseq_lennew entries evicts that many old ones — but the current queries still need those old entries. Example withwindow=4, seq_len=4, start_pos=5: a 1x buffer would overwrite positions 1-4 with 5-8, so query at position 5 can only attend to itself instead of positions 2-4. A 2x buffer (size 8) keeps positions 1-4 alive alongside 5-8, giving query 5 full access to its window. - Default
max_enc_len=750(matching the model's trained sliding window). Configurable via--max-enc-len. - Memory: 32 layers × 2 × 1500 × 32 × 64 × bytes_per_elem ≈ 786 MB (fp32), 393 MB (bf16)
- Duration: unlimited (ring buffer overwrites old entries, RoPE computed on-the-fly)
Naming note: max_enc_len in StreamingAudioEncoderExport (default
750, the --max-enc-len CLI flag) is the sliding window size for the
ring buffer. This is unrelated to max_enc_len=16384 in
CausalWhisperEncoder.__init__, which is the RoPE frequency table size
for the offline encoder.
XNNPACK/Portable: Cache writes use torch.ops.llama.update_cache_with_indices
(a custom op that scatter-writes via an indices tensor). Write indices are
computed analytically: (arange(seq_len) + start_pos) % buf_size.
Metal/CUDA: Cache writes use index_copy_ with wrapped indices
(input_pos % buf_size).
No mutable position state is needed in either variant.
Position tracking is analytic — no mutable state buffer. For buffer
slot j after total_written frames have been stored:
abs_pos[j] = j + ((total_written - 1 - j) // buf_size) * buf_size
For example, with buf_size=8 after total_written=10:
- Slot 0:
0 + ((9 - 0) // 8) * 8 = 0 + 8 = 8(wrapped) - Slot 3:
3 + ((9 - 3) // 8) * 8 = 3 + 0 = 3(not yet overwritten)
Negative results indicate unwritten slots. The sliding window mask is computed from these positions each step:
valid = (cache_pos >= 0) & (delta >= 0) & (delta < window_size)
# Metal: float additive mask
mask = torch.where(valid, 0.0, float("-inf"))
# CUDA: boolean mask (bool_mask=True returns valid directly)
mask = validThe mask is identical for all 32 layers (same input_pos), so it
is computed once in forward() and reused.
The streaming preprocessor (WhisperAudioProcessor(streaming=True))
computes mel without 30-second chunk padding. To match offline mel values
at chunk boundaries, the C++ runner uses overlapping audio windows:
- Left overlap: 320 samples (2 × hop_length, ≥ n_fft/2 = 200)
- Right look-ahead: 40 samples (2.5ms, matches vLLM's
streaming_look_ahead_ms) - Total window: 320 + 1280 + 40 = 1640 samples → 10 mel frames
- Frame extraction: skip first 2 frames (overlap region), take frames 2–9 (the 8 that align with offline mel frame positions)
For the first step, the left overlap is zero-padded (matching the
offline encoder's center=True STFT edge behavior). The 2.5ms
look-ahead introduces negligible latency.
q_r, q_i = q.float().reshape(q.shape[:-1] + (-1, 2)).unbind(-1)reshape+unbindinstead of stride-2 slicing (x[..., ::2]) — avoids strided access patterns that produce complex index expressions during export..float()upcast before rotation,.type_as()downcast after — prevents precision loss in fp16/bf16 inference.
Uses F.rms_norm(x, (self.dim,), self.weight, self.eps) with a stored
self.dim attribute for compatibility with Llama's
replace_rms_norm_with_native_rms_norm() source transformation.
Quantization is applied per-component after wrapping (following the Parakeet pattern), allowing different configs for encoder vs decoder:
# XNNPACK/Portable
--qlinear-encoder 8w # encoder linear layers
--qlinear 8da4w # decoder linear layers
--qembedding 8w # embedding layer
# Metal (use --dtype bf16 for reduced memory and improved throughput)
--qlinear-encoder fpa4w # encoder linear layers
--qlinear fpa4w # decoder linear layers
# CUDA
--qlinear-encoder 4w --qlinear-encoder-packing-format tile_packed_to_4d
--qlinear 4w --qlinear-packing-format tile_packed_to_4dThe streaming encoder references the same module objects that
quantize_model_() mutates in-place, so quantized weights are
used transparently. Conv1d layers are not quantized (not targeted
by quantize_model_). KV caches and SDPA have no trainable weights.
Metal backend uses fpa4w (floating-point activation, 4-bit weight)
quantization from TorchAO's experimental MPS ops (UIntxWeightOnlyConfig
with HQQ-based parameter selection). See export_voxtral_rt.py for the
exact configuration.
Each exported method corresponds to a thin wrapper class:
AudioEncoderExport, TextDecoderExport, and TokenEmbeddingExport
(defined in export_voxtral_rt.py). With --streaming,
AudioEncoderExport is replaced by StreamingAudioEncoderExport
(defined in model.py since it owns the ring KV caches and conv states).
All exports use torch.export.export(..., strict=True), matching the
Llama builder and optimum-executorch.
strict=True is required because the model uses .item() to extract
scalar positions from input tensors (for update_cache, custom_sdpa,
and ring buffer index computation). With strict=True, .item() produces
an unbacked SymInt — a symbolic integer that remains dynamic at runtime.
With strict=False, .item() returns the concrete sample value which
gets baked into the graph as a constant, making all cache positions
and attention masks static (the model has no temporal memory).
Each .item() call is guarded with torch._check_is_size(start_pos)
(non-negative constraint) and optionally torch._check(start_pos < max)
(upper bound for bounded caches like the decoder KV cache). The encoder
ring buffer has no upper bound since positions are unlimited.
Mistral format: params.json + consolidated.safetensors (bf16, 8.3 GB).
Tokenizer: Mistral Tekken format (tekken.json, 131K vocab).
load_model() uses the Llama pattern to halve peak memory (~17 GB instead
of ~34 GB for the full-size model):
- Meta device construction —
with torch.device("meta"):builds the model with zero-storage parameter tensors (shape/dtype metadata only). - safetensors lazy access —
safe_openloads tensors on demand, cast to the configured dtype (--dtype, default fp32; bf16 recommended for Metal and CUDA with quantization). assign=Truestate dict loading — replaces meta tensors by reference instead of copying into pre-allocated storage. No duplication.- Post-load fixups — re-tie
output.weight = tok_embeddings.weight(broken by assign), materialize remaining meta buffers (KV caches as zeros), recompute RoPE frequency tables.
| Checkpoint prefix | Model prefix |
|---|---|
mm_streams_embeddings.embedding_module.whisper_encoder.conv_layers.* |
encoder.conv_layers.* |
mm_streams_embeddings.embedding_module.whisper_encoder.transformer.layers.* |
encoder.layers.* |
mm_streams_embeddings.embedding_module.whisper_encoder.transformer.norm.* |
encoder.norm.* |
mm_streams_embeddings.embedding_module.audio_language_projection.{0,2}.weight |
adapter.w_{in,out}.weight |
mm_streams_embeddings.embedding_module.tok_embeddings.weight |
decoder.tok_embeddings.weight |
layers.* |
decoder.layers.* |
norm.weight |
decoder.norm.weight |
Weights are cast to the configured dtype during loading. decoder.output.weight is
not in the checkpoint — it is created by tying to
decoder.tok_embeddings.weight in VoxtralRealtimeModel.__init__.
During export with quantization, the tie is broken (the if args.qlinear or args.qembedding block in export_voxtral_rt.py clones the weight)
so embedding and output linear get separate quantization configs.
KV cache and RoPE frequency buffers are runtime-initialized.
- Source model: mistralai/Voxtral-Mini-4B-Realtime-2602
- Reference implementation: vLLM voxtral_realtime.py
Upstream ExecuTorch patterns:
- KV cache:
examples/models/llama/source_transformation/custom_kv_cache.py(CustomKVCache) - SDPA:
examples/models/llama/source_transformation/sdpa.py(SDPACustom) - RoPE:
examples/models/llama/rope.py(apply_rotary_emb) - Model loading:
examples/models/llama/model.py(load_model) - Export builder:
extension/llm/export/builder.py