Views
No views yet
facebook/omniASR-CTC-300M.
The top 3 of 24 Transformer encoder layers are removed, leaving a 21-layer encoder.
The final LayerNorm and the CTC projection head are kept unchanged.| Original | Pruned (this model) | |
|---|---|---|
| Encoder layers | 24 | 21 |
| Parameters | 325.5 M | 287.7 M (−37.8 M, −11.6 %) |
| Encoder compute | 100 % | −12.5 % |
| Exit layer | L17 | L18 | L20 | L23 (full) |
|---|---|---|---|---|
| WER | 45 % | 20 % | 20 % | 20 % |



fairseq2 and
omnilingual-asr.1import dataclasses, torch, torchaudio
2from huggingface_hub import hf_hub_download
3from fairseq2 import init_fairseq2
4from fairseq2.nn import BatchLayout
5from fairseq2.models.wav2vec2.asr import create_wav2vec2_asr_model
6from fairseq2.data.tokenizers.sentencepiece import load_sentencepiece_model
7from omnilingual_asr.models.wav2vec2_asr.config import get_config, Wav2Vec2AsrConfig
8
9ctx = init_fairseq2()
10
11# Build a 21-layer encoder config from the base "300m" arch
12cfg = get_config(ctx, Wav2Vec2AsrConfig, "300m")
13cfg = dataclasses.replace(cfg, encoder_config=dataclasses.replace(
14 cfg.encoder_config, num_encoder_layers=21, layer_drop_p=0.0))
15model = create_wav2vec2_asr_model(cfg).eval()
16
17# Load the pruned weights
18ckpt = hf_hub_download("ChipCracker/omniASR-CTC-300M-pruned-21L",
19 "omniASR-CTC-300M-pruned21L.pt")
20model.load_state_dict(torch.load(ckpt, map_location="cpu")["model"])
21
22# Char tokenizer (CTC blank index = 0)
23tok_path = hf_hub_download("ChipCracker/omniASR-CTC-300M-pruned-21L",
24 "omniASR_tokenizer.model")
25sp = load_sentencepiece_model(tok_path)
26
27# Transcribe a 16 kHz mono wav
28wav, sr = torchaudio.load("audio.wav")
29audio = wav.mean(0) if wav.dim() > 1 else wav.squeeze(0)
30if sr != 16000:
31 audio = torchaudio.transforms.Resample(sr, 16000)(audio)
32
33x = audio.view(1, -1)
34with torch.inference_mode():
35 logits, _ = model(x, BatchLayout.of(x))
36pred = logits[0].argmax(-1)
37keep = torch.ones_like(pred, dtype=torch.bool)
38keep[1:] = pred[1:] != pred[:-1] # CTC collapse
39ids = pred[keep]
40ids = ids[ids != 0] # drop blank
41print("".join(sp.index_to_token(int(i)) for i in ids))1# 1. load the full 24-layer model, drop encoder.layers.{21,22,23}
2# 2. build a 21-layer config, load the remaining weights strict
3# 3. verify per-layer logit-lens WER/CER (single + FLEURS multilingual)omni-viz experiment
(pruning_analysis.py, pruning_fleurs.py, prune_model.py).facebook/omniASR-CTC-300M
(Omnilingual ASR, Meta AI). Please cite the original work:1@misc{omnilingual_asr_2025,
2 title = {Omnilingual ASR},
3 author = {Meta AI},
4 year = {2025},
5 url = {https://huggingface.co/facebook/omniASR-CTC-300M}
6}