203 lines
6.0 KiB
Python
203 lines
6.0 KiB
Python
#!/usr/bin/env python3
|
|
"""Small local-only Vevo2 voice-conversion studio for Athena."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import os
|
|
import re
|
|
import shutil
|
|
import subprocess
|
|
import threading
|
|
import time
|
|
import uuid
|
|
from contextlib import asynccontextmanager
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
|
|
from fastapi.responses import FileResponse, HTMLResponse
|
|
from starlette.background import BackgroundTask
|
|
|
|
import models.svc.vevo2.infer_vevo2_fm as vevo
|
|
|
|
|
|
DATA_DIR = Path(os.environ.get("VOICE_DATA_DIR", "/data"))
|
|
PROFILE_DIR = DATA_DIR / "profiles"
|
|
JOB_DIR = DATA_DIR / "jobs"
|
|
INDEX = Path("/app/index.html")
|
|
MAX_UPLOAD_BYTES = int(os.environ.get("MAX_UPLOAD_BYTES", str(100 * 1024 * 1024)))
|
|
MODEL_LOCK = threading.Lock()
|
|
PIPELINE = None
|
|
MODEL_LOAD_SECONDS: float | None = None
|
|
|
|
|
|
def safe_name(value: str) -> str:
|
|
value = re.sub(r"[^A-Za-z0-9_. -]+", "-", value.strip())
|
|
value = re.sub(r"\s+", "-", value).strip("-.")
|
|
return value[:64] or "voice"
|
|
|
|
|
|
def load_pipeline() -> None:
|
|
global PIPELINE, MODEL_LOAD_SECONDS
|
|
if PIPELINE is not None:
|
|
return
|
|
with MODEL_LOCK:
|
|
if PIPELINE is not None:
|
|
return
|
|
started = time.monotonic()
|
|
PIPELINE = vevo.load_inference_pipeline()
|
|
vevo.inference_pipeline = PIPELINE
|
|
MODEL_LOAD_SECONDS = time.monotonic() - started
|
|
|
|
|
|
def to_wav(source: Path, target: Path) -> None:
|
|
completed = subprocess.run(
|
|
["ffmpeg", "-hide_banner", "-loglevel", "error", "-y", "-i", str(source),
|
|
"-ac", "1", "-ar", "24000", "-c:a", "pcm_s16le", str(target)],
|
|
capture_output=True,
|
|
text=True,
|
|
timeout=180,
|
|
check=False,
|
|
)
|
|
if completed.returncode:
|
|
raise ValueError(completed.stderr.strip() or "Audio konnte nicht gelesen werden")
|
|
|
|
|
|
async def save_upload(upload: UploadFile, target: Path) -> None:
|
|
size = 0
|
|
with target.open("wb") as handle:
|
|
while chunk := await upload.read(1024 * 1024):
|
|
size += len(chunk)
|
|
if size > MAX_UPLOAD_BYTES:
|
|
raise HTTPException(413, "Audiodatei ist zu groß")
|
|
handle.write(chunk)
|
|
|
|
|
|
def profile_path(name: str) -> Path:
|
|
target = PROFILE_DIR / f"{safe_name(name)}.wav"
|
|
if not target.is_file():
|
|
raise HTTPException(404, "Referenzstimme wurde nicht gefunden")
|
|
return target
|
|
|
|
|
|
def run_conversion(source: Path, reference: Path, output: Path, pitch_shift: bool) -> tuple[float, float]:
|
|
"""Run the GPU-bound conversion off the API event loop."""
|
|
load_pipeline()
|
|
with MODEL_LOCK:
|
|
torch.cuda.reset_peak_memory_stats()
|
|
started = time.monotonic()
|
|
vevo.vevo2_fm(str(source), str(reference), str(output), shifted_src=pitch_shift)
|
|
elapsed = time.monotonic() - started
|
|
peak = torch.cuda.max_memory_allocated() / 1048576
|
|
return elapsed, peak
|
|
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(_app: FastAPI):
|
|
PROFILE_DIR.mkdir(parents=True, exist_ok=True)
|
|
JOB_DIR.mkdir(parents=True, exist_ok=True)
|
|
load_pipeline()
|
|
yield
|
|
|
|
|
|
app = FastAPI(title="Athena Voice Studio", lifespan=lifespan)
|
|
|
|
|
|
@app.get("/", response_class=HTMLResponse)
|
|
def index() -> str:
|
|
return INDEX.read_text(encoding="utf-8")
|
|
|
|
|
|
@app.get("/health")
|
|
def health() -> dict:
|
|
return {
|
|
"status": "ok" if PIPELINE is not None else "starting",
|
|
"model": "RMSnow/Vevo2",
|
|
"sample_rate": 24000,
|
|
"model_load_seconds": MODEL_LOAD_SECONDS,
|
|
}
|
|
|
|
|
|
@app.get("/api/profiles")
|
|
def profiles() -> dict:
|
|
return {"profiles": [path.stem for path in sorted(PROFILE_DIR.glob("*.wav"))]}
|
|
|
|
|
|
@app.post("/api/profiles")
|
|
async def create_profile(
|
|
name: str = Form(...),
|
|
consent: bool = Form(False),
|
|
audio: UploadFile = File(...),
|
|
) -> dict:
|
|
if not consent:
|
|
raise HTTPException(400, "Bestätige, dass du die Stimme verwenden darfst")
|
|
clean = safe_name(name)
|
|
job = JOB_DIR / f"profile-{uuid.uuid4().hex}"
|
|
job.mkdir(parents=True)
|
|
raw = job / f"upload-{safe_name(audio.filename or 'reference.audio')}"
|
|
try:
|
|
await save_upload(audio, raw)
|
|
target = PROFILE_DIR / f"{clean}.wav"
|
|
temporary = job / "reference.wav"
|
|
to_wav(raw, temporary)
|
|
os.replace(temporary, target)
|
|
return {"status": "ok", "profile": clean}
|
|
except HTTPException:
|
|
raise
|
|
except Exception as exc:
|
|
raise HTTPException(400, str(exc)) from exc
|
|
finally:
|
|
shutil.rmtree(job, ignore_errors=True)
|
|
|
|
|
|
@app.delete("/api/profiles/{name}")
|
|
def delete_profile(name: str) -> dict:
|
|
target = profile_path(name)
|
|
target.unlink()
|
|
return {"status": "ok", "profile": target.stem}
|
|
|
|
|
|
@app.post("/api/convert")
|
|
async def convert(
|
|
source: UploadFile = File(...),
|
|
profile: str = Form(...),
|
|
pitch_shift: bool = Form(True),
|
|
) -> FileResponse:
|
|
reference = profile_path(profile)
|
|
job_id = uuid.uuid4().hex
|
|
job = JOB_DIR / f"convert-{job_id}"
|
|
job.mkdir(parents=True)
|
|
raw = job / f"upload-{safe_name(source.filename or 'source.audio')}"
|
|
source_wav = job / "source.wav"
|
|
output = JOB_DIR / f"voice-{job_id}.wav"
|
|
try:
|
|
await save_upload(source, raw)
|
|
to_wav(raw, source_wav)
|
|
elapsed, peak = await asyncio.to_thread(
|
|
run_conversion, source_wav, reference, output, pitch_shift
|
|
)
|
|
return FileResponse(
|
|
output,
|
|
media_type="audio/wav",
|
|
filename=f"{safe_name(profile)}-{job_id[:8]}.wav",
|
|
headers={
|
|
"X-Conversion-Seconds": f"{elapsed:.3f}",
|
|
"X-Peak-VRAM-MiB": f"{peak:.1f}",
|
|
},
|
|
background=BackgroundTask(output.unlink, missing_ok=True),
|
|
)
|
|
except HTTPException:
|
|
raise
|
|
except Exception as exc:
|
|
output.unlink(missing_ok=True)
|
|
raise HTTPException(500, f"Stimmenwandlung fehlgeschlagen: {exc}") from exc
|
|
finally:
|
|
shutil.rmtree(job, ignore_errors=True)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
import uvicorn
|
|
|
|
uvicorn.run(app, host="0.0.0.0", port=8008)
|