#!/usr/bin/env python3 from __future__ import annotations import argparse import json import logging import time from pathlib import Path import numpy as np import soundfile as sf from scripts.common_infer import extract_prompt_features, load_tokenizer from scripts.common_infer import load_vocoder, vocoder_decode_loaded from scripts.text_processing import build_cat_tokens, build_segments from scripts.text_processing import load_text, normalize_punctuation, setup_jieba_cache from scripts.zipvoice_decoder4_runtime import Decoder4ZipVoiceBoardRuntime def parse_args() -> argparse.Namespace: repo_dir = Path(__file__).resolve().parent p = argparse.ArgumentParser( description="ZipVoice encoder + decoder4 inference on AXERA boards", formatter_class=argparse.ArgumentDefaultsHelpFormatter, ) p.add_argument( "--model-name", default="zipvoice_ax650", choices=["zipvoice_ax650", "zipvoice_distill_ax650", "zipvoice_distill_ax630C"], help="Model folder under models/", ) p.add_argument( "--model-dir", default=None, help="Override model directory. If unset, models/ is used.", ) p.add_argument("--text", default=None, help="Text to synthesize") p.add_argument("--text-file", default=None, help="UTF-8 text file containing text") p.add_argument("--prompt-text", required=True) p.add_argument("--prompt-wav", required=True) p.add_argument("--output-wav", default="outputs/zipvoice_axera.wav") p.add_argument("--num-step", type=int, default=None) p.add_argument("--guidance-scale", type=float, default=None) p.add_argument("--speed", type=float, default=None) p.add_argument("--t-shift", type=float, default=None) p.add_argument("--seed", type=int, default=42) p.add_argument("--max-tokens", type=int, default=384) p.add_argument("--max-feat-len", type=int, default=1024) p.add_argument("--feat-scale", type=float, default=0.1) p.add_argument("--target-rms", type=float, default=0.1) p.add_argument("--min-generated-frames", type=int, default=360) p.add_argument("--max-generated-frames", type=int, default=620) p.add_argument( "--max-raw-feat-ratio", type=float, default=1.2, help="Allow raw features_len up to this ratio of max_feat_len before splitting", ) p.add_argument("--silence-ms", type=int, default=140) return p.parse_args() def load_runtime_defaults(model_dir: Path) -> dict: config_path = model_dir / "runtime_config.json" if not config_path.exists(): return {} return json.loads(config_path.read_text()) def main() -> None: args = parse_args() logging.basicConfig(level=logging.INFO, format="%(message)s") repo_dir = Path(__file__).resolve().parent setup_jieba_cache(repo_dir) model_dir = Path(args.model_dir) if args.model_dir else repo_dir / "models" / args.model_name runtime_defaults = load_runtime_defaults(model_dir) args.max_tokens = int(runtime_defaults.get("max_tokens", args.max_tokens)) args.max_feat_len = int(runtime_defaults.get("max_feat_len", args.max_feat_len)) if args.num_step is None: args.num_step = int(runtime_defaults.get("num_step", 12)) if args.guidance_scale is None: args.guidance_scale = float(runtime_defaults.get("guidance_scale", 1.0)) if args.speed is None: args.speed = float(runtime_defaults.get("speed", 1.0)) if args.t_shift is None: args.t_shift = float(runtime_defaults.get("t_shift", 0.5)) output_wav = Path(args.output_wav) output_wav.parent.mkdir(parents=True, exist_ok=True) t_target_text_start = time.perf_counter() original_text = load_text(args) text = normalize_punctuation(original_text) target_text_preprocess_sec = time.perf_counter() - t_target_text_start prompt_text = normalize_punctuation(args.prompt_text) if text != original_text: logging.debug("Normalized text punctuation: %s", text) if prompt_text != args.prompt_text: logging.debug("Normalized prompt punctuation: %s", prompt_text) logging.debug("Loading tokenizer...") tokenizer = load_tokenizer(repo_dir) prompt_tokens = tokenizer.texts_to_token_ids([prompt_text])[0] max_text_tokens = args.max_tokens - len(prompt_tokens) - 1 if max_text_tokens <= 0: raise ValueError("prompt_tokens leaves no room for text") logging.debug("Extracting prompt features from %s...", args.prompt_wav) prompt_features, prompt_rms = extract_prompt_features( args.prompt_wav, repo_dir=repo_dir, sampling_rate=24000, feat_scale=args.feat_scale, target_rms=args.target_rms, ) prompt_frames = int(prompt_features.shape[1]) logging.debug( "prompt_tokens=%d, prompt_frames=%d", len(prompt_tokens), prompt_frames, ) t_segment_planning_start = time.perf_counter() segments = build_segments( tokenizer=tokenizer, text=text, prompt_frames=prompt_frames, prompt_tokens_len=len(prompt_tokens), speed=args.speed, max_feat_len=args.max_feat_len, max_text_tokens=max_text_tokens, min_generated_frames=args.min_generated_frames, max_generated_frames=args.max_generated_frames, max_raw_feat_ratio=args.max_raw_feat_ratio, ) segment_planning_sec = time.perf_counter() - t_segment_planning_start logging.debug("Built %d segments", len(segments)) for index, segment in enumerate(segments, start=1): logging.debug( "segment %02d: text_tokens=%d generated_frames=%d features_len=%d clamped=%s text=%s", index, segment["text_tokens"], segment["generated_frames"], segment["features_len"], segment["clamped"], segment["text"], ) logging.info("正在加载模型...") runtime = Decoder4ZipVoiceBoardRuntime( config_dir=model_dir, models_dir=model_dir, max_feat_len=args.max_feat_len, max_tokens=args.max_tokens, num_step=args.num_step, t_shift=args.t_shift, ) logging.debug("Loading vocoder...") vocoder = load_vocoder(repo_dir) logging.info("模型加载完成") audios: list[np.ndarray] = [] silence = np.zeros(int(24000 * args.silence_ms / 1000), dtype=np.float32) total_segment_sec = 0.0 for index, segment in enumerate(segments, start=1): logging.info("开始推理第 %d/%d 句", index, len(segments)) t_segment_start = time.perf_counter() seg_text = str(segment["text"]) text_tokens = tokenizer.texts_to_token_ids([seg_text])[0] cat_tokens = build_cat_tokens(tokenizer, prompt_tokens, text_tokens, args.max_tokens) pred_features, timing = runtime.sample( cat_tokens=cat_tokens, prompt_tokens_len=len(prompt_tokens), text_tokens_len=len(text_tokens), prompt_features=prompt_features, prompt_features_len=prompt_frames, speed=args.speed, guidance_scale=args.guidance_scale, seed=args.seed + index - 1, ) audio = vocoder_decode_loaded( vocoder, pred_features, feat_scale=args.feat_scale, target_rms=args.target_rms, prompt_rms=prompt_rms, ) segment_sec = time.perf_counter() - t_segment_start audio_sec = len(audio) / 24000 if audios: audios.append(silence) audios.append(audio.astype(np.float32, copy=False)) total_segment_sec += segment_sec logging.debug( "segment %02d done: audio=%.2fs runtime=%.2fs", index, audio_sec, segment_sec, ) final_audio = np.concatenate(audios) if audios else np.zeros(0, dtype=np.float32) sf.write(str(output_wav), final_audio, samplerate=24000) final_audio_sec = len(final_audio) / 24000 runtime_sec_for_rtf = target_text_preprocess_sec + segment_planning_sec + total_segment_sec rtf = ( runtime_sec_for_rtf / final_audio_sec if final_audio_sec > 0 else float("inf") ) logging.info("推理耗时: %.3fs", runtime_sec_for_rtf) logging.info("生成语音时长: %.3fs", final_audio_sec) logging.info("RTF: %.4f", rtf) if __name__ == "__main__": main()