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