()
| 23 | |
| 24 | |
| 25 | def main(): |
| 26 | parser = argparse.ArgumentParser(description="Fun-ASR-Nano vLLM Inference Demo") |
| 27 | parser.add_argument( |
| 28 | "--model-dir", |
| 29 | type=str, |
| 30 | default="FunAudioLLM/Fun-ASR-Nano-2512", |
| 31 | help="Model name (from hub) or local directory path", |
| 32 | ) |
| 33 | parser.add_argument("--input", type=str, default=None, help="Audio file, wav.scp, or jsonl") |
| 34 | parser.add_argument("--hub", type=str, default="ms", choices=["ms", "hf"]) |
| 35 | parser.add_argument("--device", type=str, default="cuda:0", help="Device for audio encoder") |
| 36 | parser.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16", "fp32"]) |
| 37 | parser.add_argument( |
| 38 | "--tensor-parallel-size", type=int, default=1, help="Number of GPUs for vLLM" |
| 39 | ) |
| 40 | parser.add_argument("--gpu-memory-utilization", type=float, default=0.8) |
| 41 | parser.add_argument("--max-model-len", type=int, default=2048) |
| 42 | parser.add_argument("--max-new-tokens", type=int, default=512) |
| 43 | parser.add_argument("--language", type=str, default="中文", help="Language hint") |
| 44 | parser.add_argument("--hotwords", type=str, nargs="*", default=[], help="Hotwords list") |
| 45 | parser.add_argument("--no-itn", action="store_true", help="Disable inverse text normalization") |
| 46 | parser.add_argument("--batch-size", type=int, default=16, help="Batch size for inference") |
| 47 | parser.add_argument("--output", type=str, default=None, help="Output file for results") |
| 48 | args = parser.parse_args() |
| 49 | |
| 50 | from funasr.models.fun_asr_nano.inference_vllm import FunASRNanoVLLM |
| 51 | |
| 52 | print(f"=" * 60) |
| 53 | print(f"Fun-ASR-Nano vLLM Inference") |
| 54 | print(f"=" * 60) |
| 55 | print(f" Model: {args.model_dir}") |
| 56 | print(f" Tensor Parallel: {args.tensor_parallel_size} GPU(s)") |
| 57 | print(f" Dtype: {args.dtype}") |
| 58 | print(f" Language: {args.language}") |
| 59 | print(f" Hotwords: {args.hotwords or '(none)'}") |
| 60 | print() |
| 61 | |
| 62 | t_load = time.perf_counter() |
| 63 | engine = FunASRNanoVLLM.from_pretrained( |
| 64 | model=args.model_dir, |
| 65 | hub=args.hub, |
| 66 | device=args.device, |
| 67 | dtype=args.dtype, |
| 68 | tensor_parallel_size=args.tensor_parallel_size, |
| 69 | gpu_memory_utilization=args.gpu_memory_utilization, |
| 70 | max_model_len=args.max_model_len, |
| 71 | ) |
| 72 | print(f"Model loaded in {time.perf_counter() - t_load:.1f}s\n") |
| 73 | |
| 74 | # Determine input files |
| 75 | if args.input is None: |
| 76 | # Use default example audio |
| 77 | example_dir = os.path.join(engine.model_dir, "example") |
| 78 | if os.path.isdir(example_dir): |
| 79 | wav_files = [ |
| 80 | os.path.join(example_dir, f) |
| 81 | for f in sorted(os.listdir(example_dir)) |
| 82 | if f.endswith((".wav", ".mp3", ".flac")) |
no test coverage detected