Automatic Speech Recognition
MLX
English
parakeet_tdt_five_value
apple-silicon
speech-to-text
asr
stt
low-bit
quantization-aware-training
on-device
Eval Results
Instructions to use FermionResearch/Phonon-2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use FermionResearch/Phonon-2 with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir Phonon-2 FermionResearch/Phonon-2
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download reference_transformers.py from FermionResearch/Phonon-2: direct link, hf CLI and curl.
- Browser
- Download file 6.83 kB
-
https://huggingface.co/FermionResearch/Phonon-2/resolve/main/reference_transformers.py
- Command line
-
hf download hf://FermionResearch/Phonon-2/reference_transformers.py
-
curl -L -o reference_transformers.py https://huggingface.co/FermionResearch/Phonon-2/resolve/main/reference_transformers.py
6.83 kB
| """Dense `transformers` reference for the Phonon-2 container. | |
| read container -> expand every record to fp32 (five-value: exact {0,+-lo,+-hi}; intN: q*scale; | |
| fp16: as stored) -> stock ParakeetForTDT(config).load_state_dict(strict=True) -> greedy TDT | |
| generate through ParakeetProcessor, the same call the leaderboard scoring uses. | |
| Usage: python reference_transformers.py model.fermion --audio recording.wav | |
| python reference_transformers.py model.fermion --utts rows.json --out receipt.json | |
| Outputs per utterance: transcript (+ token ids) and, for --dump-encoder N utts, the encoder | |
| output in fp32 so the MLX graph can be compared against the torch graph. Pure reference: | |
| no speed claim is made or measured here. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import re | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| HERE = Path(__file__).resolve().parent | |
| sys.path.insert(0, str(HERE)) | |
| from fermion_container import read_container # noqa: E402 | |
| def container_state_dict(container: str): | |
| """HF-named fp32 torch state dict; num_batches_tracked restored to int64.""" | |
| import torch | |
| tensors, index = read_container(container, with_raw=False) | |
| sd = {} | |
| for k, v in tensors.items(): | |
| if k.endswith("num_batches_tracked"): | |
| # the writer stores this int64 BatchNorm counter as fp16 -> inf for counts >= 65,520; | |
| # BatchNorm.eval() never reads it, so a non-finite value becomes 0 | |
| v32 = float(np.asarray(v, dtype=np.float32).reshape(-1)[0]) | |
| sd[k] = torch.tensor(int(v32) if np.isfinite(v32) else 0, dtype=torch.int64) | |
| else: | |
| arr = np.ascontiguousarray(v.astype(np.float32)) | |
| # the writer squeezed the kernel-1 pointwise convs to [O, I]; stock ParakeetForTDT wants [O, I, 1] | |
| if arr.ndim == 2 and re.search(r"\.conv\.pointwise_conv[12]\.weight$", k): | |
| arr = arr[:, :, None] | |
| sd[k] = torch.from_numpy(arr) | |
| return sd, index | |
| def load_model(container: str, base_dir: str, dtype=None): | |
| import torch | |
| from transformers import AutoProcessor, ParakeetForTDT, ParakeetTDTConfig | |
| cfg = ParakeetTDTConfig.from_pretrained(base_dir) | |
| model = ParakeetForTDT(cfg) | |
| sd, index = container_state_dict(container) | |
| missing, unexpected = model.load_state_dict(sd, strict=False) | |
| receipt = {"missing": list(missing), "unexpected": list(unexpected), | |
| "container_records": len(index), "state_dict_keys": len(sd), | |
| "params": sum(p.numel() for p in model.parameters())} | |
| if missing or unexpected: | |
| raise RuntimeError(f"G-LOAD (transformers): missing {list(missing)[:5]} unexpected {list(unexpected)[:5]}") | |
| model.eval() | |
| if dtype is not None: | |
| model = model.to(dtype) | |
| # a model built from its config carries no generation_config; the base repo's is authoritative | |
| from transformers import GenerationConfig | |
| model.generation_config = GenerationConfig.from_pretrained(base_dir) | |
| receipt["generation_config"] = {k: getattr(model.generation_config, k, None) | |
| for k in ("decoder_start_token_id", "bos_token_id", "pad_token_id", "max_symbols_per_step")} | |
| processor = AutoProcessor.from_pretrained(base_dir) | |
| return model, processor, receipt | |
| def main(): | |
| import argparse | |
| import torch | |
| import soundfile as sf | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("container") | |
| ap.add_argument("--base-dir", default="nvidia/parakeet-tdt-0.6b-v3", | |
| help="the base repo (config, processor, generation config); a Hub id or a local directory") | |
| ap.add_argument("--audio", nargs="+", metavar="FILE", | |
| help="audio file(s) to transcribe (any sample rate; resampled to 16 kHz, mixed to mono)") | |
| ap.add_argument("--utts", default=None, help="JSON list of {id, path, ref} rows (the evaluation rows); ignored with --audio") | |
| ap.add_argument("--n", type=int, default=40) | |
| ap.add_argument("--dtype", default="float32", choices=["float32", "bfloat16"]) | |
| ap.add_argument("--dump-encoder", type=int, default=None, | |
| help="dump the fp32 encoder output of the first N utterances (default 4 with --utts, 0 with --audio)") | |
| ap.add_argument("--out", default=None, help="write transcripts (+ token ids) as JSON here; default: print only") | |
| a = ap.parse_args() | |
| if not a.audio and not a.utts: | |
| ap.error("give --audio FILE [FILE ...] or --utts rows.json") | |
| if a.dump_encoder is None: | |
| a.dump_encoder = 0 if a.audio else 4 | |
| torch.set_num_threads(max(1, torch.get_num_threads() // 2)) | |
| dtype = getattr(torch, a.dtype) | |
| t0 = time.perf_counter() | |
| model, processor, rc = load_model(a.container, a.base_dir, dtype=dtype) | |
| rc["load_s"] = round(time.perf_counter() - t0, 1) | |
| rc["torch"] = torch.__version__ | |
| import transformers; rc["transformers"] = transformers.__version__ | |
| rc["dtype"] = a.dtype | |
| if a.audio: | |
| utts = [{"id": Path(f).stem, "path": f, "split": "file", "ref": None} for f in a.audio] | |
| else: | |
| utts = json.load(open(a.utts))[: a.n] | |
| rows = [] | |
| enc_dump = {} | |
| t1 = time.perf_counter() | |
| for i, u in enumerate(utts): | |
| w, sr = sf.read(u["path"], dtype="float32", always_2d=True) | |
| w = w.mean(axis=1) if w.shape[1] > 1 else w[:, 0] | |
| if sr != 16000: | |
| try: | |
| import librosa | |
| w = librosa.resample(w, orig_sr=sr, target_sr=16000).astype(np.float32) | |
| except ImportError: | |
| import torchaudio.functional as taf | |
| w = taf.resample(torch.from_numpy(w), sr, 16000).numpy().astype(np.float32) | |
| sr = 16000 | |
| inp = processor([w], sampling_rate=16000, return_tensors="pt", padding=True) | |
| feats = inp["input_features"].to(dtype) | |
| with torch.no_grad(): | |
| if i < a.dump_encoder: | |
| enc = model.encoder(input_features=feats, attention_mask=inp.get("attention_mask")) | |
| enc_dump[u["id"]] = enc.last_hidden_state[0].float().numpy() | |
| out = model.generate(input_features=feats, attention_mask=inp.get("attention_mask")) | |
| seq = getattr(out, "sequences", out) | |
| text = processor.batch_decode(seq, skip_special_tokens=True)[0].strip() | |
| ids = [int(x) for x in seq[0].tolist()] | |
| rows.append({"id": u["id"], "split": u["split"], "text": text, "token_ids": ids, "ref": u["ref"]}) | |
| print(f"{u['id']:>18} {text[:90]}") | |
| rc["decode_wall_s"] = round(time.perf_counter() - t1, 1) | |
| if a.out: | |
| json.dump({"load": rc, "rows": rows}, open(a.out, "w"), indent=1) | |
| if enc_dump: | |
| np.savez(a.out.replace(".json", "_encoder.npz"), **enc_dump) | |
| print(json.dumps(rc, indent=1)) | |
| if __name__ == "__main__": | |
| main() | |