"""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_clock_time(hour_text: str, minute_text: str) -> str: """Return the canonical Mandarin reading for a validated clock time.""" hour = int(hour_text) minute = int(minute_text) minute_spoken = "整" if minute == 0 else f"{_zh_integer(minute)}分" return f"{_zh_integer(hour)}點{minute_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 _nonspace_neighbor(text: str, index: int, step: int) -> str | None: """Return the nearest non-space character on one side of ``index``.""" position = int(index) while 0 <= position < len(text): character = text[position] if not character.isspace(): return character position += int(step) return None def _semantic_neighbor(text: str, index: int, step: int) -> str | None: """Return a simplified alphanumeric/CJK context anchor.""" position = int(index) while 0 <= position < len(text): character = text[position] if character.isalnum() or _is_cjk(character): return _T2S_CONVERTER.convert(character).casefold() position += int(step) return None def _structured_numeric_neighbor(character: str | None) -> bool: if character is None: return False return bool( character in _ASR_STRUCTURED_NUMERIC_NEIGHBORS or character.isascii() and character.isalnum() ) def _target_numeric_proofs( normalized_target: str, reading: str, ) -> list[tuple[int, int]]: """Find unstructured, exact target occurrences of one derived reading.""" proofs: list[tuple[int, int]] = [] start = 0 while True: index = normalized_target.find(reading, start) if index < 0: return proofs end = index + len(reading) left = _nonspace_neighbor(normalized_target, index - 1, -1) right = _nonspace_neighbor(normalized_target, end, 1) if not ( _structured_numeric_neighbor(left) or _structured_numeric_neighbor(right) ): proofs.append((index, end)) start = index + 1 def _canonicalize_target_proven_arabic_digits( normalized_text: str, normalized_target: str, ) -> str: """Replace a bare ASR digit run only with its exact target-proven reading. Both possible Mandarin readings are derived from the run itself: a cardinal value (``1250`` -> ``一千二百五十``) and a digit sequence (``1250`` -> ``一二五零``). A derived reading must occur in the target at matching semantic left/right anchors. Structured numeric domains are deliberately excluded because dates, clocks, versions, network values, and model codes have their own stricter normalizers. """ matches = tuple(_ASR_ARABIC_DIGIT_RUN_RE.finditer(normalized_text)) if not matches: return normalized_text used_proofs: set[tuple[int, int, str]] = set() output: list[str] = [] cursor = 0 for match in matches: output.append(normalized_text[cursor : match.start()]) run = match.group(0) left = _nonspace_neighbor(normalized_text, match.start() - 1, -1) right = _nonspace_neighbor(normalized_text, match.end(), 1) replacement: str | None = None if not ( _structured_numeric_neighbor(left) or _structured_numeric_neighbor(right) ): readings: list[str] = [] try: readings.append(_zh_integer(int(run))) except ValueError: pass digit_reading = _zh_digit_sequence(run) if digit_reading not in readings: readings.append(digit_reading) text_context = ( _semantic_neighbor(normalized_text, match.start() - 1, -1), _semantic_neighbor(normalized_text, match.end(), 1), ) for reading in readings: for proof_start, proof_end in _target_numeric_proofs( normalized_target, reading, ): proof_key = (proof_start, proof_end, reading) if proof_key in used_proofs: continue target_context = ( _semantic_neighbor( normalized_target, proof_start - 1, -1, ), _semantic_neighbor(normalized_target, proof_end, 1), ) if target_context == text_context: replacement = reading used_proofs.add(proof_key) break if replacement is not None: break output.append(replacement if replacement is not None else run) cursor = match.end() output.append(normalized_text[cursor:]) return "".join(output) 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: return _zh_clock_time(match.group(1), match.group(2)) 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_asr_spoken_forms( text: str, target_text: str, *, locale: str = "zh-TW", ) -> str: """Canonicalize only target-proven ASR numeric readings.""" normalized_target = normalize_spoken_forms(target_text, locale=locale) normalized_text = normalize_spoken_forms(text, locale=locale) locale_key = str(locale or "").replace("_", "-").lower() if locale_key not in {"zh", "zh-tw", "zh-hant", "zh-cn", "zh-hans"}: return normalized_text normalized_text = _canonicalize_target_proven_arabic_digits( normalized_text, normalized_target, ) # Only an unambiguous target reading may prove that Whisper's numeric # ``15點30(分)`` rendering means a clock time. In particular, do not # rewrite the target's own numeric form first: that would let ambiguous # decimals such as ``15點05分貝`` or scores self-authorize a clock # canonicalization. Colon times have already been expanded by # ``normalize_spoken_forms`` and fully-spoken clock targets already contain # the canonical phrase, so both remain valid proof sources. target_clock_text = normalized_target.translate( str.maketrans({"点": "點", "時": "點", "时": "點"}) ) if ( _ASR_EXPLICIT_ZH_TIME_RE.search(normalized_target) is not None or _ASR_BARE_ZH_TIME_RE.search(normalized_target) is not None ): # A numeric ``X點Y`` construction in the target is itself ambiguous. # Do not let a separate, fully-spoken clock elsewhere in the sentence # globally authorize rewriting that decimal, duration, or score. return normalized_text def replace_target_proven_time(match: re.Match[str]) -> str: canonical = _zh_clock_time(match.group(1), match.group(2)) return canonical if canonical in target_clock_text else match.group(0) normalized_text = _ASR_EXPLICIT_ZH_TIME_RE.sub( replace_target_proven_time, normalized_text, ) normalized_text = _ASR_BARE_ZH_TIME_RE.sub( replace_target_proven_time, normalized_text, ) return normalize_tts_text(normalized_text) 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 active_pace_correction_speed( active_duration_seconds: float, text: str, *, target_cps: float, prior_speed: float = 1.0, min_total_speed: float = 0.80, ) -> float: """Return a second stretch rate from active-voice duration. ``prior_speed`` is the rate already applied by ``target_pace_speed``. Successive pitch-preserving stretch rates multiply, so the second rate is floored at ``min_total_speed / prior_speed``. Invalid evidence and audio that is already at or below the target active CPS fail safely to ``1.0``; this helper never speeds audio up. """ units = count_speech_units(text) try: active_seconds = float(active_duration_seconds) target = float(target_cps) previous_rate = float(prior_speed) total_floor = float(min_total_speed) except (TypeError, ValueError, OverflowError): return 1.0 if ( units <= 0 or not math.isfinite(active_seconds) or active_seconds <= 0.0 or not math.isfinite(target) or target <= 0.0 or not math.isfinite(previous_rate) or not 0.0 < previous_rate <= 1.0 or not math.isfinite(total_floor) or not 0.0 < total_floor <= 1.0 ): return 1.0 desired_rate = active_seconds * target / float(units) if not math.isfinite(desired_rate) or desired_rate >= 1.0: return 1.0 remaining_floor = total_floor / previous_rate if remaining_floor >= 1.0: return 1.0 return min(1.0, max(remaining_floor, desired_rate)) 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, target_text: str | None = None, ) -> str: spoken = ( normalize_asr_spoken_forms(text, target_text, locale=locale) if target_text is not None else normalize_spoken_forms(text, locale=locale) ) normalized = normalize_tts_eval_text( spoken ).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, target_text=target, ) transcript_text = _comparison_text( transcript, locale=locale, target_text=target, ) 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)