""" 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()