CoolFace
Modelpublic

iNLP-Lab/Myna-Hokkien

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
14likes73downloads
inference.py69 linesDownload Raw Back to root
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