| """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 |
| 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 |
| 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: |
| |
| print(f"[{frames[0] * 0.04:.2f}s .. {frames[-1] * 0.04:.2f}s, {len(ids)} tokens]") |
|
|