| |
|
|
| 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/<model-name> 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() |
|
|