inoryQwQ's picture
三芯片合并:AX620E/AX637 升级 encoder+decoder 全 NPU,新增新一代 SDK;AX650 保持老 SDK
5eee449 verified
Raw
History Blame Contribute Delete
6.94 kB
"""InflectTTS: two-stage (encoder/decoder AXMODEL) VITS TTS pipeline.
Pipeline (per text chunk):
text -> [Host: eSpeak phonemize + symbol ids + intersperse(add_blank)]
-> tokens[1,256] int (0-padded) + x_lengths[1]
-> encoder.axmodel -> m_p/logs_p[1,192,256], logw[1,1,256]
-> [Host: first x_lengths frames; w=exp(logw)*length_scale; ceil;
generate_path; attn matmul -> T' frames;
z_p = m_p' + randn(seed)*exp(logs_p')*variation]
-> Tp=512 chunks (64-frame overlap crossfade) -> decoder.axmodel
-> waveform blocks, tail trimmed
-> edge fade 5 ms + sentence pauses + clip [-1,1]
-> 24 kHz mono float32
The same class drives both hardware targets: AX620E and AX637 AXMODELs share
the identical shape contract, so switching target = switching model paths.
"""
from __future__ import annotations
from pathlib import Path
import numpy as np
from . import host_chain
from .backend import ModelSession
from .frontend import intersperse, text_to_token_ids
from .wav_io import write_wav
# Deterministic EXPORT parity baseline (export/model_meta.json numeric_baseline):
# cleaned_text_to_sequence("ðə kwɪk"), 7 phonemes.
DUMMY_PHONEME_IDS = [81, 83, 16, 53, 65, 102, 53]
class InflectTTS:
ENCODER_T = host_chain.ENCODER_T
DECODER_TP = host_chain.DECODER_TP
SAMPLE_RATE = host_chain.SAMPLE_RATE
def __init__(
self,
encoder_path: str | Path,
decoder_path: str | Path,
backend: str = "auto",
) -> None:
self.encoder = ModelSession(encoder_path, backend)
self.decoder = ModelSession(decoder_path, backend)
# ------------------------------------------------------------------
# Encoder + host chain
# ------------------------------------------------------------------
def _run_encoder(
self, token_ids: list[int], length_scale: float
) -> tuple[np.ndarray, np.ndarray, int]:
"""token ids (already interspersed) -> (m_p_e [T',C], logs_p_e, T')."""
x_len = len(token_ids)
if x_len > self.ENCODER_T:
raise ValueError(
f"token sequence length {x_len} exceeds encoder static T={self.ENCODER_T}"
)
tok_dtype = self.encoder.input_dtype("tokens")
len_dtype = self.encoder.input_dtype("x_lengths")
tokens = np.zeros((1, self.ENCODER_T), dtype=tok_dtype)
tokens[0, :x_len] = np.asarray(token_ids, dtype=tok_dtype)
x_lengths = np.asarray([x_len], dtype=len_dtype)
out = self.encoder.run({"tokens": tokens, "x_lengths": x_lengths})
m_p, logs_p, logw = out["m_p"], out["logs_p"], out["logw"]
# Only the first x_len frames are valid (rest is padding, x_mask-zeroed).
return host_chain.expand_priors(
logw, m_p, logs_p, x_len, length_scale=length_scale
)
def _decode_z_p(self, z_p: np.ndarray) -> np.ndarray:
"""z_p [C, T'] -> wav [T'*256] via chunked decoder."""
z_dtype = self.decoder.input_dtype("z_p")
def run_chunk(z_chunk: np.ndarray) -> np.ndarray:
out = self.decoder.run({"z_p": z_chunk[None].astype(z_dtype)})
return np.asarray(out["wav"], dtype=np.float32).reshape(-1)
return host_chain.decode_waveform(
z_p, run_chunk, self.DECODER_TP, host_chain.DECODER_OVERLAP
)
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def synthesize_tokens(
self,
token_ids: list[int],
*,
speed: float = 1.0,
variation: float = 0.667,
seed: int = 0,
) -> tuple[int, np.ndarray]:
"""Synthesize from raw phoneme ids (NOT interspersed; blank added here).
Skips the eSpeak text frontend — useful for tests or when the frontend
runs elsewhere. Mirrors one chunk of origin model.infer.
"""
if not 0.5 <= speed <= 2.0:
raise ValueError("speed must be between 0.5 and 2.0")
if not 0.0 <= variation <= 1.0:
raise ValueError("variation must be between 0.0 and 1.0")
sequence = intersperse(list(token_ids), 0)
m_p_e, logs_p_e, _ = self._run_encoder(sequence, 1.0 / speed)
z_p = host_chain.inject_noise(m_p_e, logs_p_e, variation, seed)
wav = self._decode_z_p(z_p)
return self.SAMPLE_RATE, np.clip(wav, -1.0, 1.0)
def _token_chunks(self, text: str) -> list[list[int]]:
"""text -> list of interspersed token id lists, each <= ENCODER_T-1."""
ids = text_to_token_ids(text) # interspersed, 2N+1
if len(ids) <= self.ENCODER_T - 1:
return [ids]
words = text.split()
if len(words) < 2:
raise ValueError(
f"single unbreakable chunk produces {len(ids)} tokens "
f"(> {self.ENCODER_T - 1}); shorten the text"
)
mid = len(words) // 2
return self._token_chunks(" ".join(words[:mid])) + self._token_chunks(
" ".join(words[mid:])
)
def synthesize(
self,
text: str,
*,
speed: float = 1.0,
variation: float = 0.667,
seed: int = 0,
) -> tuple[int, np.ndarray]:
"""Full text -> 24 kHz mono waveform (mirrors origin inference.py)."""
normalized = " ".join(text.split())
if not normalized:
raise ValueError("Text must not be empty.")
if not 0.5 <= speed <= 2.0:
raise ValueError("speed must be between 0.5 and 2.0")
if not 0.0 <= variation <= 1.0:
raise ValueError("variation must be between 0.0 and 1.0")
length_scale = 1.0 / speed
chunks = host_chain.split_text(normalized)
pieces: list[np.ndarray] = []
prev_ending = ""
index = 0
for chunk in chunks:
for token_ids in self._token_chunks(chunk):
if index:
pause = host_chain.boundary_pause_seconds(prev_ending)
pieces.append(
np.zeros(
host_chain.seconds_to_samples(pause), dtype=np.float32
)
)
m_p_e, logs_p_e, _ = self._run_encoder(token_ids, length_scale)
z_p = host_chain.inject_noise(m_p_e, logs_p_e, variation, seed + index)
wav = self._decode_z_p(z_p)
pieces.append(host_chain.edge_fade(wav, self.SAMPLE_RATE))
index += 1
prev_ending = chunk
waveform = np.clip(np.concatenate(pieces), -1.0, 1.0)
return self.SAMPLE_RATE, waveform
def save(self, text: str, output: str | Path, **kwargs) -> Path:
sample_rate, waveform = self.synthesize(text, **kwargs)
return write_wav(output, waveform, sample_rate)