BlueMagpie-TTS-Demo / production.py
voidful's picture
Add fail-closed stable speaker inference
7e7df2a
Raw History Blame
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
@dataclass(frozen=True)
class AsrComparison:
"""Character-level semantic completion evidence for one TTS candidate."""
target_text: str
transcript_text: str
edit_distance: int
cer: float
prefix_cer: float
suffix_cer: float
prefix_deletions: int
suffix_deletions: int
extra_tail: str
extra_tail_units: int
passed: bool
def compare_asr_text(
target: str,
transcript: str,
*,
locale: str = "zh-TW",
prefix_units: int = 6,
suffix_units: int = 6,
max_cer: float = 0.10,
max_prefix_cer: float = 0.0,
max_suffix_cer: float = 0.0,
max_extra_tail_units: int = 0,
) -> AsrComparison:
"""Compare an ASR transcript with strict onset, completion, and tail gates.
Punctuation and whitespace are ignored, while written dates/numbers are
normalized before alignment. Empty targets, unsupported locales, invalid
limits, and non-finite limits all fail closed via ``passed=False``.
"""
try:
target_text = _comparison_text(target, locale=locale)
transcript_text = _comparison_text(transcript, locale=locale)
except (TypeError, ValueError):
target_text = ""
transcript_text = ""
try:
prefix_count = int(prefix_units)
suffix_count = int(suffix_units)
except (TypeError, ValueError, OverflowError):
prefix_count = 0
suffix_count = 0
prefix_limit = min(max(0, prefix_count), len(target_text))
suffix_start = max(0, len(target_text) - max(0, suffix_count))
operations = _levenshtein_alignment(target_text, transcript_text)
edit_distance = sum(operation != "equal" for operation, _, _ in operations)
cer = edit_distance / len(target_text) if target_text else math.inf
prefix_errors = 0
suffix_errors = 0
prefix_deletions = 0
suffix_deletions = 0
tail_characters: list[str] = []
for operation, source_index, hyp_index in operations:
if operation == "equal":
continue
if operation == "insert":
if source_index < prefix_limit:
prefix_errors += 1
if source_index >= suffix_start:
suffix_errors += 1
if source_index >= len(target_text) and 0 <= hyp_index < len(transcript_text):
tail_characters.append(transcript_text[hyp_index])
continue
if source_index < prefix_limit:
prefix_errors += 1
prefix_deletions += operation == "delete"
if source_index >= suffix_start:
suffix_errors += 1
suffix_deletions += operation == "delete"
prefix_cer = prefix_errors / prefix_limit if prefix_limit else math.inf
suffix_length = len(target_text) - suffix_start
suffix_cer = suffix_errors / suffix_length if suffix_length else math.inf
extra_tail = "".join(tail_characters)
limits: list[float] = []
for value in (max_cer, max_prefix_cer, max_suffix_cer, max_extra_tail_units):
try:
limits.append(float(value))
except (TypeError, ValueError, OverflowError):
limits.append(math.nan)
valid_limits = all(math.isfinite(value) and value >= 0.0 for value in limits)
passed = bool(
target_text
and transcript_text
and prefix_limit > 0
and suffix_length > 0
and valid_limits
and cer <= limits[0]
and prefix_cer <= limits[1]
and suffix_cer <= limits[2]
and len(extra_tail) <= limits[3]
)
return AsrComparison(
target_text=target_text,
transcript_text=transcript_text,
edit_distance=edit_distance,
cer=cer,
prefix_cer=prefix_cer,
suffix_cer=suffix_cer,
prefix_deletions=prefix_deletions,
suffix_deletions=suffix_deletions,
extra_tail=extra_tail,
extra_tail_units=len(extra_tail),
passed=passed,
)
def _finite_number(value, *, minimum: float | None = None, maximum: float | None = None) -> float | None:
if isinstance(value, (bool, np.bool_)):
return None
try:
number = float(value)
except (TypeError, ValueError, OverflowError):
return None
if not math.isfinite(number):
return None
if minimum is not None and number < minimum:
return None
if maximum is not None and number > maximum:
return None
return number
def candidate_local_score(
*,
cer: float,
speaker_similarity: float,
boundary_speaker_drop: float,
prefix_cer: float = 0.0,
suffix_cer: float = 0.0,
pace_penalty: float = 0.0,
style_penalty: float = 0.0,
asr_passed: bool = True,
truncated: bool = False,
extra_tail: bool = False,
speaker_weight: float = 0.05,
boundary_weight: float = 0.10,
prefix_weight: float = 0.50,
suffix_weight: float = 0.50,
pace_weight: float = 1.0,
style_weight: float = 1.0,
max_cer: float | None = None,
min_speaker_similarity: float | None = None,
max_boundary_speaker_drop: float | None = None,
) -> float:
"""Return a finite local candidate cost or ``inf`` for unsafe evidence."""
if asr_passed is not True or truncated is not False or extra_tail is not False:
return math.inf
metric_values = [
_finite_number(cer, minimum=0.0),
_finite_number(speaker_similarity, minimum=-1.0, maximum=1.0),
_finite_number(boundary_speaker_drop, minimum=0.0),
_finite_number(prefix_cer, minimum=0.0),
_finite_number(suffix_cer, minimum=0.0),
_finite_number(pace_penalty, minimum=0.0),
_finite_number(style_penalty, minimum=0.0),
]
weights = [
_finite_number(speaker_weight, minimum=0.0),
_finite_number(boundary_weight, minimum=0.0),
_finite_number(prefix_weight, minimum=0.0),
_finite_number(suffix_weight, minimum=0.0),
_finite_number(pace_weight, minimum=0.0),
_finite_number(style_weight, minimum=0.0),
]
if any(value is None for value in metric_values + weights):
return math.inf
candidate_cer, similarity, boundary_drop, onset_cer, ending_cer, pace, style = metric_values
gate_max_cer = None if max_cer is None else _finite_number(max_cer, minimum=0.0)
gate_min_similarity = (
None
if min_speaker_similarity is None
else _finite_number(min_speaker_similarity, minimum=-1.0, maximum=1.0)
)
gate_max_boundary = (
None
if max_boundary_speaker_drop is None
else _finite_number(max_boundary_speaker_drop, minimum=0.0)
)
if (
max_cer is not None and gate_max_cer is None
or min_speaker_similarity is not None and gate_min_similarity is None
or max_boundary_speaker_drop is not None and gate_max_boundary is None
):
return math.inf
if gate_max_cer is not None and candidate_cer > gate_max_cer:
return math.inf
if gate_min_similarity is not None and similarity < gate_min_similarity:
return math.inf
if gate_max_boundary is not None and boundary_drop > gate_max_boundary:
return math.inf
speaker_w, boundary_w, prefix_w, suffix_w, pace_w, style_w = weights
score = (
candidate_cer
+ speaker_w * (1.0 - similarity)
+ boundary_w * boundary_drop
+ prefix_w * onset_cer
+ suffix_w * ending_cer
+ pace_w * pace
+ style_w * style
)
return score if math.isfinite(score) and score >= 0.0 else math.inf
def candidate_transition_score(
*,
speaker_similarity: float,
f0_delta: float,
rms_delta: float,
speaker_weight: float = 1.0,
f0_weight: float = 1.0,
rms_weight: float = 1.0,
) -> float:
"""Score continuity between adjacent candidates with finite-only inputs."""
values = [
_finite_number(speaker_similarity, minimum=-1.0, maximum=1.0),
_finite_number(f0_delta, minimum=0.0),
_finite_number(rms_delta, minimum=0.0),
_finite_number(speaker_weight, minimum=0.0),
_finite_number(f0_weight, minimum=0.0),
_finite_number(rms_weight, minimum=0.0),
]
if any(value is None for value in values):
return math.inf
similarity, pitch_delta, loudness_delta, speaker_w, pitch_w, loudness_w = values
score = (
speaker_w * (1.0 - similarity)
+ pitch_w * pitch_delta
+ loudness_w * loudness_delta
)
return score if math.isfinite(score) and score >= 0.0 else math.inf
@dataclass(frozen=True)
class CandidateSequenceSelection:
candidate_indices: tuple[int, ...]
total_score: float
def select_candidate_sequence(
local_scores: Sequence[Sequence[float]],
transition_scores: Sequence[Sequence[Sequence[float]]] = (),
) -> CandidateSequenceSelection | None:
"""Select the minimum-cost candidate path in ``O(N K^2)``.
Non-finite/negative scores remove only their candidate or edge. Malformed
matrices and graphs with no complete finite path return ``None`` so callers
cannot accidentally fall back to an unverified candidate.
"""
try:
local_rows = [list(row) for row in local_scores]
except TypeError:
return None
if not local_rows or any(not row for row in local_rows):
return None
safe_local = [
[
score if score is not None else math.inf
for score in (_finite_number(value, minimum=0.0) for value in row)
]
for row in local_rows
]
if len(safe_local) == 1:
try:
if len(transition_scores) != 0:
return None
except TypeError:
return None
best_index = min(range(len(safe_local[0])), key=safe_local[0].__getitem__)
best_score = safe_local[0][best_index]
if not math.isfinite(best_score):
return None
return CandidateSequenceSelection((best_index,), best_score)
try:
transitions = [[list(row) for row in matrix] for matrix in transition_scores]
except TypeError:
return None
if len(transitions) != len(safe_local) - 1:
return None
safe_transitions: list[list[list[float]]] = []
for index, matrix in enumerate(transitions):
previous_count = len(safe_local[index])
current_count = len(safe_local[index + 1])
if len(matrix) != previous_count or any(len(row) != current_count for row in matrix):
return None
safe_transitions.append(
[
[
score if score is not None else math.inf
for score in (_finite_number(value, minimum=0.0) for value in row)
]
for row in matrix
]
)
previous_costs = safe_local[0]
backpointers: list[list[int]] = []
for step in range(1, len(safe_local)):
current_costs = [math.inf] * len(safe_local[step])
current_backpointers = [-1] * len(safe_local[step])
for current_index, local_score in enumerate(safe_local[step]):
if not math.isfinite(local_score):
continue
for previous_index, previous_score in enumerate(previous_costs):
edge_score = safe_transitions[step - 1][previous_index][current_index]
if not math.isfinite(previous_score) or not math.isfinite(edge_score):
continue
total = previous_score + edge_score + local_score
if math.isfinite(total) and total < current_costs[current_index]:
current_costs[current_index] = total
current_backpointers[current_index] = previous_index
previous_costs = current_costs
backpointers.append(current_backpointers)
final_index = min(range(len(previous_costs)), key=previous_costs.__getitem__)
total_score = previous_costs[final_index]
if not math.isfinite(total_score):
return None
indices = [final_index]
for pointers in reversed(backpointers):
final_index = pointers[final_index]
if final_index < 0:
return None
indices.append(final_index)
indices.reverse()
return CandidateSequenceSelection(tuple(indices), total_score)
@torch.no_grad()
def extract_windowed_speaker_embedding(
wav_path: str,
encoder,
*,
device: str = "cpu",
min_duration_seconds: float = 3.0,
window_seconds: float = 3.0,
hop_seconds: float = 1.5,
max_windows: int = 12,
full_clip_max_seconds: float = 12.0,
) -> torch.Tensor:
"""Extract one denoised ECAPA embedding from overlapping reference windows."""
import librosa
waveform, _ = librosa.load(wav_path, sr=16000, mono=True)
waveform = np.asarray(waveform, dtype=np.float32)
duration = waveform.size / 16000.0
if duration < min_duration_seconds:
raise ValueError(
f"reference audio must be at least {min_duration_seconds:.1f} seconds; got {duration:.2f}"
)
segments: list[np.ndarray] = []
if duration <= full_clip_max_seconds:
segments.append(waveform)
window = max(1, int(round(window_seconds * 16000)))
hop = max(1, int(round(hop_seconds * 16000)))
starts = list(range(0, max(0, waveform.size - window) + 1, hop))
if max_windows > 0 and len(starts) > max_windows:
indices = np.linspace(0, len(starts) - 1, max_windows).round().astype(int)
starts = [starts[index] for index in dict.fromkeys(indices.tolist())]
segments.extend(waveform[start : start + window] for start in starts)
embeddings: list[torch.Tensor] = []
for segment in segments:
tensor = torch.from_numpy(np.ascontiguousarray(segment)).float().unsqueeze(0).to(device)
embedding = encoder.encode_batch(tensor).reshape(-1)
embeddings.append(torch.nn.functional.normalize(embedding, dim=0).cpu())
if not embeddings:
raise ValueError("reference audio did not contain a usable speech window")
return torch.nn.functional.normalize(torch.stack(embeddings).mean(dim=0), dim=0)