Spaces:
Running on Zero
Running on Zero
Download production.py from voidful/BlueMagpie-TTS-Demo: direct link, hf CLI and curl.
- Browser
- Download file 51.7 kB
-
https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/7e7df2a5f9ea6839a43a3ed880a32d7ecaea9176/production.py
- Command line
-
hf download hf://spaces/voidful/BlueMagpie-TTS-Demo@7e7df2a5f9ea6839a43a3ed880a32d7ecaea9176/production.py
-
curl -L -o production.py https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/7e7df2a5f9ea6839a43a3ed880a32d7ecaea9176/production.py
51.7 kB
| """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)(?<![a-z0-9])(?:po|bo)\s*(?:po|bo)\s*(?:mo|摸)(?![a-z0-9])" | |
| ) | |
| _ASR_HOMOPHONE_TRANSLATION = str.maketrans( | |
| { | |
| "她": "他", | |
| "它": "他", | |
| "牠": "他", | |
| "祂": "他", | |
| "妳": "你", | |
| } | |
| ) | |
| _TERMINAL_PUNCTUATION = frozenset("。!?.!?") | |
| _NONTERMINAL_TRAILING_PUNCTUATION = frozenset(",,、;;::") | |
| _TRAILING_CLOSERS = frozenset("\"'”’」』】))]}") | |
| _BOPOMOFO_READINGS = { | |
| "ㄅ": "波", "ㄆ": "坡", "ㄇ": "摸", "ㄈ": "佛", "ㄉ": "得", | |
| "ㄊ": "特", "ㄋ": "呢", "ㄌ": "了", "ㄍ": "哥", "ㄎ": "科", | |
| "ㄏ": "喝", "ㄐ": "基", "ㄑ": "欺", "ㄒ": "希", "ㄓ": "知", | |
| "ㄔ": "吃", "ㄕ": "師", "ㄖ": "日", "ㄗ": "資", "ㄘ": "雌", | |
| "ㄙ": "思", "ㄚ": "啊", "ㄛ": "喔", "ㄜ": "鵝", "ㄝ": "欸", | |
| "ㄞ": "哀", "ㄟ": "欸", "ㄠ": "凹", "ㄡ": "歐", "ㄢ": "安", | |
| "ㄣ": "恩", "ㄤ": "昂", "ㄥ": "鞥", "ㄦ": "兒", "ㄧ": "衣", | |
| "ㄨ": "烏", "ㄩ": "迂", | |
| } | |
| _ZH_DIGITS = "零一二三四五六七八九" | |
| _ZH_SMALL_UNITS = ("", "十", "百", "千") | |
| _ZH_LARGE_UNITS = ("", "萬", "億", "兆") | |
| _SPOKEN_DATE_RE = re.compile( | |
| r"(?<![A-Za-z0-9/])(\d{4})([/-])(\d{1,2})\2(\d{1,2})(?![A-Za-z0-9/])" | |
| ) | |
| _SPOKEN_ZH_DATE_RE = re.compile( | |
| r"(?<![A-Za-z0-9])(\d{4})年(\d{1,2})月(\d{1,2})(?:日|號)(?![A-Za-z0-9])" | |
| ) | |
| _SPOKEN_TIME_RE = re.compile( | |
| r"(?<![A-Za-z0-9:])([01]?\d|2[0-3]):([0-5]\d)(?![A-Za-z0-9:])" | |
| ) | |
| _SPOKEN_PERCENT_RE = re.compile( | |
| r"(?<![A-Za-z0-9.])([+-]?\d+(?:\.\d+)?)\s*[%%](?![A-Za-z0-9])" | |
| ) | |
| _SPOKEN_UNIT_RE = re.compile( | |
| r"(?<![A-Za-z0-9.])([+-]?\d+(?:\.\d+)?)\s*" | |
| r"(km/h|m/s|kHz|MHz|GHz|Hz|km|cm|mm|kg|mg|mL|ml|kW|°\s*[Cc]|℃|m|g|L|l|W|V)" | |
| r"(?![A-Za-z])" | |
| ) | |
| _SPOKEN_ZH_MEASURE_RE = re.compile( | |
| r"(?<![A-Za-z0-9.])([+-]?\d+(?:\.\d+)?)\s*" | |
| r"(公里|公尺|公分|公厘|厘米|毫米|公斤|千克|公克|毫克|公升|毫升|" | |
| r"攝氏度|千赫茲|兆赫茲|吉赫茲|赫茲|小時|分鐘|米|克|升|度|秒|天|" | |
| r"個|人|次|張|台|件|份|元)" | |
| r"(?![A-Za-z0-9])" | |
| ) | |
| _SPOKEN_URL_RE = re.compile( | |
| r"(?i)(?<![A-Za-z0-9])(?:https?://|ftp://|www\.)[^\s<>\"',。!?;]+" | |
| ) | |
| _SPOKEN_EMAIL_RE = re.compile( | |
| r"(?i)(?<![A-Za-z0-9.!#$%&'*+/=?^_`{|}~-])" | |
| r"[A-Za-z0-9.!#$%&'*+/=?^_`{|}~-]+@" | |
| r"[A-Za-z0-9-]+(?:\.[A-Za-z0-9-]+)+" | |
| r"(?![A-Za-z0-9-])" | |
| ) | |
| _SPOKEN_SEMVER_RE = re.compile( | |
| r"(?i)(?<![A-Za-z0-9.])(?:v|version\s+|版本\s*(?:v\s*)?)?" | |
| r"(\d+)\.(\d+)\.(\d+)" | |
| r"(?:-([0-9A-Za-z]+(?:[.-][0-9A-Za-z]+)*))?" | |
| r"(?:\+([0-9A-Za-z]+(?:[.-][0-9A-Za-z]+)*))?" | |
| r"(?![A-Za-z0-9.])" | |
| ) | |
| _CURRENCY_NUMBER_PATTERN = r"[+-]?(?:\d{1,3}(?:,\d{3})+|\d+)(?:\.\d+)?" | |
| _CURRENCY_CODE_PATTERN = ( | |
| r"NT\$|TWD|NTD|US\$|USD|HK\$|HKD|CN¥|CNY|RMB|JPY|EUR|GBP|KRW|" | |
| r"\$|€|¥|¥|£|₩" | |
| ) | |
| _SPOKEN_CURRENCY_PREFIX_RE = re.compile( | |
| rf"(?i)(?<![A-Za-z0-9])({_CURRENCY_CODE_PATTERN})\s*" | |
| rf"({_CURRENCY_NUMBER_PATTERN})(?![\d,.])" | |
| ) | |
| _SPOKEN_CURRENCY_SUFFIX_RE = re.compile( | |
| rf"(?i)(?<![A-Za-z0-9.])({_CURRENCY_NUMBER_PATTERN})\s*" | |
| rf"(TWD|NTD|USD|HKD|CNY|RMB|JPY|EUR|GBP|KRW)(?![A-Za-z])" | |
| ) | |
| _MODEL_CODE_RE = re.compile( | |
| r"(?<![A-Za-z0-9])([A-Z]{2,8})[ -]*(\d{1,8})(?![A-Za-z0-9.])" | |
| ) | |
| _UPPERCASE_ACRONYM_RE = re.compile(r"(?<![A-Za-z0-9])([A-Z]{2,8})(?![A-Za-z0-9])") | |
| _ZH_UNIT_READINGS = { | |
| "km/h": "公里每小時", | |
| "m/s": "公尺每秒", | |
| "khz": "千赫茲", | |
| "mhz": "兆赫茲", | |
| "ghz": "吉赫茲", | |
| "hz": "赫茲", | |
| "km": "公里", | |
| "cm": "公分", | |
| "mm": "毫米", | |
| "kg": "公斤", | |
| "mg": "毫克", | |
| "ml": "毫升", | |
| "kw": "千瓦", | |
| "°c": "攝氏度", | |
| "℃": "攝氏度", | |
| "m": "公尺", | |
| "g": "公克", | |
| "l": "公升", | |
| "w": "瓦", | |
| "v": "伏特", | |
| } | |
| _ZH_CURRENCY_READINGS = { | |
| "NT$": "新台幣", | |
| "TWD": "新台幣", | |
| "NTD": "新台幣", | |
| "US$": "美元", | |
| "USD": "美元", | |
| "$": "美元", | |
| "HK$": "港幣", | |
| "HKD": "港幣", | |
| "CN¥": "人民幣", | |
| "CNY": "人民幣", | |
| "RMB": "人民幣", | |
| "JPY": "日圓", | |
| "¥": "日圓", | |
| "¥": "日圓", | |
| "EUR": "歐元", | |
| "€": "歐元", | |
| "GBP": "英鎊", | |
| "£": "英鎊", | |
| "KRW": "韓元", | |
| "₩": "韓元", | |
| } | |
| _ZH_NETWORK_SYMBOL_READINGS = { | |
| ".": "點", | |
| "/": "斜線", | |
| ":": "冒號", | |
| "?": "問號", | |
| "=": "等於", | |
| "&": "和", | |
| "#": "井號", | |
| "@": "小老鼠", | |
| "-": "橫線", | |
| "_": "底線", | |
| "%": "百分號", | |
| "+": "加號", | |
| "~": "波浪號", | |
| } | |
| _T2S_CONVERTER = OpenCC("t2s") | |
| 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, | |
| late_start_ratio: float = 0.80, | |
| late_full_ratio: float = 1.00, | |
| ) -> 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 | |
| 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 | |
| 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) | |
| 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) | |