#!/usr/bin/env python3 """Run one request through the released Myna-Hokkien model.""" from __future__ import annotations import argparse from pathlib import Path import soundfile as sf import torch from mynahokkien import MynaHokkien def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--model", default="iNLP-Lab/MynaHokkien") source = parser.add_mutually_exclusive_group(required=True) source.add_argument("--audio", type=Path) source.add_argument("--text") parser.add_argument("--prompt", help="audio-user-turn override") parser.add_argument("--output-wav", type=Path, default=Path("output.wav")) parser.add_argument("--output-text", type=Path) parser.add_argument("--device", default="cuda:0") parser.add_argument("--dtype", choices=("float16", "bfloat16"), default="float16") parser.add_argument("--seed", type=int, default=1234) parser.add_argument("--max-new-tokens", type=int, default=256) parser.add_argument("--max-audio-tokens", type=int, default=4096) parser.add_argument("--revision") parser.add_argument("--local-files-only", action="store_true") return parser.parse_args() def main() -> None: args = parse_args() if args.prompt is not None and args.audio is None: raise ValueError("--prompt is only valid with --audio") dtype = torch.float16 if args.dtype == "float16" else torch.bfloat16 model = MynaHokkien.from_pretrained( args.model, device_map=args.device, dtype=dtype, revision=args.revision, local_files_only=args.local_files_only, ) output = model.generate( audio=args.audio, text=args.text, prompt=args.prompt, language="nan", speaker="Ethan", seed=args.seed, max_new_tokens=args.max_new_tokens, max_audio_tokens=args.max_audio_tokens, ) if output.text is None or output.audio is None: raise RuntimeError("expected both text and audio") args.output_wav.parent.mkdir(parents=True, exist_ok=True) sf.write(args.output_wav, output.audio, output.sampling_rate) if args.output_text is not None: args.output_text.parent.mkdir(parents=True, exist_ok=True) args.output_text.write_text(output.text + "\n", encoding="utf-8") print(output.text) print(f"[audio] {args.output_wav} ({output.sampling_rate} Hz)") if __name__ == "__main__": main()