SVC / app.py
AlaraRad's picture
Initial commit
6f82bd7
Raw History Blame Contribute Delete
12.8 kB
"""
TSVC β€” Gradio interface for Hugging Face Spaces.
Place your voices/ directory and onnx/ directory next to this file.
voices/
voice_1.wav ← default target
voice_2.wav
onnx/
encode.onnx
decode.onnx
ssl.onnx
"""
import collections
import logging
import os
import tempfile
from pathlib import Path
import gradio as gr
import numpy as np
import soundfile as sf
# ── Logging ───────────────────────────────────────────────────────────────────
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s", datefmt="%H:%M:%S")
log = logging.getLogger("tsvc")
# ── Constants ─────────────────────────────────────────────────────────────────
SAMPLE_RATE = 44_100
OUT_SAMPLE_RATE = 48_000
SSL_SAMPLE_RATE = 16_000
WAVLM_HOP = 320
MIN_PERIOD = 110
SPLICE_FADE = 32
TOKEN_HZ = 25
TOKEN_SAMPLES = SAMPLE_RATE // TOKEN_HZ
CHUNK_TOKENS = 12
WINDOW_TOKENS = 36
HERE = Path(__file__).parent
ONNX_DIR = Path(os.environ.get("ONNX_DIR", HERE / "onnx"))
VOICES_DIR = Path(os.environ.get("VOICES_DIR", HERE / "voices"))
def _list_voices() -> list[str]:
if not VOICES_DIR.exists():
return []
exts = {".wav", ".mp3", ".flac", ".ogg", ".m4a"}
return sorted(p.name for p in VOICES_DIR.iterdir() if p.suffix.lower() in exts)
# ── DSP helpers ───────────────────────────────────────────────────────────────
def resample_np(audio: np.ndarray, sr_in: int, sr_out: int) -> np.ndarray:
if sr_in == sr_out:
return audio
ratio = sr_out / sr_in
out_len = int(len(audio) * ratio)
idx = np.arange(out_len, dtype=np.float64) / ratio
lo = np.clip(np.floor(idx).astype(np.int64), 0, len(audio) - 1)
hi = np.clip(lo + 1, 0, len(audio) - 1)
frac = (idx - lo).astype(np.float32)
return audio[lo] * (1.0 - frac) + audio[hi] * frac
def pad_or_trim_1d(arr: np.ndarray, n: int) -> np.ndarray:
if len(arr) >= n:
return arr[:n]
return np.concatenate([arr, np.zeros(n - len(arr), dtype=arr.dtype)])
def pad_or_trim_2d(arr: np.ndarray, rows: int) -> np.ndarray:
if arr.shape[0] >= rows:
return arr[:rows]
return np.concatenate([arr, np.zeros((rows - arr.shape[0], arr.shape[1]), dtype=arr.dtype)])
def find_splice_point(prev_tail: np.ndarray, overlap: np.ndarray) -> int:
n = len(prev_tail)
margin = SPLICE_FADE + MIN_PERIOD
if n < 2 * margin:
return n // 2
best_idx, best_cost = margin, np.inf
step = max(1, MIN_PERIOD // 4)
for i in range(margin, n - margin, step):
lo = max(0, i - MIN_PERIOD // 2)
hi = min(n, i + MIN_PERIOD // 2)
diff = prev_tail[lo:hi] - overlap[lo:hi]
cost = float(np.dot(diff, diff))
if cost < best_cost:
best_cost, best_idx = cost, i
return best_idx
def splice(prev_tail: np.ndarray, overlap: np.ndarray) -> np.ndarray:
n = len(prev_tail)
idx = find_splice_point(prev_tail, overlap)
out = np.empty(n, dtype=np.float32)
out[:idx] = prev_tail[:idx]
out[idx:] = overlap[idx:]
fade_lo = max(0, idx - SPLICE_FADE)
fade_hi = min(n, idx + SPLICE_FADE)
fade_len = fade_hi - fade_lo
if fade_len > 1:
w = 0.5 - 0.5 * np.cos(np.linspace(0.0, np.pi, fade_len, dtype=np.float32))
out[fade_lo:fade_hi] = (
prev_tail[fade_lo:fade_hi] * (1 - w) + overlap[fade_lo:fade_hi] * w
)
return out
# ── Model loader (singleton) ──────────────────────────────────────────────────
class ModelBundle:
_instance: "ModelBundle | None" = None
def __init__(self):
import onnxruntime as ort
providers = (
["CUDAExecutionProvider", "CPUExecutionProvider"]
if ort.get_device() == "GPU"
else ["CPUExecutionProvider"]
)
opts = ort.SessionOptions()
opts.intra_op_num_threads = 4
opts.inter_op_num_threads = 1
opts.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
log.info(f"Loading ONNX sessions from {ONNX_DIR} …")
self.enc = ort.InferenceSession(str(ONNX_DIR / "encode.onnx"), sess_options=opts, providers=providers)
self.dec = ort.InferenceSession(str(ONNX_DIR / "decode.onnx"), sess_options=opts, providers=providers)
self.ssl = ort.InferenceSession(str(ONNX_DIR / "ssl.onnx"), sess_options=opts, providers=providers)
self.ssl_seq_len = self.enc.get_inputs()[0].shape[0]
self.fixed_16k_len = self.ssl_seq_len * WAVLM_HOP
self.fixed_audio_len = int(self.fixed_16k_len / SSL_SAMPLE_RATE * SAMPLE_RATE)
log.info("Models ready.")
@classmethod
def get(cls) -> "ModelBundle":
if cls._instance is None:
cls._instance = cls()
return cls._instance
# ── Inference helpers ─────────────────────────────────────────────────────────
def _extract_ssl(models: ModelBundle, audio_np: np.ndarray):
audio_16k = resample_np(audio_np, SAMPLE_RATE, SSL_SAMPLE_RATE)
audio_16k = (
pad_or_trim_1d(audio_16k, models.fixed_16k_len)
.reshape(1, -1)
.astype(np.float32)
)
return models.ssl.run(["local_features", "global_features"], {"audio_16k": audio_16k})
def compute_embedding(models: ModelBundle, audio: np.ndarray) -> np.ndarray:
embeddings = []
for i in range(0, len(audio), models.fixed_audio_len):
chunk = audio[i : i + models.fixed_audio_len]
if len(chunk) < models.fixed_audio_len // 2 and embeddings:
break
chunk = pad_or_trim_1d(chunk, models.fixed_audio_len)
_, tgt_g = _extract_ssl(models, chunk)
tgt_g = pad_or_trim_2d(tgt_g, models.ssl_seq_len).astype(np.float32)
_, glob_emb = models.enc.run(
["content_token_indices", "global_embedding"],
{"local_ssl_features": tgt_g, "global_ssl_features": tgt_g},
)
embeddings.append(glob_emb)
if not embeddings:
raise ValueError("Reference audio is too short.")
return np.mean(embeddings, axis=0)
def convert_audio(
models: ModelBundle,
source: np.ndarray,
global_embedding: np.ndarray,
) -> np.ndarray:
"""Process source in fixed-size chunks and stitch with crossfade splice."""
chunk_samples = CHUNK_TOKENS * TOKEN_SAMPLES
window_samples = WINDOW_TOKENS * TOKEN_SAMPLES
fade_in = np.linspace(0.0, 1.0, chunk_samples, dtype=np.float32)
buf = collections.deque(np.zeros(window_samples, dtype=np.float32), maxlen=window_samples)
out_chunks = []
prev_tail = None
pos = 0
while pos < len(source):
chunk = source[pos : pos + chunk_samples]
chunk = pad_or_trim_1d(chunk, chunk_samples)
pos += chunk_samples
buf.extend(chunk)
window = np.array(buf, dtype=np.float32)
audio = pad_or_trim_1d(window, models.fixed_audio_len)
src_local, _ = _extract_ssl(models, audio)
src_local = pad_or_trim_2d(src_local, models.ssl_seq_len).astype(np.float32)
content_indices, _ = models.enc.run(
["content_token_indices", "global_embedding"],
{"local_ssl_features": src_local, "global_ssl_features": src_local},
)
(waveform,) = models.dec.run(
["waveform"],
{"content_token_indices": content_indices, "global_embedding": global_embedding},
)
tail_region = waveform[-(chunk_samples * 2):]
overlap, output = tail_region[:chunk_samples], tail_region[chunk_samples:]
chunk_out = overlap * fade_in if prev_tail is None else splice(prev_tail, overlap)
prev_tail = output
out_chunks.append(chunk_out)
if not out_chunks:
return np.zeros(0, dtype=np.float32)
return resample_np(np.concatenate(out_chunks), SAMPLE_RATE, OUT_SAMPLE_RATE)
# ── Gradio callback ───────────────────────────────────────────────────────────
def _load_audio(path: str) -> tuple[np.ndarray, float]:
"""Read any audio file, mix to mono, normalise, resample to SAMPLE_RATE."""
audio, sr = sf.read(path, dtype="float32", always_2d=True)
audio = audio.mean(axis=1)
if sr != SAMPLE_RATE:
audio = resample_np(audio, sr, SAMPLE_RATE)
peak = np.abs(audio).max()
if peak > 1e-8:
audio /= peak
return audio, len(audio) / SAMPLE_RATE
def run_conversion(target_path: str, source_audio):
logs = []
def note(msg: str):
log.info(msg)
logs.append(msg)
if not target_path:
return None, "❌ No target voice selected."
if source_audio is None:
return None, "❌ Please upload or record a source audio clip."
try:
models = ModelBundle.get()
except Exception as e:
return None, f"❌ Failed to load models: {e}"
try:
ref_audio, ref_dur = _load_audio(target_path)
note(f"Target: {Path(target_path).name} ({ref_dur:.1f}s)")
global_embedding = compute_embedding(models, ref_audio)
note("Reference encoded.")
except Exception as e:
return None, f"❌ Reference load failed: {e}"
try:
src_audio, src_dur = _load_audio(source_audio)
note(f"Source loaded ({src_dur:.1f}s). Converting…")
except Exception as e:
return None, f"❌ Source load failed: {e}"
try:
converted = convert_audio(models, src_audio, global_embedding)
tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
sf.write(tmp.name, converted, OUT_SAMPLE_RATE)
note(f"Done! {len(converted) / OUT_SAMPLE_RATE:.1f}s output.")
return tmp.name, "\n".join(logs)
except Exception as e:
log.exception("Conversion failed")
return None, f"❌ Conversion failed: {e}"
# ── UI ────────────────────────────────────────────────────────────────────────
def build_ui() -> gr.Blocks:
voices = _list_voices()
default_target = str(VOICES_DIR / "voice_1.wav") if "voice_1.wav" in voices else None
default_source = str(VOICES_DIR / "voice_2.wav") if "voice_2.wav" in voices else None
with gr.Blocks(title="TSVC", theme=gr.themes.Base()) as demo:
gr.Markdown("# πŸŽ™οΈ TSVC β€” Voice Conversion")
with gr.Row():
with gr.Column(scale=1):
target_audio = gr.Audio(
label="Target Voice",
type="filepath",
sources=["upload"],
value=default_target,
)
source_audio = gr.Audio(
label="Source Audio",
type="filepath",
sources=["upload", "microphone"],
value=default_source,
)
gr.Markdown(
"You can upload or record your own target or source audio. "
"The conversion will output audio from source as the style of target. "
"This is all running purely on CPU and will work in a streaming setting on desktop."
)
convert_btn = gr.Button("πŸ”„ Convert", variant="primary")
with gr.Column(scale=1):
output_audio = gr.Audio(
label="πŸ”Š Converted audio",
type="filepath",
interactive=False,
)
log_box = gr.Textbox(
label="Log",
lines=8,
interactive=False,
placeholder="Log will appear here…",
)
convert_btn.click(
fn=run_conversion,
inputs=[target_audio, source_audio],
outputs=[output_audio, log_box],
)
return demo
if __name__ == "__main__":
demo = build_ui()
demo.launch()