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] hf download FermionResearch/Phonon-2 --local-dir Phonon-2
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
File size: 6,828 Bytes
b816c2e 88baeef b816c2e 88baeef b816c2e 88baeef e4d360d 88baeef | 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 | """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()
|