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()