Myna-Hokkien / inference.py
mattChrisP's picture
release myna
28e721c verified
Raw
History Blame Contribute Delete
2.47 kB
#!/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()