asr / scripts /transcribe.py
shubhexists's picture
Add Zipformer-inspired ASR model: weights, tokenizer, config, and training code
ce3c8df verified
Raw History Blame Contribute Delete
2.51 kB
"""
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()