| """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 |
|
|
| |
| |
| 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) |
|
|
| |
| |
| |
| 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"] |
| |
| 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 |
| ) |
|
|
| |
| |
| |
| 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) |
| 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) |
|
|