ZipVoice.AXERA / infer_zipvoice_axera.py
HY-2012's picture
Add ZipVoice_distill model
a17e6b8 verified
Raw
History Blame
8.47 kB
#!/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/<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()