"""De-risk before building around these models: do they load, fit, and run fast enough?""" import logging import time import numpy as np import torch logging.basicConfig(level=logging.INFO, format="%(message)s") log = logging.getLogger("probe") DEV = "cuda" if torch.cuda.is_available() else "cpu" def vram(tag): if DEV == "cuda": log.info(" VRAM %-10s alloc %.2f GB | reserved %.2f GB", tag, torch.cuda.memory_allocated() / 1e9, torch.cuda.memory_reserved() / 1e9) log.info("device=%s", DEV) if DEV == "cuda": log.info("gpu=%s total=%.1f GB", torch.cuda.get_device_name(0), torch.cuda.get_device_properties(0).total_memory / 1e9) # ── STT ────────────────────────────────────────────────────────────────────── log.info("\n[1/2] loading IndicConformer…") t0 = time.perf_counter() from transformers import AutoModel stt = AutoModel.from_pretrained("ai4bharat/indic-conformer-600m-multilingual", trust_remote_code=True) stt = stt.to(DEV).eval() log.info(" loaded in %.1fs", time.perf_counter() - t0) vram("after STT") # 3 s of quiet noise — we only care that a forward pass runs and how long it takes. wav = torch.from_numpy((np.random.randn(16000 * 3) * 0.01).astype(np.float32)).unsqueeze(0).to(DEV) for i in range(2): t0 = time.perf_counter() with torch.inference_mode(): out = stt(wav, "ta", "ctc") log.info(" pass %d: %.0f ms -> %r", i + 1, (time.perf_counter() - t0) * 1000, str(out)[:60]) # ── TTS ────────────────────────────────────────────────────────────────────── log.info("\n[2/2] loading Indic Parler-TTS…") t0 = time.perf_counter() from parler_tts import ParlerTTSForConditionalGeneration, ParlerTTSStreamer from transformers import AutoTokenizer dtype = torch.float16 if DEV == "cuda" else torch.float32 tts = ParlerTTSForConditionalGeneration.from_pretrained("ai4bharat/indic-parler-tts", torch_dtype=dtype).to(DEV).eval() tok = AutoTokenizer.from_pretrained("ai4bharat/indic-parler-tts") dtok = AutoTokenizer.from_pretrained(tts.config.text_encoder._name_or_path) log.info(" loaded in %.1fs sr=%d", time.perf_counter() - t0, tts.config.sampling_rate) vram("after TTS") desc = "Jaya speaks in a warm, clear, professional tone at a natural pace. The recording is very high quality with no background noise." prompt = "உங்கள் புதிய லீட்கள் மூன்று ஆயிரம் நானூறு." d = dtok(desc, return_tensors="pt").to(DEV) p = tok(prompt, return_tensors="pt").to(DEV) kw = dict(input_ids=d.input_ids, attention_mask=d.attention_mask, prompt_input_ids=p.input_ids, prompt_attention_mask=p.attention_mask) # Streaming: what the user actually experiences is time-to-first-audio. frame_rate = getattr(tts.audio_encoder.config, "frame_rate", 86) streamer = ParlerTTSStreamer(tts, device=DEV, play_steps=int(frame_rate / 2)) from threading import Thread t0 = time.perf_counter() Thread(target=tts.generate, kwargs={**kw, "streamer": streamer}, daemon=True).start() first_ms, total = None, 0 for chunk in streamer: if chunk is None or len(chunk) == 0: continue if first_ms is None: first_ms = (time.perf_counter() - t0) * 1000 total += len(chunk) gen_ms = (time.perf_counter() - t0) * 1000 audio_ms = 1000 * total / tts.config.sampling_rate log.info(" time to FIRST audio : %.0f ms", first_ms or -1) log.info(" full generation : %.0f ms for %.0f ms of audio", gen_ms, audio_ms) log.info(" realtime factor : %.2fx (<1 means faster than realtime)", gen_ms / max(audio_ms, 1)) vram("peak") if DEV == "cuda": log.info(" peak reserved: %.2f GB", torch.cuda.max_memory_reserved() / 1e9)