| |
| """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() |
|
|