Spaces:
Sleeping
Sleeping
File size: 5,155 Bytes
76db545 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | """
Facebook MMS-TTS engine for Bambara, Fula, French, and English.
Usage:
engine = MMSTTSEngine()
wav_np, sample_rate = engine.synthesize("Foro fɛ ji.", "bam", device="cuda")
wav_bytes = engine.text_to_audio_bytes("Foro fɛ ji.", "bam", device="cuda")
"""
from __future__ import annotations
import io
import re
from typing import Dict, Tuple
import numpy as np
import soundfile as sf
MODEL_IDS: Dict[str, str] = {
"bam": "facebook/mms-tts-bam",
"ful": "facebook/mms-tts-ful",
"fr": "facebook/mms-tts-fra",
"en": "facebook/mms-tts-eng",
}
# Fallback for unknown languages — use French
_DEFAULT_LANG = "fr"
# MMS-TTS quality degrades beyond ~15 words; split longer text at sentence boundaries
_MAX_WORDS_PER_CHUNK = 15
# Sentence-boundary split pattern (period, exclamation, question mark followed by space or end)
_SENTENCE_RE = re.compile(r"(?<=[.!?])\s+")
class MMSTTSEngine:
"""Lazy-loading MMS-TTS engine. Models are loaded on first use and cached in CPU RAM."""
def __init__(self) -> None:
# {language_code: (VitsModel, VitsTokenizer)}
self._cache: Dict[str, tuple] = {}
# ── private helpers ──────────────────────────────────────────────────────
def _get_model(self, language: str):
"""Return (model, tokenizer) for the requested language, loading if needed."""
lang = language if language in MODEL_IDS else _DEFAULT_LANG
if lang not in self._cache:
from transformers import VitsModel, VitsTokenizer # type: ignore
model_id = MODEL_IDS[lang]
tokenizer = VitsTokenizer.from_pretrained(model_id)
model = VitsModel.from_pretrained(model_id)
model.eval()
# Keep on CPU until synthesize() moves it to the target device
self._cache[lang] = (model, tokenizer)
return self._cache[lang]
@staticmethod
def _split_sentences(text: str) -> list[str]:
"""Split text into chunks of ≤ _MAX_WORDS_PER_CHUNK words."""
sentences = _SENTENCE_RE.split(text.strip())
chunks: list[str] = []
current: list[str] = []
current_words = 0
for sent in sentences:
words = sent.split()
if current_words + len(words) > _MAX_WORDS_PER_CHUNK and current:
chunks.append(" ".join(current))
current = words
current_words = len(words)
else:
current.extend(words)
current_words += len(words)
if current:
chunks.append(" ".join(current))
return chunks or [text]
def _synthesize_chunk(
self, text: str, model, tokenizer, device: str
) -> np.ndarray:
"""Synthesize a single short text chunk. Returns 1-D float32 numpy array."""
import torch
model.to(device)
inputs = tokenizer(text, return_tensors="pt")
inputs = {k: v.to(device) for k, v in inputs.items()}
with torch.no_grad():
output = model(**inputs)
waveform = output.waveform[0].cpu().numpy() # shape: (samples,)
return waveform
# ── public API ───────────────────────────────────────────────────────────
def synthesize(
self, text: str, language: str, device: str = "cuda"
) -> Tuple[np.ndarray, int]:
"""
Convert text to speech waveform.
Args:
text: Text to synthesize (any length — long text is split automatically).
language: Language code: "bam", "ful", "fr", or "en".
device: "cuda" or "cpu".
Returns:
(waveform_np, sample_rate) — float32 numpy array, sample rate in Hz.
"""
lang = language if language in MODEL_IDS else _DEFAULT_LANG
model, tokenizer = self._get_model(lang)
chunks = self._split_sentences(text)
waveforms: list[np.ndarray] = []
for chunk in chunks:
if not chunk.strip():
continue
waveforms.append(self._synthesize_chunk(chunk, model, tokenizer, device))
# Free device memory before returning
model.to("cpu")
if not waveforms:
return np.zeros(1, dtype=np.float32), model.config.sampling_rate
combined = np.concatenate(waveforms)
return combined, model.config.sampling_rate
def text_to_audio_bytes(
self, text: str, language: str, device: str = "cuda"
) -> bytes:
"""
Convert text to WAV bytes suitable for gr.Audio or HTTP response.
Returns raw WAV file bytes (16-bit PCM).
"""
waveform, sample_rate = self.synthesize(text, language, device=device)
buf = io.BytesIO()
# soundfile expects float32 in [-1, 1]; MMS output is already normalised
sf.write(buf, waveform, sample_rate, format="WAV", subtype="PCM_16")
return buf.getvalue()
|