"""Production inference helpers for the BlueMagpie-TTS Space.""" from __future__ import annotations import math import random import re import unicodedata from dataclasses import dataclass from datetime import date from typing import Sequence import numpy as np import torch from opencc import OpenCC 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_ASR_POPO_RE = re.compile( r"(?i)(?\"',。!?;]+" ) _SPOKEN_EMAIL_RE = re.compile( r"(?i)(? 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.late_start_ratio = max(0.0, float(late_start_ratio)) self.late_full_ratio = max( self.late_start_ratio + 1.0e-6, float(late_full_ratio), ) self._active = False self._min_len = 2 self._expected_steps = 0 self._hard_stop_steps = 0 self._step = 0 self._hits = 0 self.last_probabilities: list[float] = [] self.last_generated_steps = 0 self.last_stop_reason = "inactive" 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 self.last_probabilities = [] self.last_generated_steps = 0 self.last_stop_reason = "running" 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()) self.last_probabilities.append(probability) generated_steps = self._step + 1 threshold = self.threshold if self._expected_steps > 0: progress = generated_steps / float(self._expected_steps) if progress >= self.late_start_ratio: blend = min( 1.0, max( 0.0, (progress - self.late_start_ratio) / (self.late_full_ratio - self.late_start_ratio), ), ) 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 stop_reason = "stop_threshold" if should_stop else "running" if self._hard_stop_steps > 0 and generated_steps >= self._hard_stop_steps: should_stop = True stop_reason = "hard_stop" self._step += 1 self.last_generated_steps = generated_steps self.last_stop_reason = stop_reason 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 _zh_four_digit_section(value: int, *, suppress_leading_one: bool = True) -> str: """Read one non-negative, at-most-four-digit integer in Mandarin.""" output: list[str] = [] pending_zero = False for position in range(3, -1, -1): divisor = 10**position digit = value // divisor value %= divisor if digit: if pending_zero and output: output.append(_ZH_DIGITS[0]) if not (digit == 1 and position == 1 and not output and suppress_leading_one): output.append(_ZH_DIGITS[digit]) output.append(_ZH_SMALL_UNITS[position]) pending_zero = False elif output and value: pending_zero = True return "".join(output) def _zh_integer(value: int) -> str: if value == 0: return _ZH_DIGITS[0] if value < 0 or value >= 10**16: raise ValueError("Mandarin integer normalizer supports values from 0 to 10^16 - 1") sections: list[int] = [] while value: sections.append(value % 10000) value //= 10000 output: list[str] = [] pending_zero = False for section_index in range(len(sections) - 1, -1, -1): section = sections[section_index] if section == 0: if output and any(sections[:section_index]): pending_zero = True continue if output and (pending_zero or section < 1000): output.append(_ZH_DIGITS[0]) output.append(_zh_four_digit_section(section, suppress_leading_one=not output)) output.append(_ZH_LARGE_UNITS[section_index]) pending_zero = False return "".join(output) def _zh_number(number: str) -> str: """Read a validated decimal string as a Mandarin cardinal number.""" sign = "" if number.startswith(("+", "-")): sign = "正" if number[0] == "+" else "負" number = number[1:] integer, separator, fraction = number.partition(".") if not integer or not integer.isdigit() or separator and (not fraction or not fraction.isdigit()): raise ValueError("invalid decimal number") if len(integer) > 16: raise ValueError("number is too large for conservative normalization") spoken = _zh_integer(int(integer)) if separator: spoken += "點" + "".join(_ZH_DIGITS[int(digit)] for digit in fraction) return sign + spoken def _zh_digit_sequence(digits: str) -> str: if not digits or not digits.isdigit(): raise ValueError("expected a non-empty digit sequence") return "".join(_ZH_DIGITS[int(digit)] for digit in digits) def _zh_network_text(value: str) -> str: """Spell network identifiers while making separators audible.""" output: list[str] = [] buffer: list[str] = [] def flush() -> None: if not buffer: return token = "".join(buffer) if token.isdigit(): output.append(_zh_digit_sequence(token)) elif token.lower() == "www" or len(token) > 1 and token.isupper(): output.extend(token.upper()) else: output.append(token) buffer.clear() for character in value: if character.isascii() and character.isalnum(): if buffer and character.isdigit() != buffer[-1].isdigit(): flush() buffer.append(character) continue flush() reading = _ZH_NETWORK_SYMBOL_READINGS.get(character) if reading: output.append(reading) elif not character.isspace(): output.append(character) flush() return " ".join(output) def _zh_url(value: str) -> str: scheme, separator, remainder = value.partition("://") if separator: protocol = " ".join(scheme.upper()) return f"{protocol} 冒號 斜線 斜線 {_zh_network_text(remainder)}" return _zh_network_text(value) def _zh_email(value: str) -> str: local, domain = value.rsplit("@", 1) return f"{_zh_network_text(local)} 小老鼠 {_zh_network_text(domain)}" def _zh_semantic_version(match: re.Match[str]) -> str: major, minor, patch, prerelease, build = match.groups() try: core = "點".join(_zh_integer(int(part)) for part in (major, minor, patch)) except ValueError: return match.group(0) output = f"版本{core}" if prerelease: output += f" 預發布 {_zh_network_text(prerelease)}" if build: output += f" 建置 {_zh_network_text(build)}" return output def _zh_currency(code: str, number: str, original: str) -> str: reading = _ZH_CURRENCY_READINGS.get(code.upper(), _ZH_CURRENCY_READINGS.get(code)) if reading is None: return original try: spoken_number = _zh_number(number.replace(",", "")) except ValueError: return original return reading + spoken_number def normalize_spoken_forms(text: str, *, locale: str = "zh-TW") -> str: """Conservatively expand common written forms for TTS. Chinese locales expand validated dates, network addresses, semantic versions, currencies, percentages, numeric units, and standalone all-uppercase acronyms/model identifiers. Ordinary English words are deliberately left untouched. English locales only receive the existing Unicode/punctuation normalization. Unknown locales raise instead of silently applying the wrong pronunciation rules. """ raw = unicodedata.normalize("NFC", str(text or "")) locale_key = str(locale or "").replace("_", "-").lower() if locale_key in {"en", "en-us", "en-gb"}: return normalize_tts_text(raw) if locale_key not in {"zh", "zh-tw", "zh-hant", "zh-cn", "zh-hans"}: raise ValueError(f"unsupported spoken-form locale: {locale!r}") protected: list[str] = [] def protect(spoken: str) -> str: marker = f"\uf000{len(protected)}\uf001" protected.append(spoken) return marker def replace_url(match: re.Match[str]) -> str: matched = match.group(0) core = matched.rstrip(".,!?;:") trailing = matched[len(core) :] return protect(_zh_url(core)) + trailing def replace_email(match: re.Match[str]) -> str: return protect(_zh_email(match.group(0))) def replace_date(match: re.Match[str]) -> str: year_text, _, month_text, day_text = match.groups() try: date(int(year_text), int(month_text), int(day_text)) except ValueError: return match.group(0) return ( f"{_zh_digit_sequence(year_text)}年" f"{_zh_integer(int(month_text))}月{_zh_integer(int(day_text))}日" ) def replace_zh_date(match: re.Match[str]) -> str: year_text, month_text, day_text = match.groups() try: date(int(year_text), int(month_text), int(day_text)) except ValueError: return match.group(0) return ( f"{_zh_digit_sequence(year_text)}年" f"{_zh_integer(int(month_text))}月{_zh_integer(int(day_text))}日" ) def replace_number(match: re.Match[str], suffix: str) -> str: try: return _zh_number(match.group(1)) + suffix except ValueError: return match.group(0) def replace_time(match: re.Match[str]) -> str: hour = int(match.group(1)) minute = int(match.group(2)) minute_text = "整" if minute == 0 else f"{_zh_integer(minute)}分" return f"{_zh_integer(hour)}點{minute_text}" def replace_percent(match: re.Match[str]) -> str: try: return "百分之" + _zh_number(match.group(1)) except ValueError: return match.group(0) def replace_unit(match: re.Match[str]) -> str: unit_key = re.sub(r"\s+", "", match.group(2)).lower() reading = _ZH_UNIT_READINGS.get(unit_key) if reading is None: return match.group(0) return replace_number(match, reading) def replace_model_code(match: re.Match[str]) -> str: letters, digits = match.groups() return f"{' '.join(letters)} {_zh_digit_sequence(digits)}" output = _SPOKEN_URL_RE.sub(replace_url, raw) output = _SPOKEN_EMAIL_RE.sub(replace_email, output) output = _SPOKEN_ZH_DATE_RE.sub(replace_zh_date, output) output = _SPOKEN_DATE_RE.sub(replace_date, output) output = _SPOKEN_TIME_RE.sub(replace_time, output) output = _SPOKEN_PERCENT_RE.sub(replace_percent, output) output = _SPOKEN_UNIT_RE.sub(replace_unit, output) output = _SPOKEN_ZH_MEASURE_RE.sub( lambda match: replace_number(match, match.group(2)), output, ) output = _SPOKEN_SEMVER_RE.sub(_zh_semantic_version, output) output = _SPOKEN_CURRENCY_PREFIX_RE.sub( lambda match: _zh_currency(match.group(1), match.group(2), match.group(0)), output, ) output = _SPOKEN_CURRENCY_SUFFIX_RE.sub( lambda match: _zh_currency(match.group(2), match.group(1), match.group(0)), output, ) output = _MODEL_CODE_RE.sub(replace_model_code, output) output = _UPPERCASE_ACRONYM_RE.sub(lambda match: " ".join(match.group(1)), output) for index, spoken in enumerate(protected): output = output.replace(f"\uf000{index}\uf001", spoken) return normalize_tts_text(output) 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 {":", ":"}: ascii_context = ( previous.isascii() and previous.isalnum() and following.isascii() and following.isalnum() ) output.append(":" if following == "/" or ascii_context 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 normalize_tts_eval_text(text: str) -> str: """Normalize only verified ASR-equivalent forms for semantic scoring. This remains separate from model-input normalization: Mandarin pronoun homophones are acoustically indistinguishable, but their written forms must stay untouched in the text sent to the TTS model. """ normalized = normalize_tts_text(text) normalized = normalized.translate(_ASR_HOMOPHONE_TRANSLATION) if "注音" in normalized or "符號" in normalized: normalized = _BOPOMOFO_ASR_POPO_RE.sub("波坡摸", normalized) return normalized def ensure_terminal_punctuation(text: str) -> str: """Give the acoustic model an explicit endpoint cue without changing words.""" normalized = normalize_tts_text(text) if not normalized: return "" split_at = len(normalized) while split_at > 0 and normalized[split_at - 1] in _TRAILING_CLOSERS: split_at -= 1 core = normalized[:split_at].rstrip() suffix = normalized[split_at:] if core and core[-1] in _TERMINAL_PUNCTUATION: return normalized while core and core[-1] in _NONTERMINAL_TRAILING_PUNCTUATION: core = core[:-1].rstrip() if not core: return normalized punctuation = "。" if any(_is_cjk(char) for char in core) else "." return f"{core}{punctuation}{suffix}" 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 effective_generation_cfg( text: str, requested_cfg: float, *, short_text_unit_threshold: int = 6, short_text_min_cfg: float = 3.0, ) -> float: """Use stronger acoustic guidance only for empirically unstable short text.""" cfg = float(requested_cfg) units = count_speech_units(text) if 0 < units < max(1, int(short_text_unit_threshold)): return max(cfg, float(short_text_min_cfg)) return cfg 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 endpoint_generation_plan( text: str, *, generation_cps: float, step_seconds: float | None, margin_steps: int = 1, add_terminal_punctuation: bool = True, ) -> tuple[str, int, int]: """Plan a native-pace generation cap independently of output playback pace.""" model_text = ( ensure_terminal_punctuation(text) if add_terminal_punctuation else normalize_tts_text(text) ) expected_steps = target_cps_steps(model_text, generation_cps, step_seconds) hard_stop_steps = duration_hard_stop_steps( expected_steps, ratio=1.0, margin_steps=margin_steps, ) return model_text, expected_steps, hard_stop_steps def select_generation_cps( text: str, *, cjk_cps: float = 5.2, ascii_cps: float = 4.6, ) -> float: """Reserve more generation time for ASCII words than compact CJK units.""" normalized = normalize_tts_text(text) if any(char.isascii() and char.isalnum() for char in normalized): return float(ascii_cps) return float(cjk_cps) 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 split_leading_clause( text: str, *, search_chars: int = 40, min_chunk_chars: int = 12, ) -> list[str]: """Split the onset only at a real punctuation boundary. A hard onset split gives the stop head a text fragment with no endpoint cue. That can turn the artificial chunk tail into extra speech. Leave the utterance intact when no suitable punctuation exists. """ normalized = normalize_tts_text(text) if not normalized: return [] minimum = max(1, int(min_chunk_chars)) upper = min(max(0, int(search_chars)), len(normalized) - minimum) if upper < minimum: return [normalized] for index, char in enumerate(normalized[:upper], start=1): if index >= minimum and char in ",,、;;::。!?!?": return [normalized[:index], normalized[index:]] return [normalized] 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 def _comparison_text(text: str, *, locale: str) -> str: normalized = normalize_tts_eval_text( normalize_spoken_forms(text, locale=locale) ).casefold() locale_key = str(locale or "").replace("_", "-").lower() if locale_key in {"zh", "zh-tw", "zh-hant", "zh-cn", "zh-hans"}: normalized = _T2S_CONVERTER.convert(normalized) return "".join(char for char in normalized if char.isalnum()) def _levenshtein_alignment(source: str, hypothesis: str) -> list[tuple[str, int, int]]: """Return deterministic edit operations as ``(op, source_index, hyp_index)``.""" rows = len(source) + 1 columns = len(hypothesis) + 1 distance = [[0] * columns for _ in range(rows)] for source_index in range(rows): distance[source_index][0] = source_index for hyp_index in range(columns): distance[0][hyp_index] = hyp_index for source_index in range(1, rows): for hyp_index in range(1, columns): substitution = distance[source_index - 1][hyp_index - 1] + ( source[source_index - 1] != hypothesis[hyp_index - 1] ) deletion = distance[source_index - 1][hyp_index] + 1 insertion = distance[source_index][hyp_index - 1] + 1 distance[source_index][hyp_index] = min(substitution, deletion, insertion) operations: list[tuple[str, int, int]] = [] source_index = len(source) hyp_index = len(hypothesis) while source_index or hyp_index: if source_index and hyp_index: cost = source[source_index - 1] != hypothesis[hyp_index - 1] if distance[source_index][hyp_index] == distance[source_index - 1][hyp_index - 1] + cost: operations.append( ( "substitute" if cost else "equal", source_index - 1, hyp_index - 1, ) ) source_index -= 1 hyp_index -= 1 continue if source_index and distance[source_index][hyp_index] == distance[source_index - 1][hyp_index] + 1: operations.append(("delete", source_index - 1, hyp_index)) source_index -= 1 continue operations.append(("insert", source_index, hyp_index - 1)) hyp_index -= 1 operations.reverse() return operations @dataclass(frozen=True) class AsrComparison: """Character-level semantic completion evidence for one TTS candidate.""" target_text: str transcript_text: str edit_distance: int cer: float prefix_cer: float suffix_cer: float prefix_deletions: int suffix_deletions: int extra_tail: str extra_tail_units: int passed: bool def compare_asr_text( target: str, transcript: str, *, locale: str = "zh-TW", prefix_units: int = 6, suffix_units: int = 6, max_cer: float = 0.10, max_prefix_cer: float = 0.0, max_suffix_cer: float = 0.0, max_extra_tail_units: int = 0, ) -> AsrComparison: """Compare an ASR transcript with strict onset, completion, and tail gates. Punctuation and whitespace are ignored, while written dates/numbers are normalized before alignment. Empty targets, unsupported locales, invalid limits, and non-finite limits all fail closed via ``passed=False``. """ try: target_text = _comparison_text(target, locale=locale) transcript_text = _comparison_text(transcript, locale=locale) except (TypeError, ValueError): target_text = "" transcript_text = "" try: prefix_count = int(prefix_units) suffix_count = int(suffix_units) except (TypeError, ValueError, OverflowError): prefix_count = 0 suffix_count = 0 prefix_limit = min(max(0, prefix_count), len(target_text)) suffix_start = max(0, len(target_text) - max(0, suffix_count)) operations = _levenshtein_alignment(target_text, transcript_text) edit_distance = sum(operation != "equal" for operation, _, _ in operations) cer = edit_distance / len(target_text) if target_text else math.inf prefix_errors = 0 suffix_errors = 0 prefix_deletions = 0 suffix_deletions = 0 tail_characters: list[str] = [] for operation, source_index, hyp_index in operations: if operation == "equal": continue if operation == "insert": if source_index < prefix_limit: prefix_errors += 1 if source_index >= suffix_start: suffix_errors += 1 if source_index >= len(target_text) and 0 <= hyp_index < len(transcript_text): tail_characters.append(transcript_text[hyp_index]) continue if source_index < prefix_limit: prefix_errors += 1 prefix_deletions += operation == "delete" if source_index >= suffix_start: suffix_errors += 1 suffix_deletions += operation == "delete" prefix_cer = prefix_errors / prefix_limit if prefix_limit else math.inf suffix_length = len(target_text) - suffix_start suffix_cer = suffix_errors / suffix_length if suffix_length else math.inf extra_tail = "".join(tail_characters) limits: list[float] = [] for value in (max_cer, max_prefix_cer, max_suffix_cer, max_extra_tail_units): try: limits.append(float(value)) except (TypeError, ValueError, OverflowError): limits.append(math.nan) valid_limits = all(math.isfinite(value) and value >= 0.0 for value in limits) passed = bool( target_text and transcript_text and prefix_limit > 0 and suffix_length > 0 and valid_limits and cer <= limits[0] and prefix_cer <= limits[1] and suffix_cer <= limits[2] and len(extra_tail) <= limits[3] ) return AsrComparison( target_text=target_text, transcript_text=transcript_text, edit_distance=edit_distance, cer=cer, prefix_cer=prefix_cer, suffix_cer=suffix_cer, prefix_deletions=prefix_deletions, suffix_deletions=suffix_deletions, extra_tail=extra_tail, extra_tail_units=len(extra_tail), passed=passed, ) def _finite_number(value, *, minimum: float | None = None, maximum: float | None = None) -> float | None: if isinstance(value, (bool, np.bool_)): return None try: number = float(value) except (TypeError, ValueError, OverflowError): return None if not math.isfinite(number): return None if minimum is not None and number < minimum: return None if maximum is not None and number > maximum: return None return number def candidate_local_score( *, cer: float, speaker_similarity: float, boundary_speaker_drop: float, prefix_cer: float = 0.0, suffix_cer: float = 0.0, pace_penalty: float = 0.0, style_penalty: float = 0.0, asr_passed: bool = True, truncated: bool = False, extra_tail: bool = False, speaker_weight: float = 0.05, boundary_weight: float = 0.10, prefix_weight: float = 0.50, suffix_weight: float = 0.50, pace_weight: float = 1.0, style_weight: float = 1.0, max_cer: float | None = None, min_speaker_similarity: float | None = None, max_boundary_speaker_drop: float | None = None, ) -> float: """Return a finite local candidate cost or ``inf`` for unsafe evidence.""" if asr_passed is not True or truncated is not False or extra_tail is not False: return math.inf metric_values = [ _finite_number(cer, minimum=0.0), _finite_number(speaker_similarity, minimum=-1.0, maximum=1.0), _finite_number(boundary_speaker_drop, minimum=0.0), _finite_number(prefix_cer, minimum=0.0), _finite_number(suffix_cer, minimum=0.0), _finite_number(pace_penalty, minimum=0.0), _finite_number(style_penalty, minimum=0.0), ] weights = [ _finite_number(speaker_weight, minimum=0.0), _finite_number(boundary_weight, minimum=0.0), _finite_number(prefix_weight, minimum=0.0), _finite_number(suffix_weight, minimum=0.0), _finite_number(pace_weight, minimum=0.0), _finite_number(style_weight, minimum=0.0), ] if any(value is None for value in metric_values + weights): return math.inf candidate_cer, similarity, boundary_drop, onset_cer, ending_cer, pace, style = metric_values gate_max_cer = None if max_cer is None else _finite_number(max_cer, minimum=0.0) gate_min_similarity = ( None if min_speaker_similarity is None else _finite_number(min_speaker_similarity, minimum=-1.0, maximum=1.0) ) gate_max_boundary = ( None if max_boundary_speaker_drop is None else _finite_number(max_boundary_speaker_drop, minimum=0.0) ) if ( max_cer is not None and gate_max_cer is None or min_speaker_similarity is not None and gate_min_similarity is None or max_boundary_speaker_drop is not None and gate_max_boundary is None ): return math.inf if gate_max_cer is not None and candidate_cer > gate_max_cer: return math.inf if gate_min_similarity is not None and similarity < gate_min_similarity: return math.inf if gate_max_boundary is not None and boundary_drop > gate_max_boundary: return math.inf speaker_w, boundary_w, prefix_w, suffix_w, pace_w, style_w = weights score = ( candidate_cer + speaker_w * (1.0 - similarity) + boundary_w * boundary_drop + prefix_w * onset_cer + suffix_w * ending_cer + pace_w * pace + style_w * style ) return score if math.isfinite(score) and score >= 0.0 else math.inf def candidate_transition_score( *, speaker_similarity: float, f0_delta: float, rms_delta: float, speaker_weight: float = 1.0, f0_weight: float = 1.0, rms_weight: float = 1.0, ) -> float: """Score continuity between adjacent candidates with finite-only inputs.""" values = [ _finite_number(speaker_similarity, minimum=-1.0, maximum=1.0), _finite_number(f0_delta, minimum=0.0), _finite_number(rms_delta, minimum=0.0), _finite_number(speaker_weight, minimum=0.0), _finite_number(f0_weight, minimum=0.0), _finite_number(rms_weight, minimum=0.0), ] if any(value is None for value in values): return math.inf similarity, pitch_delta, loudness_delta, speaker_w, pitch_w, loudness_w = values score = ( speaker_w * (1.0 - similarity) + pitch_w * pitch_delta + loudness_w * loudness_delta ) return score if math.isfinite(score) and score >= 0.0 else math.inf @dataclass(frozen=True) class CandidateSequenceSelection: candidate_indices: tuple[int, ...] total_score: float def select_candidate_sequence( local_scores: Sequence[Sequence[float]], transition_scores: Sequence[Sequence[Sequence[float]]] = (), ) -> CandidateSequenceSelection | None: """Select the minimum-cost candidate path in ``O(N K^2)``. Non-finite/negative scores remove only their candidate or edge. Malformed matrices and graphs with no complete finite path return ``None`` so callers cannot accidentally fall back to an unverified candidate. """ try: local_rows = [list(row) for row in local_scores] except TypeError: return None if not local_rows or any(not row for row in local_rows): return None safe_local = [ [ score if score is not None else math.inf for score in (_finite_number(value, minimum=0.0) for value in row) ] for row in local_rows ] if len(safe_local) == 1: try: if len(transition_scores) != 0: return None except TypeError: return None best_index = min(range(len(safe_local[0])), key=safe_local[0].__getitem__) best_score = safe_local[0][best_index] if not math.isfinite(best_score): return None return CandidateSequenceSelection((best_index,), best_score) try: transitions = [[list(row) for row in matrix] for matrix in transition_scores] except TypeError: return None if len(transitions) != len(safe_local) - 1: return None safe_transitions: list[list[list[float]]] = [] for index, matrix in enumerate(transitions): previous_count = len(safe_local[index]) current_count = len(safe_local[index + 1]) if len(matrix) != previous_count or any(len(row) != current_count for row in matrix): return None safe_transitions.append( [ [ score if score is not None else math.inf for score in (_finite_number(value, minimum=0.0) for value in row) ] for row in matrix ] ) previous_costs = safe_local[0] backpointers: list[list[int]] = [] for step in range(1, len(safe_local)): current_costs = [math.inf] * len(safe_local[step]) current_backpointers = [-1] * len(safe_local[step]) for current_index, local_score in enumerate(safe_local[step]): if not math.isfinite(local_score): continue for previous_index, previous_score in enumerate(previous_costs): edge_score = safe_transitions[step - 1][previous_index][current_index] if not math.isfinite(previous_score) or not math.isfinite(edge_score): continue total = previous_score + edge_score + local_score if math.isfinite(total) and total < current_costs[current_index]: current_costs[current_index] = total current_backpointers[current_index] = previous_index previous_costs = current_costs backpointers.append(current_backpointers) final_index = min(range(len(previous_costs)), key=previous_costs.__getitem__) total_score = previous_costs[final_index] if not math.isfinite(total_score): return None indices = [final_index] for pointers in reversed(backpointers): final_index = pointers[final_index] if final_index < 0: return None indices.append(final_index) indices.reverse() return CandidateSequenceSelection(tuple(indices), total_score) @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)