iNLP-Lab/Myna-Hokkien
1473
1#!/usr/bin/env python32"""Run one request through the released Myna-Hokkien model."""3from __future__ import annotations4 5import argparse6from pathlib import Path7 8import soundfile as sf9import torch10 11from mynahokkien import MynaHokkien12 13 14def parse_args() -> argparse.Namespace:15 parser = argparse.ArgumentParser()16 parser.add_argument("--model", default="iNLP-Lab/MynaHokkien")17 source = parser.add_mutually_exclusive_group(required=True)18 source.add_argument("--audio", type=Path)19 source.add_argument("--text")20 parser.add_argument("--prompt", help="audio-user-turn override")21 parser.add_argument("--output-wav", type=Path, default=Path("output.wav"))22 parser.add_argument("--output-text", type=Path)23 parser.add_argument("--device", default="cuda:0")24 parser.add_argument("--dtype", choices=("float16", "bfloat16"), default="float16")25 parser.add_argument("--seed", type=int, default=1234)26 parser.add_argument("--max-new-tokens", type=int, default=256)27 parser.add_argument("--max-audio-tokens", type=int, default=4096)28 parser.add_argument("--revision")29 parser.add_argument("--local-files-only", action="store_true")30 return parser.parse_args()31 32 33def main() -> None:34 args = parse_args()35 if args.prompt is not None and args.audio is None:36 raise ValueError("--prompt is only valid with --audio")37 dtype = torch.float16 if args.dtype == "float16" else torch.bfloat1638 model = MynaHokkien.from_pretrained(39 args.model,40 device_map=args.device,41 dtype=dtype,42 revision=args.revision,43 local_files_only=args.local_files_only,44 )45 output = model.generate(46 audio=args.audio,47 text=args.text,48 prompt=args.prompt,49 language="nan",50 speaker="Ethan",51 seed=args.seed,52 max_new_tokens=args.max_new_tokens,53 max_audio_tokens=args.max_audio_tokens,54 )55 if output.text is None or output.audio is None:56 raise RuntimeError("expected both text and audio")57 58 args.output_wav.parent.mkdir(parents=True, exist_ok=True)59 sf.write(args.output_wav, output.audio, output.sampling_rate)60 if args.output_text is not None:61 args.output_text.parent.mkdir(parents=True, exist_ok=True)62 args.output_text.write_text(output.text + "\n", encoding="utf-8")63 print(output.text)64 print(f"[audio] {args.output_wav} ({output.sampling_rate} Hz)")65 66 67if __name__ == "__main__":68 main()69 