Automatic Speech Recognition
Transformers
Safetensors
English
voxtral
audio
speculative-decoding
neuron
trainium
distillation
Instructions to use jburtoft/Voxtral-Mini-3B-2507-draft-4layer with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use jburtoft/Voxtral-Mini-3B-2507-draft-4layer with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("automatic-speech-recognition", model="jburtoft/Voxtral-Mini-3B-2507-draft-4layer")# Load model directly from transformers import AutoProcessor, AutoModelForMultimodalLM processor = AutoProcessor.from_pretrained("jburtoft/Voxtral-Mini-3B-2507-draft-4layer") model = AutoModelForMultimodalLM.from_pretrained("jburtoft/Voxtral-Mini-3B-2507-draft-4layer", device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 12,827 Bytes
f9e3832 c734631 f9e3832 c734631 f9e3832 c734631 f9e3832 cbbd993 f9e3832 cbbd993 c734631 cbbd993 f9e3832 c734631 f9e3832 cbbd993 f9e3832 c734631 f9e3832 c734631 f9e3832 c734631 f9e3832 c734631 f9e3832 c734631 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 | # Training scripts
Full reproduction pipeline for training the 4-layer distilled Voxtral draft in this repository. See the [main model card](../README.md#reproducibility) for the end-to-end recipe. This document describes each script individually.
All scripts are self-contained and driven by command-line arguments. Defaults assume `/mnt/data/` for working data and `/mnt/drafts/` for output checkpoints, but every path is overridable via a flag.
## Two training paths
The scripts support two Trainium software stacks:
1. **torch-neuronx XLA (SDK 2.31 DLAMI)** -- the default. Uses `torch_xla` with `xm.mark_step()`. Runs in the standard SDK 2.31 DLAMI's `pytorch_2_9` venv. Faster steady-state per-step, produced the weights in this repo.
2. **PyTorch Native Beta 4 (SageMaker Training DLC)** -- eager mode + `torch.compile(backend='neuron')`. Runs in a separate venv created by `beta35-eval`'s `deploy-beta35.sh --variant beta4` script. Same convergence, slightly slower per-step, simpler port from vanilla PyTorch training code, and the direction AWS is investing in for future Trainium training.
Both paths produce the same trained weights (within run-to-run seed variance). The `_beta4` variants of the sanity and training scripts are the Beta 4 ports.
## Pipeline diagram
```
LibriSpeech .flac -----> precompute_pseudo_labels.py ----> {clip}.pt per clip
| (input_ids,
audio_encoder (CPU) + Voxtral greedy --+ audio_embeds,
teacher_token_ids)
|
Trainium training loop: v
for step in range(total_steps): |
batch = load({clip}.pt files, pad to max_seq_len) <----------+
inputs_embeds = splice audio_embeds into text embeds
logits = student(inputs_embeds)
loss = CE(logits, teacher_token_ids) # teacher forcing
backward + step
|
v
/mnt/drafts/best <---- best checkpoint by
val top-1 agreement
on held-out audio
```
## Scripts
| Script | Purpose | Stack |
|---|---|---|
| `sanity_voxtral_xla.py` | Verify Voxtral loads on trn2 via `torch_xla`. **Expected to fail** on the audio path (documents the trn2 `sort` op limitation). | XLA |
| `sanity_voxtral_xla_v2.py` | Verified working XLA path: audio encoder on CPU, LLM decoder on trn2 via `torch_xla`. | XLA |
| `sanity_voxtral_beta4.py` | Beta 4 equivalent of `sanity_voxtral_xla_v2.py`. Uses `torch.device("neuron:0")` eager mode + `torch.compile(backend='neuron')`. Prints wall time for both eager and compiled variants. | Beta 4 |
| `build_draft.py` | Take a Voxtral checkpoint and prune to N decoder layers. Saves to `--out-dir`. Used at training-start to create the initial student. | Any |
| `build_librispeech_manifest.py` | Build a CSV manifest from a single LibriSpeech split (`dev-clean` etc.). | Any |
| `build_librispeech_combined_manifest.py` | Build a manifest combining multiple LibriSpeech splits, with de-duplication against a previously-precomputed pseudo-label directory. | Any |
| `precompute_pseudo_labels.py` | For each audio clip in a manifest: run Voxtral greedy on CPU, save `(input_ids, audio_embeds, teacher_token_ids)` as `<stem>.pt`. | Any (CPU only) |
| `train_distill.py` (v1) | Simple training loop, no validation. Kept for reference; use v2 instead. | XLA |
| `train_distill_v2.py` (v2) | Training with periodic validation on a held-out set + best-checkpoint tracking. **This is the script that produced this repo's weights.** | XLA |
| `train_distill_v2_beta4.py` | Beta 4 port of v2. Identical arguments + adds `--compile` flag to wrap the LLM decoder in `torch.compile(backend='neuron')`. | Beta 4 |
| `validate_draft.py` | Standalone quality check for a trained draft. Runs greedy generation with the draft alone, then computes teacher-forced per-position top-1 agreement between draft and target. | Any (CPU) |
## Key design choices
### Layer selection: `[0, 10, 20, 29]`
- Includes layer 0 (initial embedding processing) and layer 29 (final projection) for stability.
- Even spacing in between (10, 20) samples different depths of the target's representation.
- 4 layers is ~13% of the target's decoder FLOPs (`4/30`).
- Not systematically swept. Alternatives like `[0, 7, 15, 22, 29]` (5 layers) may give a better acceptance/cost tradeoff and are trivial to try by changing `--keep-layers`.
### Pseudo-labels via greedy teacher
- Runs Voxtral end-to-end on each audio clip with greedy decoding.
- Saves target's transcript token IDs; student learns to reproduce them exactly under teacher forcing (cross-entropy loss).
- Alternatives considered:
- **KL on soft logits**: potentially better quality but requires storing full 131072-vocab distributions per position (~256 MB per 128-token clip, 500+ GB total). Not run.
- **Ground-truth transcripts (LibriSpeech gold)**: risks the student learning phrasings different from Voxtral's, hurting downstream spec-decode acceptance.
- Chose hard-label distillation (teacher's argmax) as the simplest reasonable option.
### Frozen audio path
- Audio encoder + projector + `embed_tokens` + final `norm` + `lm_head` are byte-identical between student and target -- no reason to train them.
- Only the 4 decoder layers get gradients (~428M trainable of the ~3B target total).
### Batch size 2
- Chose conservatively for a first run. Larger batches (4, 8) are likely feasible and would give more stable gradients. Untested; open for future work.
- Batch size 1 also works but shows higher loss variance.
### Cosine LR schedule
- Standard for short distillation runs.
- 100-step linear warmup + cosine decay to 10% of peak LR over remaining steps.
- Peak LR: 1e-4. Reasonable for LLM fine-tuning; not swept.
## Trainium-specific workaround (XLA path only)
**Root cause -- XLA path only.** On the `torch-neuronx` XLA path, Voxtral's end-to-end forward compile fails with `[NCC_EVRF029] Operation sort is not supported on trn2`. The op being sorted is `s32[1176576]` -- a 1D int32 index tensor of size `1 * 383 * 3072` -- which is the flattened form of Voxtral's `inputs_embeds` `[1, 383, 3072]`. The `sort` comes from `torch_xla`'s lowering of `masked_scatter` in `VoxtralModel.forward`, specifically the splicing of audio embeddings into `inputs_embeds` at the audio-token positions. **It is not the audio encoder itself.** The audio encoder (`audio_tower` + `multi_modal_projector`) compiles cleanly on Neuron in both XLA and Beta 4 eager modes.
**Verified 2026-07-31**:
- XLA `model.model.audio_tower(input_features)` -> compiles + returns `[1, 1500, 1280]` correctly.
- XLA `inputs_embeds.masked_scatter(mask, audio_embeds)` on `[1, 383, 3072]` -> fails with `sort` error.
- Beta 4 eager `inputs_embeds.masked_scatter(...)` on the same tensor -> works.
- Beta 4 eager `model(**inputs)` end-to-end on `neuron:0` -> works, ~1150 ms per forward.
**What the workaround does (used by every script here):** precompute audio embeddings on CPU once per training clip, splice them into `inputs_embeds` also on CPU in the DataLoader collate function, and only send the audio-conditioned `inputs_embeds` to Neuron. The audio encoder + projector never runs on Neuron in the training loop.
1. `precompute_pseudo_labels.py` runs `model.model.get_audio_features(input_features)` on CPU. This produces `audio_embeds` of shape `[375, 3072]` per clip.
2. `train_distill_v2.py` (and its Beta 4 variant) splices `audio_embeds` into `inputs_embeds` at the audio-token positions in the collate function (on CPU, per batch).
3. The Trainium graph only sees `inputs_embeds` (already audio-conditioned), not `input_features`. No audio encoder involved.
Since the audio path is frozen (never trained), this precomputation only runs once per training clip -- it's cheaper than re-running the frozen encoder every training step regardless of whether it would compile.
**Beta 4 note:** you could rewrite `train_distill_v2_beta4.py` to compute `audio_embeds` on-device inside the training loop and skip the precompute step. It works (see the sanity test), but it's slower per step (~1150 ms including audio encoder) than the precompute approach (~200-590 ms per step, audio-encoder cost paid once offline). The precompute approach is retained in `train_distill_v2_beta4.py` for parity with the XLA script and for training efficiency.
## Observations from the reference run (XLA path, 3000 steps)
Final training summary (2000 LibriSpeech clips, batch=2):
```
val step 200: top-1 agreement = 12.82%
val step 400: top-1 agreement = 19.23%
val step 600: top-1 agreement = 25.05%
val step 800: top-1 agreement = 31.36%
val step 1000: top-1 agreement = 36.59%
val step 1200: top-1 agreement = 41.62%
val step 1400: top-1 agreement = 49.01%
val step 1600: top-1 agreement = 51.78%
val step 1800: top-1 agreement = 53.75%
val step 2000: top-1 agreement = 57.40%
val step 2200: top-1 agreement = 62.43%
val step 2400: top-1 agreement = 64.40%
val step 2600: top-1 agreement = 61.93%
val step 2800: top-1 agreement = 66.67% <-- best
val step 3000: top-1 agreement = 65.78%
```
- Curve was still climbing at step 2400; plateaued at 62-67% for the last 600 steps. Training data appears to be the bottleneck at this scale.
- Data-scaling law observed in earlier runs: 4x more training data (500 -> 2000 clips) produced 1.7x higher final val top-1 (39% -> 67%). Scaling to 10k+ LibriSpeech clips (e.g., `train-clean-100` = ~28k clips) should push val top-1 into the 75-85% range.
- Compile bursts happen periodically as new sequence shapes are encountered by the XLA HLO compiler. After the first ~60 s compile the training loop runs at ~130 ms/step steady-state. Save-checkpoint triggers a flush every N steps but is quick.
## Beta 4 comparison (500 steps, `--compile`)
Same hyperparameters, same data, same seed. 500-step run:
| Step | XLA v2 | Beta 4 compiled |
|---:|---:|---:|
| 100 | 8.28% | 6.61% |
| 200 | 9.17% | 9.76% |
| 300 | 10.06% | 10.16% |
| 400 | 14.40% | 15.88% |
| 500 | 18.64% | 18.34% |
**Convergence is essentially identical**. The two paths produce interchangeable weights.
Per-step wall clock (steady state, batch=2, max_seq_len=512):
| Path | Compile time | Steady-state step |
|---|---:|---:|
| `torch-neuronx` XLA | ~60 s | ~130 ms |
| Beta 4 `--compile` | ~49 s | ~200-590 ms (some recompiles on val batches) |
| Beta 4 eager (no `--compile`) | ~107 s | ~3700 ms (28x slower) |
**XLA is currently the fastest path for this scale.** Beta 4 with `--compile` is 1.5-2x slower total wall time. Beta 4 eager is not competitive.
**Recommendation**: use `train_distill_v2.py` (XLA) unless your environment mandates the Beta 4 stack (some SageMaker training clusters, TorchTitan integration, or if you want to try `torch.compile` fusion on more advanced setups). The Beta 4 script is a drop-in replacement with the same arguments plus `--compile`.
## Known gotchas
1. **First iteration compile is slow** (both paths): 60-150 s per new tensor shape encountered. Subsequent iterations hit the cache. Keep `max_seq_len` constant across all batches for best cache hit rate.
2. **Checkpoint save is expensive**: ~30-40 s per save (writes 3.6 GB safetensors + re-loads the full target for state-dict merging via `load_state_dict(..., strict=False)`). Use `--save-every 500` or larger to minimize interruptions.
3. **CPU pseudo-labeling is slow but embarrassingly parallel**: ~5-10 s per clip on a single CPU. 2000 clips takes ~2 hours on one instance. If precompute is the bottleneck, run it on a large-core CPU instance (c7i.16xlarge is much faster than trn2's ~48 CPUs).
4. **Voxtral's audio path is byte-identical between student and target**, so the `.pt` files produced by `precompute_pseudo_labels.py` are reusable across draft architectures. If you train a 6-layer variant later, the precompute step doesn't need to re-run.
5. **Beta 4-specific**: the first log-print says step time is very high (e.g., "step 1 | step time 24768ms"). That's the cumulative-average step time carrying the compile cost. Look at the deltas between successive `elapsed` timestamps to see the true per-step wall time.
6. **Beta 4-specific**: `torch.compile(dynamic=False)` locks shapes. If your batch shapes vary (e.g., variable-length audio), you'll trigger recompiles and lose the fusion advantage. Pad every batch to a single fixed shape.
|