gigaam-v3-e2e-rnnt-coreml / example_infer.py
smkrv's picture
self-contained: tokenizer.model, full upstream config yaml, runnable example_infer.py, card update
9782676 verified
Raw
History Blame Contribute Delete
2.77 kB
"""Minimal reference inference for GigaAM-v3 e2e RNNT Core ML packages.
Usage: python example_infer.py audio.wav
Audio: any format ffmpeg reads; downmixed/resampled to 16 kHz mono; first 30 s.
Deps: pip install coremltools torch torchaudio numpy; ffmpeg binary on PATH
(audio is loaded via an ffmpeg pipe, mirroring gigaam.preprocess.load_audio -
torchaudio is used only for MelSpectrogram, so torchcodec is NOT required).
"""
import json
import pathlib
import subprocess
import sys
import coremltools as ct
import numpy as np
import torch
import torchaudio
HERE = pathlib.Path(__file__).parent
CU = ct.ComputeUnit.CPU_AND_GPU # token-exact vs the PyTorch reference
BLANK = 1024
SR = 16000
WINDOW = 30 * SR
MAX_SYMBOLS = 10
enc = ct.models.MLModel(str(HERE / "GigaAMv3Encoder.mlpackage"), compute_units=CU)
dec = ct.models.MLModel(str(HERE / "GigaAMv3DecoderStep.mlpackage"), compute_units=CU)
joint = ct.models.MLModel(str(HERE / "GigaAMv3JointStep.mlpackage"), compute_units=CU)
pieces = json.loads((HERE / "tokens.json").read_text())
mel = torchaudio.transforms.MelSpectrogram(
sample_rate=SR, n_mels=64, win_length=320, hop_length=160, n_fft=320,
center=False, mel_scale="htk", norm=None,
)
raw = subprocess.run(
["ffmpeg", "-nostdin", "-i", sys.argv[1], "-f", "s16le", "-ac", "1",
"-acodec", "pcm_s16le", "-ar", str(SR), "-"],
capture_output=True, check=True,
).stdout
wav = torch.frombuffer(bytearray(raw), dtype=torch.int16).float() / 32768.0
n = min(wav.shape[-1], WINDOW)
padded = torch.zeros(1, WINDOW)
padded[0, :n] = wav[:n]
feats = mel(padded).clamp(1e-9, 1e9).log().numpy().astype(np.float32)
feat_len = (n - 320) // 160 + 1
out = enc.predict({"features": feats, "length": np.array([feat_len], dtype=np.int32)})
encoded = out["encoded"]
enc_len = int(np.array(out["encoded_len"]).reshape(-1)[0])
ids, frames = [], []
h = np.zeros((1, 1, 320), dtype=np.float32)
c = h.copy()
last = BLANK # blank embedding row is zeros == the reference fresh start
for t in range(enc_len):
enc_t = encoded[:, :, t].astype(np.float32)
for _ in range(MAX_SYMBOLS):
d = dec.predict({"token": np.array([[last]], dtype=np.int32), "h_in": h, "c_in": c})
logits = joint.predict(
{"enc_t": enc_t, "dec_t": d["dec_out"].astype(np.float32)}
)["logits"]
k = int(np.argmax(logits))
if k == BLANK:
break
ids.append(k)
frames.append(t)
h, c = d["h_out"].astype(np.float32), d["c_out"].astype(np.float32)
last = k
text = "".join(pieces[i] for i in ids).replace("▁", " ").strip()
print(text)
if frames:
# 4x subsampling, 10 ms hop -> 40 ms per encoder frame
print(f"[{frames[0] * 0.04:.2f}s .. {frames[-1] * 0.04:.2f}s, {len(ids)} tokens]")