Download scripts/transcribe.py from shubhexists/asr: direct link, hf CLI and curl.
- Browser
- Download file 2.51 kB
-
https://huggingface.co/shubhexists/asr/resolve/main/scripts/transcribe.py
- Command line
-
hf download hf://shubhexists/asr/scripts/transcribe.py
-
curl -L -o transcribe.py https://huggingface.co/shubhexists/asr/resolve/main/scripts/transcribe.py
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() | |