""" Transcribe a single audio file with a checkpoint, for poking at the model by hand. Not for benchmarking -- see src/evaluate.py for WER over a split. Usage (from project root, venv active): python -m scripts.transcribe --config configs/zipformer_s.yaml \ --checkpoint checkpoints/latest.pt --audio /path/to/clip.wav """ import os os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") import argparse import warnings import torch import torchaudio import yaml from src.evaluate import ctc_greedy_decode from src.model import ASRModel from src.tokenizer import ASRTokenizer, BOS_ID, EOS_ID warnings.filterwarnings("ignore", message=".*output with one or more elements was resized.*") def load_and_resample(path: str, target_sr: int = 16000) -> torch.Tensor: waveform, sr = torchaudio.load(path) # (channels, L) if waveform.size(0) > 1: waveform = waveform.mean(dim=0, keepdim=True) # downmix to mono if sr != target_sr: waveform = torchaudio.transforms.Resample(orig_freq=sr, new_freq=target_sr)(waveform) return waveform.squeeze(0) # (L,) def main(): parser = argparse.ArgumentParser() parser.add_argument("--config", required=True) parser.add_argument("--checkpoint", required=True) parser.add_argument("--audio", required=True, help="Path to any audio file (wav/flac/m4a/mp3/...)") args = parser.parse_args() with open(args.config) as f: cfg = yaml.safe_load(f) device = torch.device("mps" if torch.backends.mps.is_available() else "cpu") tokenizer = ASRTokenizer(cfg["tokenizer_model"]) model = ASRModel(vocab_size=tokenizer.vocab_size, **cfg["model"]).to(device) ckpt = torch.load(args.checkpoint, map_location=device) model.load_state_dict(ckpt["model"]) model.eval() print(f"Loaded checkpoint {args.checkpoint} (epoch {ckpt.get('epoch')}, step {ckpt.get('step')})") waveform = load_and_resample(args.audio).unsqueeze(0).to(device) # (1, L) wave_lengths = torch.tensor([waveform.size(1)], device=device) with torch.no_grad(): enc_out, enc_lengths, ctc_log_probs = model.forward_eval(waveform, wave_lengths) ctc_ids = ctc_greedy_decode(ctc_log_probs, enc_lengths, blank_id=model.blank_id)[0] attn_ids = model.decoder.greedy_decode(enc_out, enc_lengths, bos_id=BOS_ID, eos_id=EOS_ID, max_len=200)[0] print(f"\nCTC : {tokenizer.decode(ctc_ids)}") print(f"ATTN : {tokenizer.decode(attn_ids)}") if __name__ == "__main__": main()