CoolFace
Modelpublic

RASMUS/Finnish-ASR-Canary-v2

sourceHugging Facemitupdated 7mo agoView on Hugging Face
0likes1.2kdownloads
inference_example.py61 linesDownload Raw Back to root
1from nemo.collections.asr.models import EncDecMultiTaskModel2from omegaconf import OmegaConf3import os4import argparse5 6def main():7    parser = argparse.ArgumentParser(description="Finnish ASR Inference Example")8    parser.add_argument("--audio", type=str, required=True, help="Path to the audio file (.wav)")9    parser.add_argument("--model", type=str, default="models/canary-finnish.nemo", help="Path to the finetuned .nemo model")10    parser.add_argument("--kenlm", type=str, default="models/kenlm_5M.nemo", help="Path to the KenLM model")11    parser.add_argument("--beam_size", type=int, default=4, help="Beam size for decoding")12    parser.add_argument("--pnc", type=str, default="yes", help="Enable Punctuation and Capitalization (yes/no)")13    14    args = parser.parse_args()15 16    # 1. Load Model and KenLM Bundle17    if not os.path.exists(args.model):18        print(f"Error: Model not found at {args.model}")19        return20 21    print(f"Loading model from {args.model}...")22    model = EncDecMultiTaskModel.restore_from(args.model)23 24    # Configure KenLM if provided25    if args.kenlm and os.path.exists(args.kenlm):26        print(f"Configuring decoding strategy with KenLM from {args.kenlm}...")27        model.change_decoding_strategy(28            decoding_cfg=OmegaConf.create({29                'strategy': 'beam',30                'beam': {31                    'beam_size': args.beam_size,32                    'ngram_lm_model': args.kenlm,33                    'ngram_lm_alpha': 0.2,34                },35                'batch_size': 136            })37        )38    else:39        print("Using greedy decoding (no KenLM found or specified).")40 41    # 2. Transcribe with Finnish Prompts42    if not os.path.exists(args.audio):43        print(f"Error: Audio sample not found at {args.audio}")44        return45 46    print(f"Transcribing {args.audio}...")47    transcription = model.transcribe(48        audio=[args.audio],49        taskname="asr",50        source_lang="fi",51        target_lang="fi",52        pnc=args.pnc53    )54 55    print("-" * 30)56    print(f"Result: {transcription[0]}")57    print("-" * 30)58 59if __name__ == "__main__":60    main()61