from __future__ import annotations import argparse import logging import sys from pathlib import Path from typing import Optional, Sequence sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) import numpy as np _HERE = Path(__file__).resolve().parent def parse_args(argv: Optional[Sequence[str]] = None) -> argparse.Namespace: p = argparse.ArgumentParser( description=( "MOSS-TTS-Nano:" "默认走全 AXModel 快速路径;也可仅将 " "local_fixed_sampled_frame 切换为 ONNX Runtime CPU," "用于验证量化后的离散 Token 分叉。" ) ) p.add_argument("--config-dir", default=str(_HERE.parent / "config"), help="config directory (default: ./configs)") p.add_argument("--axmodel-dir", default=str(_HERE.parent / "models" / "axmodels_650"), help="axmodel directory (default: ./models/axmodels)") p.add_argument("--onnx-dir", default=str(_HERE.parent / "models" / "onnxmodels"), help="ONNX directory (default: ./models/onnxmodels)") p.add_argument( "--prefill-backend", choices=("axmodel", "onnx"), default="axmodel", help="prefill backend (default: axmodel)", ) p.add_argument( "--decode-backend", choices=("axmodel", "onnx"), default="axmodel", help="decode_step backend (default: axmodel)", ) p.add_argument( "--local-fixed-backend", choices=("axmodel", "onnx"), default="axmodel", help="local_fixed_sampled_frame backend (default: axmodel)", ) p.add_argument("--output-audio-path", default=str(_HERE / "outputs" / "board_axmodel_decode.wav"), help="output WAV path") text_group = p.add_mutually_exclusive_group(required=True) text_group.add_argument("--text", help="待合成文本") text_group.add_argument("--text-file", help="UTF-8 文本文件路径") p.add_argument("--voice", default="Junhao", help="内置音色名称") p.add_argument("--prompt-audio-path", "--reference-audio-path", dest="prompt_audio_path", default=None, help="参考音频路径(语音克隆,覆盖 --voice)") p.add_argument( "--sample-mode", choices=("greedy", "fixed", "full"), default="fixed", help="fixed=单次 local-fixed 快速路径;full/greedy=逐 codebook local-decoder 慢路径", ) p.add_argument( "--do-sample", type=int, default=1, choices=[0, 1], help="兼容参数;默认与 fixed 快速路径配置一致", ) p.add_argument("--streaming", type=int, default=0, choices=[0, 1], help="按文本 chunk 增量写 WAV (1);不改变 codec 全量解码方式") p.add_argument("--max-new-frames", type=int, default=150, help="最大生成帧数") p.add_argument("--voice-clone-max-text-tokens", type=int, default=75, help="语音克隆每段最大 token 数") p.add_argument("--seed", type=int, default=None, help="随机种子") p.add_argument( "--greedy-prefix-frames", type=int, default=4, help="前 N 帧使用 top-1 音频采样,之后恢复随机采样", ) p.add_argument( "--assistant-random-u", type=float, default=None, help="可选:固定每帧继续/停止采样值,取值范围 [0, 1)", ) p.add_argument("--debug", action="store_true", help="启用 debug 日志") args = p.parse_args(argv) if args.assistant_random_u is not None and not 0.0 <= args.assistant_random_u < 1.0: p.error("--assistant-random-u must be in [0, 1)") return args def main(argv: Optional[Sequence[str]] = None) -> dict: args = parse_args(argv) logging.basicConfig( format="%(asctime)s %(levelname)s %(name)s: %(message)s", level=logging.DEBUG if args.debug else logging.INFO, stream=sys.stderr, ) raw_text = ( str(args.text) if args.text is not None else Path(args.text_file).read_text(encoding="utf-8") ) logging.info("config_dir: %s", args.config_dir) logging.info("axmodel_dir: %s", args.axmodel_dir) logging.info("onnx_dir: %s", args.onnx_dir) logging.info("text: %r", raw_text[:60]) from scripts.tts_runtime import AxTtsRuntime runtime = AxTtsRuntime( config_dir=args.config_dir, axmodel_dir=args.axmodel_dir, onnx_dir=args.onnx_dir, use_onnx_prefill=args.prefill_backend == "onnx", use_onnx_decode=args.decode_backend == "onnx", use_onnx_local_fixed=args.local_fixed_backend == "onnx", max_new_frames=args.max_new_frames, do_sample=bool(args.do_sample), sample_mode=args.sample_mode, ) logging.info("decode_step backend: %s", args.decode_backend.upper()) logging.info("prefill backend: %s", args.prefill_backend.upper()) local_stage = ( f"local_fixed_sampled_frame({args.local_fixed_backend}, fast)" if args.sample_mode == "fixed" else "local_decoder(axmodel, slow x17/frame)" ) logging.info( "chain: prefill(%s) -> %s -> decode_step(%s) -> codec_decode(axmodel)", args.prefill_backend, local_stage, args.decode_backend, ) if args.prompt_audio_path: logging.info("参考音频: %s", args.prompt_audio_path) else: logging.info("内置音色: %s", args.voice) result = runtime.synthesize( text=raw_text, voice=args.voice, prompt_audio_path=args.prompt_audio_path, output_audio_path=args.output_audio_path, sample_mode=args.sample_mode, do_sample=bool(args.do_sample), streaming=bool(args.streaming), max_new_frames=args.max_new_frames, voice_clone_max_text_tokens=args.voice_clone_max_text_tokens, seed=args.seed, greedy_prefix_frames=args.greedy_prefix_frames, assistant_random_u=args.assistant_random_u, ) token_path = Path(args.output_audio_path).expanduser().resolve().with_suffix(".tokens.npy") token_path.parent.mkdir(parents=True, exist_ok=True) np.save(token_path, np.asarray(result["audio_token_ids"], dtype=np.int32)) result["audio_token_path"] = str(token_path) t = result["timing"] print(f"\n{'='*55}") result_kind = ( "板端混合 ONNX/AXModel 推理结果" if ( args.local_fixed_backend == "onnx" or args.prefill_backend == "onnx" or args.decode_backend == "onnx" ) else "板端全 AXModel 推理结果" ) print(f" {result_kind}") print(f"{'='*55}") print(f" 输出: {result['audio_path']}") print(f" Token: {result['audio_token_path']}") print(f" 音频时长: {t['audio_duration_sec']:.2f}s") print(f" 推理耗时: {t['total_infer_time_sec']:.2f}s") print(f" RTF: {t['rtf']:.4f}") print(f" 生成模型 RTF: {t['generation_model_rtf']:.4f}") print( " local calls: " f"fixed={t['per_model_calls'].get('local_fixed_sampled_frame', 0)} " f"decoder={t['per_model_calls'].get('local_decoder', 0)}" ) print(f"{'='*55}\n") return result if __name__ == "__main__": main()