"""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]")