"""Production inference helpers for the BlueMagpie-TTS Space.""" from __future__ import annotations import math import random import re import unicodedata import numpy as np import torch from torch import nn _SPACE_RE = re.compile(r"\s+") _PUNCT_NO_LEFT_SPACE_RE = re.compile(r"\s+([,。!?;:、,.!?;:])") _CJK_PUNCT_RIGHT_SPACE_RE = re.compile(r"([,。!?;:、])\s+") _BOPOMOFO_TONES = {"ˊ", "ˇ", "ˋ", "˙"} _BOPOMOFO_READINGS = { "ㄅ": "波", "ㄆ": "坡", "ㄇ": "摸", "ㄈ": "佛", "ㄉ": "得", "ㄊ": "特", "ㄋ": "呢", "ㄌ": "了", "ㄍ": "哥", "ㄎ": "科", "ㄏ": "喝", "ㄐ": "基", "ㄑ": "欺", "ㄒ": "希", "ㄓ": "知", "ㄔ": "吃", "ㄕ": "師", "ㄖ": "日", "ㄗ": "資", "ㄘ": "雌", "ㄙ": "思", "ㄚ": "啊", "ㄛ": "喔", "ㄜ": "鵝", "ㄝ": "欸", "ㄞ": "哀", "ㄟ": "欸", "ㄠ": "凹", "ㄡ": "歐", "ㄢ": "安", "ㄣ": "恩", "ㄤ": "昂", "ㄥ": "鞥", "ㄦ": "兒", "ㄧ": "衣", "ㄨ": "烏", "ㄩ": "迂", } class StopHysteresisController(nn.Module): """Apply probability-threshold hysteresis around a legacy stop head. The pinned public model predates native ``stop_threshold`` and ``stop_consecutive`` arguments. This adapter preserves the learned logits and only changes the final stop decision used by legacy generation. """ def __init__( self, stop_head: nn.Module, threshold: float = 0.65, late_threshold: float = 0.50, consecutive: int = 2, ) -> None: super().__init__() self.stop_head = stop_head self.threshold = float(threshold) self.late_threshold = float(late_threshold) self.consecutive = max(1, int(consecutive)) self._active = False self._min_len = 2 self._expected_steps = 0 self._hard_stop_steps = 0 self._step = 0 self._hits = 0 def begin( self, min_len: int, *, expected_steps: int = 0, hard_stop_steps: int = 0, ) -> None: self._active = True self._min_len = max(0, int(min_len)) self._expected_steps = max(0, int(expected_steps)) self._hard_stop_steps = max(0, int(hard_stop_steps)) self._step = 0 self._hits = 0 def end(self) -> None: self._active = False def forward(self, hidden: torch.Tensor) -> torch.Tensor: logits = self.stop_head(hidden) if not self._active: return logits probability = float(torch.softmax(logits.float(), dim=-1)[0, 1].detach().cpu()) generated_steps = self._step + 1 threshold = self.threshold if self._expected_steps > 0: progress = generated_steps / float(self._expected_steps) if progress >= 0.8: blend = min(1.0, max(0.0, (progress - 0.8) / 0.2)) threshold = self.threshold + blend * (self.late_threshold - self.threshold) eligible = self._step > self._min_len if eligible and probability >= threshold: self._hits += 1 else: self._hits = 0 should_stop = eligible and self._hits >= self.consecutive if self._hard_stop_steps > 0 and generated_steps >= self._hard_stop_steps: should_stop = True self._step += 1 decision = torch.zeros_like(logits) decision[..., 1 if should_stop else 0] = 1.0 return decision def set_generation_seed(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def _is_cjk(char: str) -> bool: codepoint = ord(char) return ( 0x3400 <= codepoint <= 0x4DBF or 0x4E00 <= codepoint <= 0x9FFF or 0xF900 <= codepoint <= 0xFAFF or 0x3040 <= codepoint <= 0x30FF or 0xAC00 <= codepoint <= 0xD7AF ) def _is_bopomofo(char: str) -> bool: codepoint = ord(char) return 0x3100 <= codepoint <= 0x312F or 0x31A0 <= codepoint <= 0x31BF or char in _BOPOMOFO_TONES def normalize_tts_text(text: str) -> str: """Normalize common punctuation and pronounce standalone Bopomofo symbols.""" raw = unicodedata.normalize("NFC", str(text or "")) has_cjk = any(_is_cjk(char) for char in raw) output: list[str] = [] index = 0 while index < len(raw): char = raw[index] if char.isspace() or char in {"\r", "\n", "\t", "\u3000"}: output.append(" ") index += 1 continue if unicodedata.category(char)[0] == "C": index += 1 continue if raw.startswith("...", index): output.append("。" if has_cjk else ".") index += 3 while index < len(raw) and raw[index] == ".": index += 1 continue if char in {"…", "⋯"}: output.append("。" if has_cjk else ".") index += 1 while index < len(raw) and raw[index] in {"…", "⋯"}: index += 1 continue if _is_bopomofo(char): while index < len(raw) and _is_bopomofo(raw[index]): reading = _BOPOMOFO_READINGS.get(raw[index]) if reading: output.append(reading) index += 1 continue previous = raw[index - 1] if index else "" following = raw[index + 1] if index + 1 < len(raw) else "" folded = unicodedata.normalize("NFKC", char) if folded.isascii() and folded.isalnum(): char = folded if char in {",", ",", "﹐", "、"}: output.append("," if previous.isdigit() and following.isdigit() else ("," if has_cjk else ",")) elif char in {".", "。", "。", "."}: inside_ascii = previous.isascii() and previous.isalnum() and following.isascii() and following.isalnum() output.append("." if inside_ascii or not has_cjk else "。") elif char in {"?", "?", "﹖"}: output.append("?" if has_cjk else "?") elif char in {"!", "!", "﹗"}: output.append("!" if has_cjk else "!") elif char in {";", ";"}: output.append(";" if has_cjk else ";") elif char in {":", ":"}: output.append(":" if following == "/" or previous.isdigit() and following.isdigit() else (":" if has_cjk else ":")) elif char in {"“", "”", "„", """}: output.append('"') elif char in {"‘", "’", "'"}: output.append("'") elif char == "(": output.append("(") elif char == ")": output.append(")") else: output.append(char) index += 1 normalized = _SPACE_RE.sub(" ", "".join(output)) normalized = _PUNCT_NO_LEFT_SPACE_RE.sub(r"\1", normalized) normalized = _CJK_PUNCT_RIGHT_SPACE_RE.sub(r"\1", normalized) return normalized.strip() def count_speech_units(text: str) -> int: """Count CJK characters and compressed ASCII runs for pace control.""" units = 0 ascii_buffer: list[str] = [] def flush_ascii() -> None: nonlocal units if not ascii_buffer: return token = "".join(ascii_buffer) divisor = 2 if token.isdigit() else 4 units += max(1, math.ceil(len(token) / divisor)) ascii_buffer.clear() for char in normalize_tts_text(text).lower(): if char.isascii() and char.isalnum(): ascii_buffer.append(char) elif _is_cjk(char): flush_ascii() units += 1 else: flush_ascii() flush_ascii() return units def estimate_step_seconds(model, sample_rate: int) -> float | None: patch_size = int( getattr(model, "patch_size", 0) or getattr(getattr(model, "config", None), "patch_size", 0) or 0 ) audio_vae = getattr(model, "audio_vae", None) decode_chunk = int( getattr(model, "_decode_chunk_size", 0) or getattr(audio_vae, "decode_chunk_size", 0) or getattr(audio_vae, "chunk_size", 0) or 0 ) if patch_size <= 0 or decode_chunk <= 0 or sample_rate <= 0: return None return float(patch_size * decode_chunk / sample_rate) def target_cps_min_len( text: str, target_cps: float, step_seconds: float | None, base_min_len: int = 2, stop_consecutive: int = 1, ) -> int: if target_cps <= 0.0 or not step_seconds or step_seconds <= 0.0: return int(base_min_len) units = count_speech_units(text) if units <= 0: return int(base_min_len) target_steps = math.ceil((units / target_cps) / step_seconds) earliest_stop_offset = max(1, int(stop_consecutive)) + 1 return max(int(base_min_len), max(0, target_steps - earliest_stop_offset)) def target_cps_steps(text: str, target_cps: float, step_seconds: float | None) -> int: if target_cps <= 0.0 or not step_seconds or step_seconds <= 0.0: return 0 units = count_speech_units(text) if units <= 0: return 0 return max(1, math.ceil((units / target_cps) / step_seconds)) def target_pace_speed( audio_samples: int, sample_rate: int, text: str, *, target_cps: float, min_speed: float = 0.80, ) -> float: """Return a pitch-preserving stretch rate without extending generation.""" units = count_speech_units(text) if audio_samples <= 0 or sample_rate <= 0 or units <= 0 or target_cps <= 0.0: return 1.0 actual_seconds = float(audio_samples) / float(sample_rate) target_seconds = float(units) / float(target_cps) if actual_seconds >= target_seconds: return 1.0 return min(1.0, max(float(min_speed), actual_seconds / target_seconds)) def duration_hard_stop_steps( expected_steps: int, *, ratio: float = 1.08, margin_steps: int = 3, fallback: int = 2000, ) -> int: expected_steps = max(0, int(expected_steps)) if expected_steps <= 0: return max(1, int(fallback)) return max( expected_steps + max(0, int(margin_steps)), math.ceil(expected_steps * max(1.0, float(ratio))), ) def finish_audio( audio: np.ndarray, sample_rate: int, *, fade_ms: float = 60.0, trailing_silence_ms: float = 180.0, ) -> np.ndarray: """Fade a forced endpoint and leave a short, unambiguous final pause.""" signal = np.asarray(audio, dtype=np.float32).reshape(-1).copy() fade_samples = min( signal.size, max(0, int(round(float(fade_ms) * int(sample_rate) / 1000.0))), ) if fade_samples > 0: signal[-fade_samples:] *= np.linspace(1.0, 0.0, fade_samples, dtype=np.float32) silence_samples = max( 0, int(round(float(trailing_silence_ms) * int(sample_rate) / 1000.0)), ) if silence_samples > 0: signal = np.pad(signal, (0, silence_samples)) return signal def _hard_split_text(text: str, max_chars: int, min_chunk_chars: int) -> list[str]: chunks: list[str] = [] remaining = text.strip() min_chunk_chars = max(1, min(int(min_chunk_chars), max(1, int(max_chars)))) while len(remaining) > max_chars: window = remaining[: max_chars + 1] split_at = max((window.rfind(char) for char in ",,、;;:: "), default=-1) cut = max_chars if split_at < min_chunk_chars else split_at + (not window[split_at].isspace()) tail_len = len(remaining) - cut if 0 < tail_len < min_chunk_chars and len(remaining) >= 2 * min_chunk_chars: cut = len(remaining) - min_chunk_chars chunk = remaining[:cut].strip() if chunk: chunks.append(chunk) remaining = remaining[cut:].lstrip() if remaining: chunks.append(remaining) return chunks def split_text_for_tts(text: str, max_chars: int = 80, min_chunk_chars: int = 12) -> list[str]: """Split at punctuation while preserving it on the preceding chunk.""" text = normalize_tts_text(text) if not text or max_chars <= 0 or len(text) <= max_chars: return [text] if text else [] units: list[str] = [] buffer: list[str] = [] for char in text: buffer.append(char) if char in "。!?!?;;": unit = "".join(buffer).strip() if unit: units.append(unit) buffer.clear() tail = "".join(buffer).strip() if tail: units.append(tail) chunks: list[str] = [] current = "" for unit in units: pieces = _hard_split_text(unit, max_chars, min_chunk_chars) if len(unit) > max_chars else [unit] for piece in pieces: if current and len(current) + len(piece) > max_chars: chunks.append(current) current = piece else: current = f"{current}{piece}" if current else piece if current: chunks.append(current) return chunks or [text] def punctuation_pause_seconds(text: str, fallback: float = 0.25) -> float: stripped = text.rstrip() if not stripped: return max(0.0, float(fallback)) if stripped[-1] in ",,、": return 0.15 if stripped[-1] in ";;::": return 0.23 if stripped[-1] in "。!?.!?": return 0.35 return max(0.0, float(fallback)) def _peak_rms(audio: np.ndarray) -> tuple[float, float]: audio = np.asarray(audio, dtype=np.float32).reshape(-1) peak = float(np.max(np.abs(audio))) if audio.size else 0.0 rms = float(np.sqrt(np.mean(np.square(audio, dtype=np.float64)))) if audio.size else 0.0 return peak, rms def match_chunk_rms(reference: np.ndarray, chunk: np.ndarray, max_adjust_db: float = 4.0) -> np.ndarray: output = np.asarray(chunk, dtype=np.float32).reshape(-1).copy() reference_peak, reference_rms = _peak_rms(reference) chunk_peak, chunk_rms = _peak_rms(output) del reference_peak if output.size == 0 or reference_rms <= 0.0 or chunk_rms <= 0.0: return output bound = 10.0 ** (max(0.0, float(max_adjust_db)) / 20.0) gain = float(np.clip(reference_rms / chunk_rms, 1.0 / bound, bound)) if chunk_peak > 0.0: gain = min(gain, 0.95 / chunk_peak) output *= gain return output def fade_internal_edges(chunks: list[np.ndarray], sample_rate: int, fade_ms: float = 80.0) -> list[np.ndarray]: outputs = [np.asarray(chunk, dtype=np.float32).reshape(-1).copy() for chunk in chunks] requested = max(0, int(round(float(fade_ms) * sample_rate / 1000.0))) if requested <= 0 or len(outputs) < 2: return outputs for index, output in enumerate(outputs): count = min(requested, output.size) if index > 0: output[:count] *= np.linspace(0.0, 1.0, count, endpoint=True, dtype=np.float32) if index + 1 < len(outputs): output[-count:] *= np.linspace(1.0, 0.0, count, endpoint=True, dtype=np.float32) return outputs def join_audio_chunks( chunks: list[np.ndarray], pauses: list[int], crossfade_samples: int = 0, ) -> np.ndarray: if not chunks: return np.zeros(0, dtype=np.float32) output = np.asarray(chunks[0], dtype=np.float32).copy() for index, chunk in enumerate(chunks[1:]): next_chunk = np.asarray(chunk, dtype=np.float32).copy() pause = max(0, int(pauses[index] if index < len(pauses) else 0)) crossfade = min(max(0, int(crossfade_samples)), output.size, next_chunk.size) if pause > 0: if crossfade > 0: output[-crossfade:] *= np.linspace(1.0, 0.0, crossfade, dtype=np.float32) next_chunk[:crossfade] *= np.linspace(0.0, 1.0, crossfade, dtype=np.float32) output = np.concatenate((output, np.zeros(pause, dtype=np.float32), next_chunk)) elif crossfade > 0: fade_out = np.linspace(1.0, 0.0, crossfade, endpoint=False, dtype=np.float32) overlap = output[-crossfade:] * fade_out + next_chunk[:crossfade] * (1.0 - fade_out) output = np.concatenate((output[:-crossfade], overlap, next_chunk[crossfade:])) else: output = np.concatenate((output, next_chunk)) return output.astype(np.float32, copy=False) def apply_loudness_floor( audio: np.ndarray, min_rms: float = 0.07, peak_limit: float = 0.95, max_gain: float = 3.0, ) -> np.ndarray: output = np.asarray(audio, dtype=np.float32).reshape(-1).copy() peak, rms = _peak_rms(output) if min_rms > 0.0 and 0.0 < rms < min_rms: gain = min(float(max_gain), float(min_rms) / rms) output *= gain peak *= gain if peak_limit > 0.0 and peak > peak_limit: output *= float(peak_limit) / peak return output @torch.no_grad() def extract_windowed_speaker_embedding( wav_path: str, encoder, *, device: str = "cpu", min_duration_seconds: float = 3.0, window_seconds: float = 3.0, hop_seconds: float = 1.5, max_windows: int = 12, full_clip_max_seconds: float = 12.0, ) -> torch.Tensor: """Extract one denoised ECAPA embedding from overlapping reference windows.""" import librosa waveform, _ = librosa.load(wav_path, sr=16000, mono=True) waveform = np.asarray(waveform, dtype=np.float32) duration = waveform.size / 16000.0 if duration < min_duration_seconds: raise ValueError( f"reference audio must be at least {min_duration_seconds:.1f} seconds; got {duration:.2f}" ) segments: list[np.ndarray] = [] if duration <= full_clip_max_seconds: segments.append(waveform) window = max(1, int(round(window_seconds * 16000))) hop = max(1, int(round(hop_seconds * 16000))) starts = list(range(0, max(0, waveform.size - window) + 1, hop)) if max_windows > 0 and len(starts) > max_windows: indices = np.linspace(0, len(starts) - 1, max_windows).round().astype(int) starts = [starts[index] for index in dict.fromkeys(indices.tolist())] segments.extend(waveform[start : start + window] for start in starts) embeddings: list[torch.Tensor] = [] for segment in segments: tensor = torch.from_numpy(np.ascontiguousarray(segment)).float().unsqueeze(0).to(device) embedding = encoder.encode_batch(tensor).reshape(-1) embeddings.append(torch.nn.functional.normalize(embedding, dim=0).cpu()) if not embeddings: raise ValueError("reference audio did not contain a usable speech window") return torch.nn.functional.normalize(torch.stack(embeddings).mean(dim=0), dim=0)