voxcpm2-tools / voxcpm2_api.py
SWAG456's picture
Upload voxcpm2_api.py
99088ba verified
Raw
History Blame Contribute Delete
10 kB
#!/usr/bin/env python3
"""
VoxCPM2 REST API Server for VPS
Runs on Hostinger KVM 4 (or similar) with CPU inference.
Endpoints:
POST /tts - Basic text-to-speech
POST /design - Voice design TTS
POST /clone - Voice cloning (upload reference audio)
POST /multilingual - Multilingual TTS
POST /multilingual-clone ⭐ NEW - Clone voice + speak any language
POST /multilingual-design ⭐ NEW - Design voice + speak any language
GET /health - Health check
Run: python voxcpm2_api.py
(uses PORT env var, defaults to 5000)
Requires: pip install flask voxcpm soundfile
"""
import os
import sys
import time
import tempfile
import warnings
from functools import lru_cache
warnings.filterwarnings("ignore")
from flask import Flask, request, jsonify, send_file
app = Flask(__name__)
# ── Config ──
MODEL_ID = os.environ.get("VOXCPM_MODEL", "openbmb/VoxCPM2")
PORT = int(os.environ.get("PORT", 5000))
HOST = os.environ.get("HOST", "0.0.0.0")
DEFAULT_TIMESTEPS = int(os.environ.get("VOXCPM_TIMESTEPS", "10"))
DEFAULT_CFG = float(os.environ.get("VOXCPM_CFG", "2.0"))
@lru_cache(maxsize=1)
def get_model():
"""Lazy-load model once and cache it."""
from voxcpm import VoxCPM
import torch
print(f"\n[LOADING] Model: {MODEL_ID}")
print(f"[LOADING] This may take 2-5 minutes on first run...")
start = time.time()
model = VoxCPM.from_pretrained(
MODEL_ID,
load_denoiser=False,
optimize=False,
local_files_only=False,
)
# CPU safety: convert bfloat16 → float32
model.tts_model = model.tts_model.to(torch.float32)
elapsed = time.time() - start
print(f"[OK] Model loaded in {elapsed:.1f}s")
print(f"[OK] Sample rate: {model.tts_model.sample_rate} Hz")
return model
def _save_wav(wav, prefix="output"):
"""Save waveform to a temp WAV file."""
import soundfile as sf
model = get_model()
fd, path = tempfile.mkstemp(suffix=".wav", prefix=prefix + "_")
os.close(fd)
sf.write(path, wav, model.tts_model.sample_rate)
return path
@app.route("/health", methods=["GET"])
def health():
return jsonify({"status": "ok", "model": MODEL_ID, "device": "cpu"})
@app.route("/tts", methods=["POST"])
def tts():
"""Basic TTS. JSON: {"text": "Hello", "timesteps": 10, "cfg_value": 2.0}"""
data = request.get_json() or {}
text = data.get("text", "")
if not text:
return jsonify({"error": "Missing 'text' field"}), 400
timesteps = data.get("timesteps", DEFAULT_TIMESTEPS)
cfg = data.get("cfg_value", DEFAULT_CFG)
model = get_model()
start = time.time()
wav = model.generate(text=text, cfg_value=float(cfg),
inference_timesteps=int(timesteps), max_len=4096)
elapsed = time.time() - start
path = _save_wav(wav, "tts")
return send_file(path, mimetype="audio/wav", as_attachment=True,
download_name="tts_output.wav")
@app.route("/design", methods=["POST"])
def design():
"""Voice Design TTS. JSON: {"text": "Hello", "description": "warm female voice"}"""
data = request.get_json() or {}
text = data.get("text", "")
description = data.get("description", "A warm professional voice")
if not text:
return jsonify({"error": "Missing 'text' field"}), 400
timesteps = data.get("timesteps", DEFAULT_TIMESTEPS)
cfg = data.get("cfg_value", DEFAULT_CFG)
full_text = f"({description}) {text}"
model = get_model()
start = time.time()
wav = model.generate(text=full_text, cfg_value=float(cfg),
inference_timesteps=int(timesteps), max_len=4096)
elapsed = time.time() - start
path = _save_wav(wav, "design")
return send_file(path, mimetype="audio/wav", as_attachment=True,
download_name="design_output.wav")
@app.route("/clone", methods=["POST"])
def clone():
"""Voice Cloning. multipart: text, reference (WAV file upload)"""
text = request.form.get("text", "")
if not text:
return jsonify({"error": "Missing 'text' field"}), 400
if "reference" not in request.files:
return jsonify({"error": "Missing 'reference' audio file"}), 400
ref_file = request.files["reference"]
timesteps = int(request.form.get("timesteps", DEFAULT_TIMESTEPS))
cfg = float(request.form.get("cfg_value", DEFAULT_CFG))
fd, ref_path = tempfile.mkstemp(suffix=".wav", prefix="ref_")
os.close(fd)
ref_file.save(ref_path)
model = get_model()
start = time.time()
wav = model.generate(text=text, reference_wav_path=ref_path,
cfg_value=cfg, inference_timesteps=timesteps, max_len=4096)
elapsed = time.time() - start
os.remove(ref_path)
path = _save_wav(wav, "clone")
return send_file(path, mimetype="audio/wav", as_attachment=True,
download_name="clone_output.wav")
@app.route("/multilingual", methods=["POST"])
def multilingual():
"""Multilingual TTS. JSON: {"text": "你好"}"""
data = request.get_json() or {}
text = data.get("text", "")
if not text:
return jsonify({"error": "Missing 'text' field"}), 400
timesteps = data.get("timesteps", DEFAULT_TIMESTEPS)
cfg = data.get("cfg_value", DEFAULT_CFG)
model = get_model()
start = time.time()
wav = model.generate(text=text, cfg_value=float(cfg),
inference_timesteps=int(timesteps), max_len=4096)
elapsed = time.time() - start
path = _save_wav(wav, "multi")
return send_file(path, mimetype="audio/wav", as_attachment=True,
download_name="multilingual_output.wav")
# ⭐ NEW ENDPOINT: Multilingual Voice Clone
@app.route("/multilingual-clone", methods=["POST"])
def multilingual_clone():
"""Clone a voice and speak in ANY language! multipart: text (any lang), reference (WAV)"""
text = request.form.get("text", "")
if not text:
return jsonify({"error": "Missing 'text' field"}), 400
if "reference" not in request.files:
return jsonify({"error": "Missing 'reference' audio file"}), 400
ref_file = request.files["reference"]
timesteps = int(request.form.get("timesteps", DEFAULT_TIMESTEPS))
cfg = float(request.form.get("cfg_value", DEFAULT_CFG))
fd, ref_path = tempfile.mkstemp(suffix=".wav", prefix="ref_")
os.close(fd)
ref_file.save(ref_path)
model = get_model()
start = time.time()
wav = model.generate(text=text, reference_wav_path=ref_path,
cfg_value=cfg, inference_timesteps=timesteps, max_len=4096)
elapsed = time.time() - start
os.remove(ref_path)
path = _save_wav(wav, "multi_clone")
return send_file(path, mimetype="audio/wav", as_attachment=True,
download_name="multilingual_clone_output.wav")
# ⭐ NEW ENDPOINT: Multilingual Voice Design
@app.route("/multilingual-design", methods=["POST"])
def multilingual_design():
"""Design a voice and speak in ANY language! JSON: {"text": "你好", "description": "warm female voice"}"""
data = request.get_json() or {}
text = data.get("text", "")
description = data.get("description", "A warm professional voice")
if not text:
return jsonify({"error": "Missing 'text' field"}), 400
timesteps = data.get("timesteps", DEFAULT_TIMESTEPS)
cfg = data.get("cfg_value", DEFAULT_CFG)
full_text = f"({description}) {text}"
model = get_model()
start = time.time()
wav = model.generate(text=full_text, cfg_value=float(cfg),
inference_timesteps=int(timesteps), max_len=4096)
elapsed = time.time() - start
path = _save_wav(wav, "multi_design")
return send_file(path, mimetype="audio/wav", as_attachment=True,
download_name="multilingual_design_output.wav")
@app.route("/", methods=["GET"])
def index():
return jsonify({
"service": "VoxCPM2 TTS API",
"model": MODEL_ID,
"endpoints": {
"GET /health": "Health check",
"POST /tts": "Basic TTS (JSON: text, timesteps, cfg_value)",
"POST /design": "Voice design (JSON: text, description, timesteps, cfg_value)",
"POST /clone": "Voice cloning (multipart: text, reference, timesteps, cfg_value)",
"POST /multilingual": "Multilingual TTS (JSON: text, timesteps, cfg_value)",
"POST /multilingual-clone": "⭐ Clone voice + speak ANY language (multipart: text, reference)",
"POST /multilingual-design": "⭐ Design voice + speak ANY language (JSON: text, description)",
},
"timesteps_guide": {
"4-5": "Draft / fast",
"8-10": "Balanced (default)",
"15-20": "High quality",
"25-30": "Best quality / slow"
},
"cfg_guide": {
"1.0-1.5": "More natural",
"2.0": "Balanced (default)",
"2.5-3.0": "Tighter control"
}
})
if __name__ == "__main__":
print("=" * 60)
print("VoxCPM2 REST API Server")
print("=" * 60)
print(f"Model: {MODEL_ID}")
print(f"Host: {HOST}:{PORT}")
print(f"Timesteps default: {DEFAULT_TIMESTEPS}")
print(f"CFG default: {DEFAULT_CFG}")
print()
print("Endpoints:")
print(" GET /health")
print(" POST /tts")
print(" POST /design")
print(" POST /clone")
print(" POST /multilingual")
print(" POST /multilingual-clone ⭐ NEW")
print(" POST /multilingual-design ⭐ NEW")
print()
print("Environment variables:")
print(" VOXCPM_MODEL - Model ID (default: openbmb/VoxCPM2)")
print(" VOXCPM_TIMESTEPS- Default timesteps (default: 10)")
print(" VOXCPM_CFG - Default CFG (default: 2.0)")
print(" PORT - Server port (default: 5000)")
print(" HOST - Bind address (default: 0.0.0.0)")
print("=" * 60)
app.run(host=HOST, port=PORT, debug=False, threaded=True)