"""TDT probe — can timestamps_asr_model handle audio standalone, including >40 s? Background: EncDecMultiTaskModel.timestamps_asr_model is a Token-Duration-Transducer (RNN-T variant) with hybrid CTC head, FastConformer encoder enc_hidden=1024. Bundled inside Canary 1B v2 specifically to provide forced-alignment timestamps for the AED head, but it is itself a complete ASR model. Open question: does the relative positional encoding lift the 40 s wall that AED hits, or does TDT inherit the same training-time max_duration limit? Test set: - 20 s : sanity check, well within both models - 80 s : full source — Canary path1 truncates here, baseline expectation - 160 s : concat(80, 80), no re-encode — pushes 4x past TDT training horizon Per case: - TDT.transcribe(timestamps=False) - TDT.transcribe(timestamps=True) — also verify timestamp shape - Canary path1 — reference, will truncate >40 s """ import json import os import subprocess import tempfile import time import wave import torch from nemo.collections.asr.models import EncDecMultiTaskModel AUDIO_DIR = "/sweep/threshold_audios" RESULTS_PATH = "/sweep/probe_tdt.jsonl" LOG_PATH = "/sweep/probe_tdt.log" def log(msg: str) -> None: line = f"{time.strftime('%H:%M:%S')} {msg}" print(line, flush=True) with open(LOG_PATH, "a") as f: f.write(line + "\n") def transcode(src: str, duration_secs: float | None = None) -> str: dst = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) dst.close() cmd = ["ffmpeg", "-y", "-loglevel", "error", "-i", src, "-ac", "1", "-ar", "16000"] if duration_secs is not None: cmd += ["-t", str(duration_secs)] cmd += ["-f", "wav", dst.name] subprocess.run(cmd, check=True) return dst.name def concat_wavs(wavs: list[str], repeats: int = 1) -> str: list_file = tempfile.NamedTemporaryFile(suffix=".txt", mode="w", delete=False) for w in wavs * repeats: list_file.write(f"file '{w}'\n") list_file.close() dst = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) dst.close() subprocess.run(["ffmpeg", "-y", "-loglevel", "error", "-f", "concat", "-safe", "0", "-i", list_file.name, "-c", "copy", dst.name], check=True) os.unlink(list_file.name) return dst.name def wav_duration(p: str) -> float: with wave.open(p, "rb") as w: return w.getnframes() / float(w.getframerate()) def vram_mb() -> float: return torch.cuda.max_memory_allocated() / (1024 * 1024) if torch.cuda.is_available() else 0.0 def run_tdt(tdt, wav: str, timestamps: bool) -> dict: torch.cuda.reset_peak_memory_stats() t0 = time.monotonic() try: with torch.amp.autocast("cuda", enabled=True, dtype=torch.bfloat16): with torch.no_grad(): hyps = tdt.transcribe([wav], timestamps=timestamps) elapsed = time.monotonic() - t0 h = hyps[0] text = (getattr(h, "text", None) or "").strip() ts = getattr(h, "timestamp", None) ts_summary = None if isinstance(ts, dict): ts_summary = {k: (len(v) if hasattr(v, "__len__") else str(type(v).__name__)) for k, v in ts.items()} return {"ok": True, "text": text, "chars": len(text), "duration_infer": round(elapsed, 2), "vram_peak_mb": round(vram_mb(), 0), "timestamp_summary": ts_summary} except Exception as e: return {"ok": False, "error": f"{type(e).__name__}: {e}", "duration_infer": round(time.monotonic() - t0, 2)} def run_canary(model, wav: str) -> dict: torch.cuda.reset_peak_memory_stats() t0 = time.monotonic() try: with torch.amp.autocast("cuda", enabled=True, dtype=torch.bfloat16): with torch.no_grad(): hyps = model.transcribe([wav], source_lang="de", target_lang="de", pnc="yes", batch_size=4) elapsed = time.monotonic() - t0 text = (getattr(hyps[0], "text", None) or str(hyps[0])).strip() return {"ok": True, "text": text, "chars": len(text), "duration_infer": round(elapsed, 2), "vram_peak_mb": round(vram_mb(), 0)} except Exception as e: return {"ok": False, "error": f"{type(e).__name__}: {e}", "duration_infer": round(time.monotonic() - t0, 2)} def find_source() -> str | None: for c in ("trim_full.m4a", "full.m4a", "source.m4a"): p = os.path.join(AUDIO_DIR, c) if os.path.exists(p): return p if not os.path.isdir(AUDIO_DIR): return None m4as = [f for f in os.listdir(AUDIO_DIR) if f.endswith(".m4a")] if not m4as: return None m4as.sort(key=lambda f: os.path.getsize(os.path.join(AUDIO_DIR, f)), reverse=True) return os.path.join(AUDIO_DIR, m4as[0]) def main(): open(RESULTS_PATH, "w").close() open(LOG_PATH, "w").close() log("loading EncDecMultiTaskModel (Canary 1B v2)") model = EncDecMultiTaskModel.from_pretrained("nvidia/canary-1b-v2") model = model.to("cuda").to(torch.bfloat16) model.eval() log(f" vram after load: {vram_mb():.0f} MB") tdt = getattr(model, "timestamps_asr_model", None) if tdt is None: log("FATAL: model has no timestamps_asr_model attribute") return log(f" tdt class: {type(tdt).__name__}") log(f" tdt device: {next(tdt.parameters()).device}, dtype: {next(tdt.parameters()).dtype}") src = find_source() if src is None: log(f"FATAL: no source audio found in {AUDIO_DIR}") return log(f"source audio: {src}") wav_full = transcode(src) dur_full = wav_duration(wav_full) log(f" full wav: {dur_full:.1f} s") wav_20 = transcode(src, duration_secs=20.0) log(f" 20 s wav: {wav_duration(wav_20):.1f} s") wav_80 = transcode(src, duration_secs=80.0) if dur_full >= 80 else wav_full log(f" 80 s wav: {wav_duration(wav_80):.1f} s") wav_160 = concat_wavs([wav_80], repeats=2) log(f" 160 s concat wav: {wav_duration(wav_160):.1f} s") cases = [("20s", wav_20), ("80s", wav_80), ("160s_concat", wav_160)] def record(d: dict) -> None: with open(RESULTS_PATH, "a") as f: f.write(json.dumps(d, ensure_ascii=False) + "\n") for label, wav in cases: dur = wav_duration(wav) log(f"\n=== case {label} ({dur:.1f} s) ===") log(" TDT timestamps=False") r = run_tdt(tdt, wav, timestamps=False) record({"case": label, "duration_audio": dur, "model": "tdt", "timestamps": False, **r}) if r["ok"]: log(f" -> {r['chars']} chars in {r['duration_infer']} s, " f"{r['vram_peak_mb']} MB") else: log(f" FAIL: {r['error']}") log(" TDT timestamps=True") r = run_tdt(tdt, wav, timestamps=True) record({"case": label, "duration_audio": dur, "model": "tdt", "timestamps": True, **r}) if r["ok"]: log(f" -> {r['chars']} chars, ts_summary={r['timestamp_summary']}") else: log(f" FAIL: {r['error']}") log(" Canary path1 (reference)") r = run_canary(model, wav) record({"case": label, "duration_audio": dur, "model": "canary_path1", **r}) if r["ok"]: log(f" -> {r['chars']} chars in {r['duration_infer']} s") else: log(f" FAIL: {r['error']}") log("done") if __name__ == "__main__": main()