Spaces:
Running on Zero
Running on Zero
Download production.py from voidful/BlueMagpie-TTS-Demo: direct link, hf CLI and curl.
- Browser
- Download file 18.7 kB
-
https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/744bf7a117e042397382e4948e665653169a1c9e/production.py
- Command line
-
hf download hf://spaces/voidful/BlueMagpie-TTS-Demo@744bf7a117e042397382e4948e665653169a1c9e/production.py
-
curl -L -o production.py https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/744bf7a117e042397382e4948e665653169a1c9e/production.py
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 | |
| 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) | |