Files

239 lines
9.2 KiB
Python

"""Local Athena adaptation of the public X-VC Gradio demo."""
import logging
import os
import sys
import tempfile
import time
from pathlib import Path
from typing import Tuple
import gradio as gr
import numpy as np
import soundfile as sf
import torch
from huggingface_hub import hf_hub_download
from omegaconf import OmegaConf
HERE = "/opt/xvc"
sys.path.insert(0, HERE)
from bins.infer_utils import precompute_conditions, run_offline, run_streaming, to_numpy_audio
from models.codec.sac.model import XVC
from utils.audio import audio_highpass_filter, audio_volume_normalize, load_audio
logging.basicConfig(level=logging.INFO)
log = logging.getLogger("xvc-local")
MODEL_REPO = "chenxie95/X-VC"
SPACE_REPO = "hugging-apps/x-vc-voice-conversion"
SPEAKER_SUBDIR = "pretrained/speech_eres2net_sv_en_voxceleb_16k"
SAMPLE_RATE = 16000
ENHANCED_SAMPLE_RATE = 44100
LATENT_HOP_LENGTH = 1280
MAX_SECONDS = 20.0
RESEMBLE_REVISION = "4e3510ce4a8391159f665903544c5150bee7b2cb"
RESEMBLE_RUN_DIR = Path("/models/huggingface/resemble-enhance/enhancer_stage2")
MODE_OFFLINE = "Offline (höchste Qualität)"
MODE_STREAMING = "Streaming (simulierte Echtzeit)"
def _load_model() -> XVC:
speaker_config = hf_hub_download(
repo_id=SPACE_REPO,
repo_type="space",
filename=f"{SPEAKER_SUBDIR}/configuration.json",
)
hf_hub_download(
repo_id=SPACE_REPO,
repo_type="space",
filename=f"{SPEAKER_SUBDIR}/pretrained_eres2net.ckpt",
)
checkpoint = hf_hub_download(repo_id=MODEL_REPO, filename="xvc.pt")
cfg = OmegaConf.load(os.path.join(HERE, "configs", "xvc.yaml"))
cfg["model"]["generator"].pop("loss_config", None)
cfg["model"].pop("discriminator", None)
cfg["model"]["generator"]["speaker_encoder"]["pretrained_dir"] = os.path.dirname(speaker_config)
infer_cfg = os.path.join(tempfile.gettempdir(), "xvc_inference.yaml")
OmegaConf.save(cfg, infer_cfg)
loaded = XVC.load_from_checkpoint(infer_cfg, checkpoint, device=torch.device("cuda"))
loaded.remove_weight_norm()
loaded = loaded.eval().to("cuda")
log.info("X-VC model ready on %s", torch.cuda.get_device_name(0))
return loaded
MODEL = _load_model()
def _prepare_wav(path: str) -> np.ndarray:
wav = load_audio(path, sampling_rate=SAMPLE_RATE, volume_normalize=False)
if wav is None or len(wav) == 0:
raise gr.Error("Die Audiodatei konnte nicht gelesen werden.")
wav = wav[: int(MAX_SECONDS * SAMPLE_RATE)]
wav = audio_volume_normalize(wav)
wav = audio_highpass_filter(wav, SAMPLE_RATE, 40)
remainder = len(wav) % LATENT_HOP_LENGTH
if remainder:
wav = np.pad(wav, (0, LATENT_HOP_LENGTH - remainder), mode="constant")
return wav.astype(np.float32)
def _tensor(wav: np.ndarray) -> torch.Tensor:
return torch.from_numpy(wav).unsqueeze(0).unsqueeze(1).float().to("cuda")
def _write_wav(audio: np.ndarray, sample_rate: int = SAMPLE_RATE, suffix: str = "native-16k") -> str:
os.makedirs("/output", exist_ok=True)
path = os.path.join("/output", f"xvc-{suffix}-{int(time.time() * 1000)}.wav")
sf.write(path, np.clip(np.asarray(audio, dtype=np.float32), -1.0, 1.0), sample_rate, subtype="PCM_16")
return path
@torch.inference_mode()
def _enhance_wav(audio: np.ndarray) -> tuple[str, float]:
if not (RESEMBLE_RUN_DIR / "hparams.yaml").is_file():
raise gr.Error("Resemble-Enhance-Gewichte fehlen; bitte den Athena-Operator prüfen lassen.")
from resemble_enhance.enhancer.inference import enhance
started = time.time()
source = torch.from_numpy(np.asarray(audio, dtype=np.float32).reshape(-1))
restored, sample_rate = enhance(
source,
SAMPLE_RATE,
"cuda",
nfe=32,
solver="midpoint",
lambd=0.1,
tau=0.5,
run_dir=RESEMBLE_RUN_DIR,
)
if int(sample_rate) != ENHANCED_SAMPLE_RATE:
raise RuntimeError(f"Unerwartete Resemble-Enhance-Abtastrate: {sample_rate}")
return _write_wav(restored.cpu().numpy(), int(sample_rate), "enhanced-44k1"), time.time() - started
@torch.inference_mode()
def convert(
source_audio: str,
reference_audio: str,
enhance_44k1: bool = True,
mode: str = MODE_OFFLINE,
chunk_ms: int = 2400,
current_ms: int = 120,
future_ms: int = 100,
smooth_ms: int = 20,
progress=gr.Progress(track_tqdm=True),
) -> Tuple[str, str | None, str]:
if not source_audio:
raise gr.Error("Bitte eine Quelldatei mit dem zu erhaltenden Inhalt hochladen.")
if not reference_audio:
raise gr.Error("Bitte eine Referenzdatei mit der Zielstimme hochladen.")
source_np = _prepare_wav(source_audio)
reference_np = _prepare_wav(reference_audio)
source_wav = _tensor(source_np)
target_wav = _tensor(reference_np)
seconds = len(source_np) / SAMPLE_RATE
started = time.time()
if mode == MODE_STREAMING:
history_ms = int(chunk_ms) - int(current_ms) - int(smooth_ms) - int(future_ms)
if history_ms < 0:
raise gr.Error("Fenster muss mindestens Current + Lookahead + Crossfade umfassen.")
speaker_condition, frame_condition = precompute_conditions(MODEL, target_wav, target_wav)
recon, latency_ms = run_streaming(
model=MODEL,
source_wav=source_wav,
speaker_condition=speaker_condition,
frame_condition=frame_condition,
sample_rate=SAMPLE_RATE,
chunk_ms=int(chunk_ms),
current_ms=int(current_ms),
future_ms=int(future_ms),
smooth_ms=int(smooth_ms),
)
elapsed = time.time() - started
latency = np.asarray(latency_ms, dtype=np.float64)
report = (
f"**Streaming** · {len(latency)} Chunks · Mittel **{latency.mean():.0f} ms**, "
f"P95 **{np.percentile(latency, 95):.0f} ms** · insgesamt {elapsed:.2f} s "
f"für {seconds:.2f} s Audio (RTF {elapsed / seconds:.2f})"
)
else:
recon = run_offline(MODEL, source_wav, target_wav, target_wav)
elapsed = time.time() - started
report = (
f"**Offline** · {seconds:.2f} s Audio in {elapsed:.2f} s "
f"(RTF {elapsed / seconds:.2f})"
)
recon_np = to_numpy_audio(recon)
native_path = _write_wav(recon_np)
enhanced_path = None
if enhance_44k1:
del source_wav, target_wav, recon
torch.cuda.empty_cache()
enhanced_path, enhancement_seconds = _enhance_wav(recon_np)
report += (
f" · Resemble Enhance **{enhancement_seconds:.2f} s**, "
f"Ausgabe **44,1 kHz** (rekonstruierte Bandbreite)"
)
return native_path, enhanced_path, report
CSS = "#col-container { max-width: 1100px; margin: 0 auto; }"
HEADER = """# X-VC — Voice Changer
Die **Quelle** liefert Text, Aussprache und Timing. Die **Referenz** liefert die
Zielstimme. X-VC arbeitet Audio-zu-Audio ohne Transkript oder Training.
[Paper](https://arxiv.org/abs/2604.12456) · [Modell](https://huggingface.co/chenxie95/X-VC) ·
[Code](https://github.com/Jerrister/X-VC)
"""
with gr.Blocks(title="X-VC Voice Changer", theme=gr.themes.Citrus(), css=CSS) as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(HEADER)
with gr.Row():
source = gr.Audio(label="Quelle — Inhalt und Sprechweise", type="filepath", sources=["upload", "microphone"])
reference = gr.Audio(label="Referenz — gewünschte Zielstimme", type="filepath", sources=["upload", "microphone"])
run = gr.Button("Stimme umwandeln", variant="primary")
enhance_44k1 = gr.Checkbox(
value=True,
label="Zusätzlich mit Resemble Enhance auf 44,1 kHz restaurieren",
info="Rekonstruiert fehlende Höhen per KI; das native 16-kHz-Ergebnis bleibt zum Vergleich erhalten.",
)
with gr.Row():
output_native = gr.Audio(label="X-VC Original · 16 kHz", type="filepath", autoplay=False)
output_enhanced = gr.Audio(label="Enhanced · 44,1 kHz", type="filepath", autoplay=False)
report = gr.Markdown()
with gr.Accordion("Erweiterte Einstellungen", open=False):
mode = gr.Radio([MODE_OFFLINE, MODE_STREAMING], value=MODE_OFFLINE, label="Verarbeitungsmodus")
with gr.Row():
current_ms = gr.Slider(40, 640, value=120, step=40, label="Aktueller Chunk (ms)")
chunk_ms = gr.Slider(800, 4800, value=2400, step=200, label="Gesamtfenster (ms)")
with gr.Row():
future_ms = gr.Slider(0, 400, value=100, step=20, label="Lookahead (ms)")
smooth_ms = gr.Slider(0, 100, value=20, step=10, label="Crossfade (ms)")
gr.Markdown(
"Die ersten 20 Sekunden jeder Datei werden verarbeitet. X-VC bleibt nativ bei 16 kHz; "
"die optionale zweite Datei rekonstruiert die Sprachbandbreite auf 44,1 kHz."
)
run.click(
convert,
inputs=[source, reference, enhance_44k1, mode, chunk_ms, current_ms, future_ms, smooth_ms],
outputs=[output_native, output_enhanced, report],
api_name="convert",
)
if __name__ == "__main__":
demo.queue(default_concurrency_limit=1).launch(
server_name="0.0.0.0",
server_port=8009,
show_error=True,
)