Phonon-2 / reference_transformers.py
FermionResearch's picture
Phonon-2: reference script help text
e4d360d verified
Raw History Blame Contribute Delete
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()