diff --git a/.gitignore b/.gitignore index f7fa981..5450105 100644 --- a/.gitignore +++ b/.gitignore @@ -1,7 +1,13 @@ launch.json .venv/ __pycache__ +.numba_cache/ voxcpm.egg-info .DS_Store ./pretrained_models/ +models/ +lora/ +outputs/ +checkpoints/ +logs/ app_local.py diff --git a/app.py b/app.py index 99a3c5c..76c38b6 100644 --- a/app.py +++ b/app.py @@ -3,21 +3,40 @@ import sys import logging import random -import numpy as np -import gradio as gr -from typing import Optional, Tuple -from funasr import AutoModel +import tempfile +import threading +from typing import Callable, Optional, Tuple from pathlib import Path os.environ["TOKENIZERS_PARALLELISM"] = "false" +PROJECT_ROOT = Path(__file__).resolve().parent +os.chdir(PROJECT_ROOT) +# Librosa imports Numba during Parakeet ASR setup. Keep Numba's cache rooted in +# the project so launching app.py from another directory cannot stall startup. +os.environ.setdefault("NUMBA_CACHE_DIR", str(PROJECT_ROOT / ".numba_cache")) +SRC_DIR = PROJECT_ROOT / "src" +if SRC_DIR.exists() and str(SRC_DIR) not in sys.path: + sys.path.insert(0, str(SRC_DIR)) + +import numpy as np +import gradio as gr +from funasr import AutoModel import voxcpm from voxcpm.model.utils import resolve_runtime_device +for _stream in (sys.stdout, sys.stderr): + if hasattr(_stream, "reconfigure"): + try: + _stream.reconfigure(encoding="utf-8", errors="replace") + except Exception: + pass + logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s", handlers=[logging.StreamHandler(sys.stdout)], + force=True, ) logger = logging.getLogger(__name__) @@ -163,6 +182,12 @@ "VoxCPM2 is a creative multilingual TTS model from ModelBest, " "designed to generate highly realistic speech." ) +ASR_BACKENDS = {"auto", "sensevoice", "parakeet"} +PARAKEET_ASR_MODEL_ID = "nvidia/parakeet-tdt-0.6b-v3" +PARAKEET_LOCAL_MODEL_DIRNAME = PARAKEET_ASR_MODEL_ID.replace("/", "__") +ProgressCallback = Optional[Callable[[float, str], None]] +GenerationProgressCallback = Optional[Callable[[int, int], None]] + _CUSTOM_CSS = """ .logo-container { text-align: center; @@ -224,53 +249,301 @@ # ---------- Model ---------- +def _coerce_audio_filepath(audio_input) -> Optional[str]: + if audio_input is None or audio_input == "": + return None + if isinstance(audio_input, (str, os.PathLike)): + return os.fspath(audio_input) + if isinstance(audio_input, dict): + path = audio_input.get("path") + return os.fspath(path) if path else None + path = getattr(audio_input, "path", None) + if path: + return os.fspath(path) + return str(audio_input) + + +def _extract_asr_text(result) -> str: + if not result: + return "" + first = result[0] if isinstance(result, list) else result + raw_text = first.get("text", "") if isinstance(first, dict) else str(first) + return re.sub(r"<\|.*?\|>", "", raw_text).strip() + + +def _extract_parakeet_asr_text(result) -> str: + if not result: + return "" + if isinstance(result, str): + return result.strip() + if isinstance(result, (list, tuple)): + return " ".join(str(item).strip() for item in result if str(item).strip()).strip() + return str(result).strip() + + +def _prepare_asr_audio(audio_path: str, sample_rate: int = 16000) -> Tuple[str, Optional[str]]: + """Return an ASR-friendly 16 kHz mono file and optional temp path to remove.""" + import librosa + import soundfile as sf + + source_path = os.fspath(audio_path) + try: + info = sf.info(source_path) + suffix = Path(source_path).suffix.lower() + if info.samplerate == sample_rate and info.channels == 1 and suffix in {".wav", ".flac"}: + return source_path, None + except Exception: + pass + + audio, _ = librosa.load(source_path, sr=sample_rate, mono=True) + if audio.size == 0: + raise ValueError("Reference audio contains no readable samples.") + + with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as tmp: + temp_path = tmp.name + sf.write(temp_path, audio, sample_rate, subtype="PCM_16") + return temp_path, temp_path + + +def _normalize_asr_backend(asr_backend: str) -> str: + backend = (asr_backend or "auto").strip().lower() + if backend not in ASR_BACKENDS: + raise ValueError(f"Unknown ASR backend: {asr_backend!r}. Expected one of: {', '.join(sorted(ASR_BACKENDS))}.") + return backend + + +def _emit_progress(callback: ProgressCallback, fraction: float, message: str) -> None: + if callback is None: + return + callback(max(0.0, min(1.0, fraction)), message) + + +def _resolve_generation_inputs( + demo, + ref_wav, + use_prompt_text: bool, + prompt_text_value: str, + control_instruction: str, + progress_callback: ProgressCallback = None, +) -> Tuple[Optional[str], str, str]: + audio_path = _coerce_audio_filepath(ref_wav) + actual_prompt_text = (prompt_text_value or "").strip() if use_prompt_text else "" + if use_prompt_text: + if not audio_path: + raise gr.Error("Upload reference audio before using Ultimate Cloning Mode.") + if not actual_prompt_text: + logger.info("Auto-transcribing reference audio before generation...") + actual_prompt_text = demo.prompt_wav_recognition(audio_path, progress_callback=progress_callback).strip() + if not actual_prompt_text: + raise gr.Error( + "Auto-transcription returned no text. Enter the reference transcript or disable Ultimate Cloning Mode." + ) + return audio_path, actual_prompt_text, "" + return audio_path, "", control_instruction + + class VoxCPMDemo: - def __init__(self, model_id: str = "openbmb/VoxCPM2", device: str = "auto") -> None: + def __init__(self, model_id: str = "openbmb/VoxCPM2", device: str = "auto", asr_backend: str = "auto") -> None: self.device = resolve_runtime_device(device, "cuda") logger.info(f"Running VoxCPM on device: {self.device}") self.optimize = self.device.startswith("cuda") - - self.asr_model_id = "iic/SenseVoiceSmall" + self.asr_backend = _normalize_asr_backend(os.environ.get("VOXCPM_ASR_BACKEND", asr_backend)) + + project_root = Path(__file__).parent + local_asr_model = project_root / "models" / "iic__SenseVoiceSmall" + local_parakeet_model = project_root / "models" / PARAKEET_LOCAL_MODEL_DIRNAME + local_zipenhancer_model = project_root / "models" / "iic__speech_zipenhancer_ans_multiloss_16k_base" + + self.asr_model_id = str(local_asr_model) if local_asr_model.exists() else "iic/SenseVoiceSmall" + self.parakeet_model_id = str(local_parakeet_model) if local_parakeet_model.exists() else None + self.zipenhancer_model_id = ( + str(local_zipenhancer_model) + if local_zipenhancer_model.exists() + else "iic/speech_zipenhancer_ans_multiloss_16k_base" + ) self.asr_device = "cuda:0" if self.device.startswith("cuda") else "cpu" self.asr_model: Optional[AutoModel] = None + self.parakeet_processor = None + self.parakeet_model = None + self._voxcpm_load_lock = threading.RLock() + self._asr_load_lock = threading.RLock() + self._parakeet_load_lock = threading.RLock() + logger.info("ASR backend: %s", self._resolved_asr_backend_name()) self.voxcpm_model: Optional[voxcpm.VoxCPM] = None self._model_id = model_id + def _get_load_lock(self, attr_name: str): + lock = getattr(self, attr_name, None) + if lock is None: + lock = threading.RLock() + setattr(self, attr_name, lock) + return lock + + def asr_status_text(self) -> str: + backend = self._resolved_asr_backend_name() + if backend == "parakeet": + model_name = "NVIDIA Parakeet TDT 0.6B v3" + model_path = self.parakeet_model_id or "not installed" + device = "cuda" if self.device.startswith("cuda") else "cpu" + else: + model_name = "SenseVoiceSmall" + model_path = self.asr_model_id + device = self.asr_device + return f"ASR: {model_name} | language: auto-detect | device: {device} | model: {model_path}" + + def preload_models( + self, *, preload_asr: bool = True, preload_tts: bool = True, preload_denoiser: bool = True + ) -> None: + logger.info("Preloading models...") + if preload_tts: + logger.info("Preloading VoxCPM TTS model...") + current_model = self.get_or_load_voxcpm() + if preload_denoiser: + logger.info("Preloading ZipEnhancer denoiser...") + current_model._get_or_load_denoiser() + if preload_asr: + if self._should_use_parakeet_asr(): + logger.info("Preloading Parakeet ASR model (language=auto-detect)...") + self.get_or_load_parakeet_asr_model() + if self.asr_backend == "auto": + logger.info("Preloading SenseVoice ASR fallback (language=auto)...") + self.get_or_load_asr_model() + else: + logger.info("Preloading SenseVoice ASR model (language=auto)...") + self.get_or_load_asr_model() + logger.info("Preload complete.") + def get_or_load_voxcpm(self) -> voxcpm.VoxCPM: if self.voxcpm_model is not None: return self.voxcpm_model - logger.info(f"Loading model: {self._model_id}") - self.voxcpm_model = voxcpm.VoxCPM.from_pretrained( - self._model_id, - optimize=self.optimize, - device=self.device, - ) - logger.info("Model loaded successfully.") - return self.voxcpm_model + with self._get_load_lock("_voxcpm_load_lock"): + if self.voxcpm_model is not None: + return self.voxcpm_model + logger.info(f"Loading model: {self._model_id}") + self.voxcpm_model = voxcpm.VoxCPM.from_pretrained( + self._model_id, + zipenhancer_model_id=self.zipenhancer_model_id, + optimize=self.optimize, + device=self.device, + ) + logger.info("Model loaded successfully.") + return self.voxcpm_model def get_or_load_asr_model(self) -> AutoModel: if self.asr_model is not None: return self.asr_model - logger.info(f"Loading ASR model: {self.asr_model_id} on device: {self.asr_device}") - self.asr_model = AutoModel( - model=self.asr_model_id, - disable_update=True, - log_level="DEBUG", - device=self.asr_device, - ) - logger.info("ASR model loaded successfully.") - return self.asr_model + with self._get_load_lock("_asr_load_lock"): + if self.asr_model is not None: + return self.asr_model + logger.info(f"Loading ASR model: {self.asr_model_id} on device: {self.asr_device}") + self.asr_model = AutoModel( + model=self.asr_model_id, + disable_update=True, + log_level="DEBUG", + device=self.asr_device, + ) + logger.info("ASR model loaded successfully.") + return self.asr_model - def prompt_wav_recognition(self, prompt_wav: Optional[str]) -> str: - if prompt_wav is None: - return "" + def _should_use_parakeet_asr(self) -> bool: + if self.asr_backend == "sensevoice": + return False + if self.asr_backend == "parakeet": + return True + return self.device.startswith("cuda") and self.parakeet_model_id is not None + + def _resolved_asr_backend_name(self) -> str: + if self._should_use_parakeet_asr(): + return "parakeet" + return "sensevoice" + + def get_or_load_parakeet_asr_model(self): + if self.parakeet_processor is not None and self.parakeet_model is not None: + return self.parakeet_processor, self.parakeet_model + with self._get_load_lock("_parakeet_load_lock"): + if self.parakeet_processor is not None and self.parakeet_model is not None: + return self.parakeet_processor, self.parakeet_model + if self.parakeet_model_id is None: + raise RuntimeError( + "NVIDIA Parakeet ASR is not installed locally. Run install.bat to pre-download it, " + "or start app.py with --asr-backend sensevoice." + ) + try: + import torch + from transformers import AutoModelForTDT, AutoProcessor + except ImportError as exc: + raise RuntimeError( + "NVIDIA Parakeet ASR requires a Transformers build with AutoModelForTDT support." + ) from exc + + target_device = "cuda" if self.device.startswith("cuda") else "cpu" + logger.info("Loading Parakeet ASR model: %s on device: %s", self.parakeet_model_id, target_device) + self.parakeet_processor = AutoProcessor.from_pretrained(self.parakeet_model_id, local_files_only=True) + self.parakeet_model = AutoModelForTDT.from_pretrained( + self.parakeet_model_id, + dtype="auto", + local_files_only=True, + ) + self.parakeet_model.to(target_device) + self.parakeet_model.eval() + logger.info("Parakeet ASR model loaded successfully.") + return self.parakeet_processor, self.parakeet_model + + def _recognize_with_sensevoice(self, asr_audio_path: str, progress_callback: ProgressCallback = None) -> str: + _emit_progress(progress_callback, 0.45, "Transcribing reference audio with SenseVoice, language auto, 45%") + logger.info("Running SenseVoice ASR with language=auto on device: %s", self.asr_device) res = self.get_or_load_asr_model().generate( - input=prompt_wav, + input=asr_audio_path, language="auto", use_itn=True, ) - return res[0]["text"].split("|>")[-1] + _emit_progress(progress_callback, 0.95, "Transcribing reference audio, 95%") + return _extract_asr_text(res) + + def _recognize_with_parakeet(self, asr_audio_path: str, progress_callback: ProgressCallback = None) -> str: + import librosa + import torch + + _emit_progress(progress_callback, 0.25, "Loading Parakeet ASR, language auto-detect, 25%") + processor, model = self.get_or_load_parakeet_asr_model() + sample_rate = getattr(processor.feature_extractor, "sampling_rate", 16000) + _emit_progress(progress_callback, 0.40, "Preparing Parakeet audio features, 40%") + audio, _ = librosa.load(asr_audio_path, sr=sample_rate, mono=True) + if audio.size == 0: + return "" + inputs = processor([audio], sampling_rate=sample_rate) + inputs.to(model.device, dtype=model.dtype) + logger.info("Running Parakeet ASR with language=auto-detect on device: %s", model.device) + _emit_progress(progress_callback, 0.55, "Transcribing reference audio with Parakeet, 55%") + with torch.inference_mode(): + output = model.generate(**inputs, return_dict_in_generate=True) + sequences = getattr(output, "sequences", output) + _emit_progress(progress_callback, 0.95, "Transcribing reference audio, 95%") + return _extract_parakeet_asr_text(processor.decode(sequences, skip_special_tokens=True)) + + def prompt_wav_recognition(self, prompt_wav: Optional[str], progress_callback: ProgressCallback = None) -> str: + prompt_wav_path = _coerce_audio_filepath(prompt_wav) + if prompt_wav_path is None: + return "" + _emit_progress(progress_callback, 0.05, "Transcribing reference audio, 5%") + asr_audio_path, temp_path = _prepare_asr_audio(prompt_wav_path) + _emit_progress(progress_callback, 0.15, "Prepared 16 kHz mono ASR audio, 15%") + try: + if self._should_use_parakeet_asr(): + try: + return self._recognize_with_parakeet(asr_audio_path, progress_callback) + except Exception: + if self.asr_backend == "parakeet": + raise + logger.warning("Parakeet ASR failed; falling back to SenseVoice.", exc_info=True) + return self._recognize_with_sensevoice(asr_audio_path, progress_callback) + finally: + if temp_path and os.path.exists(temp_path): + try: + os.unlink(temp_path) + except OSError: + pass def _build_generate_kwargs( self, @@ -283,6 +556,7 @@ def _build_generate_kwargs( denoise: bool, inference_timesteps: int = 10, seed: Optional[int] = None, + progress_callback: GenerationProgressCallback = None, ) -> dict: generate_kwargs = dict( text=final_text, @@ -293,6 +567,8 @@ def _build_generate_kwargs( denoise=denoise, seed=seed, ) + if progress_callback is not None: + generate_kwargs["progress_callback"] = progress_callback if prompt_text_clean and audio_path: generate_kwargs["prompt_wav_path"] = audio_path generate_kwargs["prompt_text"] = prompt_text_clean @@ -309,6 +585,7 @@ def generate_tts_audio( denoise: bool = True, inference_timesteps: int = 10, seed: Optional[int] = None, + progress_callback: GenerationProgressCallback = None, ) -> Tuple[int, np.ndarray, Optional[int]]: current_model = self.get_or_load_voxcpm() @@ -322,7 +599,7 @@ def generate_tts_audio( control = re.sub(r"[()()]", "", control).strip() final_text = f"({control}){text}" if control else text - audio_path = reference_wav_path_input if reference_wav_path_input else None + audio_path = _coerce_audio_filepath(reference_wav_path_input) prompt_text_clean = (prompt_text or "").strip() or None if audio_path and prompt_text_clean: @@ -342,6 +619,7 @@ def generate_tts_audio( denoise=denoise, inference_timesteps=inference_timesteps, seed=seed, + progress_callback=progress_callback, ) wav = current_model.generate(**generate_kwargs) last_successful_seed = getattr(current_model.tts_model, "last_successful_seed", seed) @@ -352,7 +630,7 @@ def generate_tts_audio( def create_demo_interface(demo: VoxCPMDemo): - gr.set_static_paths(paths=[Path.cwd().absolute() / "assets"]) + gr.set_static_paths(paths=[PROJECT_ROOT / "assets"]) def _coerce_seed(seed_value) -> Optional[int]: if seed_value is None or seed_value == "": @@ -367,6 +645,9 @@ def _prepare_seed(use_random_seed: bool, seed_value): def _on_random_seed_toggle(checked): return gr.update(interactive=not checked) + def _gradio_progress_callback(progress): + return lambda fraction, message: progress(fraction, desc=message) + def _generate( text: str, control_instruction: str, @@ -378,28 +659,59 @@ def _generate( denoise: bool, dit_steps: int, seed_value, + progress=gr.Progress(track_tqdm=True), ): - actual_prompt_text = prompt_text_value.strip() if use_prompt_text else "" - actual_control = "" if use_prompt_text else control_instruction + progress(0.02, desc="Preparing generation, 2%") + + def asr_progress(fraction: float, message: str) -> None: + mapped = 0.03 + (0.24 * max(0.0, min(1.0, fraction))) + progress(mapped, desc=f"{message} / preparing generation, {int(mapped * 100)}%") + + audio_path, actual_prompt_text, actual_control = _resolve_generation_inputs( + demo, + ref_wav, + use_prompt_text, + prompt_text_value, + control_instruction, + progress_callback=asr_progress, + ) seed = _coerce_seed(seed_value) + + def tts_progress(step: int, total: int) -> None: + if total <= 0: + return + fraction = min(1.0, max(0.0, (step + 1) / total)) + mapped = 0.35 + (0.55 * fraction) + progress(mapped, desc=f"Synthesising speech, {int(mapped * 100)}%") + + progress(0.30, desc="Preparing voice prompt, 30%") sr, wav_np, last_successful_seed = demo.generate_tts_audio( text_input=text, control_instruction=actual_control, - reference_wav_path_input=ref_wav, + reference_wav_path_input=audio_path, prompt_text=actual_prompt_text, cfg_value_input=cfg_value, do_normalize=do_normalize, denoise=denoise, inference_timesteps=int(dit_steps), seed=seed, + progress_callback=tts_progress, ) - return (sr, wav_np), last_successful_seed + progress(0.95, desc="Finalising audio, 95%") + progress(1.0, desc="Complete, 100%") + return (sr, wav_np), last_successful_seed, actual_prompt_text if use_prompt_text else gr.update() - def _on_toggle_instant(checked): + def _on_toggle_instant(checked, current_prompt_text, audio_path): """Instant UI toggle — no ASR, no blocking.""" + current_prompt_text = current_prompt_text or "" if checked: + placeholder = ( + "Recognizing reference audio..." + if _coerce_audio_filepath(audio_path) and not current_prompt_text.strip() + else I18N("prompt_text_placeholder") + ) return ( - gr.update(visible=True, value="", placeholder="Recognizing reference audio..."), + gr.update(visible=True, value=current_prompt_text, placeholder=placeholder), gr.update(visible=False), ) return ( @@ -407,20 +719,59 @@ def _on_toggle_instant(checked): gr.update(visible=True, interactive=True), ) - def _run_asr_if_needed(checked, audio_path): + def _on_reference_audio_change(checked, current_prompt_text, audio_path): + if not checked or not _coerce_audio_filepath(audio_path): + return gr.update() + return gr.update( + visible=True, + value=current_prompt_text or "", + placeholder="Recognizing reference audio...", + ) + + def _run_asr_if_needed(checked, audio_path, progress=gr.Progress(track_tqdm=True)): """Run ASR after the UI has updated. Only when toggled ON.""" - if not checked or not audio_path: + audio_file = _coerce_audio_filepath(audio_path) + if not checked or not audio_file: return gr.update() try: - logger.info("Running ASR on reference audio...") - asr_text = demo.prompt_wav_recognition(audio_path) - logger.info(f"ASR result: {asr_text[:60]}...") - return gr.update(value=asr_text) + logger.info("Running ASR on reference audio using %s...", demo.asr_status_text()) + asr_text = demo.prompt_wav_recognition( + audio_file, + progress_callback=_gradio_progress_callback(progress), + ) + logger.info("ASR result: %r", asr_text[:60]) + if not asr_text: + progress(1.0, desc="Transcribing reference audio complete, 100%") + return gr.update( + value="", + placeholder="No speech was recognized. Enter the reference transcript manually.", + ) + progress(1.0, desc="Transcribing reference audio complete, 100%") + return gr.update(value=asr_text, placeholder=I18N("prompt_text_placeholder")) except Exception as e: - logger.warning(f"ASR recognition failed: {e}") - return gr.update(value="") + logger.warning("ASR recognition failed: %s", e, exc_info=True) + return gr.update(value="", placeholder=f"ASR failed: {e}") + + def _ensure_prompt_text_before_generate( + ref_wav, use_prompt_text, prompt_text_value, progress=gr.Progress(track_tqdm=True) + ): + if not use_prompt_text: + return gr.update() + progress(0.02, desc="Preparing reference transcript, 2%") + audio_path, actual_prompt_text, _ = _resolve_generation_inputs( + demo, + ref_wav, + True, + prompt_text_value, + "", + progress_callback=_gradio_progress_callback(progress), + ) + if not audio_path: + raise gr.Error("Upload reference audio before using Ultimate Cloning Mode.") + progress(1.0, desc="Reference transcript ready, 100%") + return gr.update(value=actual_prompt_text, placeholder=I18N("prompt_text_placeholder")) - with gr.Blocks() as interface: + with gr.Blocks(theme=_APP_THEME, css=_CUSTOM_CSS) as interface: gr.HTML( '
' 'VoxCPM Logo' @@ -442,6 +793,7 @@ def _run_asr_if_needed(checked, audio_path): info=I18N("show_prompt_text_info"), elem_classes=["switch-toggle"], ) + gr.Markdown(demo.asr_status_text()) prompt_text = gr.Textbox( value="", label=I18N("prompt_text_label"), @@ -513,7 +865,7 @@ def _run_asr_if_needed(checked, audio_path): show_prompt_text.change( fn=_on_toggle_instant, - inputs=[show_prompt_text], + inputs=[show_prompt_text, prompt_text, reference_wav], outputs=[prompt_text, control_instruction], ).then( fn=_run_asr_if_needed, @@ -521,6 +873,16 @@ def _run_asr_if_needed(checked, audio_path): outputs=[prompt_text], ) + reference_wav.change( + fn=_on_reference_audio_change, + inputs=[show_prompt_text, prompt_text, reference_wav], + outputs=[prompt_text], + ).then( + fn=_run_asr_if_needed, + inputs=[show_prompt_text, reference_wav], + outputs=[prompt_text], + ) + random_seed.change( fn=_on_random_seed_toggle, inputs=[random_seed], @@ -532,6 +894,10 @@ def _run_asr_if_needed(checked, audio_path): inputs=[random_seed, seed_value], outputs=[seed_value], show_progress=False, + ).then( + fn=_ensure_prompt_text_before_generate, + inputs=[reference_wav, show_prompt_text, prompt_text], + outputs=[prompt_text], ).then( fn=_generate, inputs=[ @@ -546,7 +912,7 @@ def _run_asr_if_needed(checked, audio_path): dit_steps, seed_value, ], - outputs=[audio_output, seed_value], + outputs=[audio_output, seed_value, prompt_text], show_progress=True, api_name="generate", ) @@ -555,21 +921,27 @@ def _run_asr_if_needed(checked, audio_path): def run_demo( - server_name: str = "0.0.0.0", + server_name: str = "127.0.0.1", server_port: int = 8808, show_error: bool = True, model_id: str = "openbmb/VoxCPM2", device: str = "auto", + asr_backend: str = "auto", + preload: bool = True, + preload_denoiser: bool = True, + open_browser: bool = True, ): - demo = VoxCPMDemo(model_id=model_id, device=device) + demo = VoxCPMDemo(model_id=model_id, device=device, asr_backend=asr_backend) + if preload: + demo.preload_models(preload_asr=True, preload_tts=True, preload_denoiser=preload_denoiser) interface = create_demo_interface(demo) + logger.info("Launching web UI at http://%s:%s", server_name, server_port) interface.queue(max_size=10, default_concurrency_limit=1).launch( server_name=server_name, server_port=server_port, show_error=show_error, + inbrowser=open_browser, i18n=I18N, - theme=_APP_THEME, - css=_CUSTOM_CSS, ) @@ -583,6 +955,12 @@ def run_demo( default="openbmb/VoxCPM2", help="Local path or HuggingFace repo ID (default: openbmb/VoxCPM2)", ) + parser.add_argument( + "--host", + type=str, + default="127.0.0.1", + help="Server host/interface (default: 127.0.0.1; use 0.0.0.0 for LAN access)", + ) parser.add_argument("--port", type=int, default=8808, help="Server port") parser.add_argument( "--device", @@ -590,5 +968,43 @@ def run_demo( default="auto", help="Runtime device: auto, cpu, mps, cuda, or cuda:N (default: auto)", ) + parser.add_argument( + "--asr-backend", + type=str, + default="auto", + choices=sorted(ASR_BACKENDS), + help="Reference audio transcription backend: auto, sensevoice, or parakeet (default: auto)", + ) + parser.add_argument( + "--no-preload", + action="store_true", + help="Skip loading models before launching the web UI.", + ) + parser.add_argument( + "--no-browser", + action="store_true", + help="Start the web UI without opening a browser window.", + ) + parser.add_argument( + "--preload-denoiser", + action="store_true", + default=True, + help="Deprecated: ZipEnhancer is loaded before launch by default.", + ) + parser.add_argument( + "--no-preload-denoiser", + action="store_false", + dest="preload_denoiser", + help="Skip loading ZipEnhancer before launching the web UI.", + ) args = parser.parse_args() - run_demo(model_id=args.model_id, server_port=args.port, device=args.device) + run_demo( + model_id=args.model_id, + server_name=args.host, + server_port=args.port, + device=args.device, + asr_backend=args.asr_backend, + preload=not args.no_preload, + preload_denoiser=args.preload_denoiser, + open_browser=not args.no_browser, + ) diff --git a/install.bat b/install.bat new file mode 100644 index 0000000..17aa076 --- /dev/null +++ b/install.bat @@ -0,0 +1,446 @@ +@echo off +setlocal EnableExtensions EnableDelayedExpansion + +set "ROOT=%~dp0" +cd /d "%ROOT%" || exit /b 1 + +set "VENV_DIR=.venv" +set "INSTALL_DEV=1" +set "INSTALL_TIMESTAMPS=1" +set "DOWNLOAD_MODEL=1" +set "DOWNLOAD_MS_MODELS=1" +set "DOWNLOAD_PARAKEET_MODEL=1" +set "DOWNLOAD_TIMESTAMP_MODEL=1" +set "RUN_SMOKE_CHECKS=1" +set "DRY_RUN=0" +set "TORCH_BACKEND=auto" +set "PYTORCH_INDEX_URL=" +set "MODEL_ID=openbmb/VoxCPM2" +set "MODEL_DIR=models\openbmb__VoxCPM2" +set "PARAKEET_MODEL_ID=nvidia/parakeet-tdt-0.6b-v3" +set "PARAKEET_MODEL_DIR=models\nvidia__parakeet-tdt-0.6b-v3" +set "ZIPENHANCER_MODEL_DIR=models\iic__speech_zipenhancer_ans_multiloss_16k_base" +set "ASR_MODEL_DIR=models\iic__SenseVoiceSmall" +set "PYTHON_CMD=" + +goto parse_args + +:usage +echo VoxCPM Windows installer +echo. +echo Usage: +echo install.bat [options] +echo. +echo Options: +echo --cuda Force CUDA-enabled torch/torchaudio wheels. +echo --cpu Force CPU torch/torchaudio wheels. +echo --pytorch-index-url U Use a custom PyTorch wheel index URL. +echo --model-id ID Hugging Face model to download (default: openbmb/VoxCPM2). +echo --model-dir DIR Local model directory (default: models\openbmb__VoxCPM2). +echo --download-parakeet Pre-download NVIDIA Parakeet ASR (enabled by default). +echo --skip-parakeet Skip NVIDIA Parakeet ASR pre-download. +echo --parakeet-model-dir D Local Parakeet ASR directory (default: models\nvidia__parakeet-tdt-0.6b-v3). +echo --skip-models Skip all model pre-downloads. +echo --skip-modelscope Skip ModelScope denoiser and ASR model pre-downloads. +echo --skip-timestamp-model Skip stable-ts Whisper base model pre-download. +echo --no-dev Skip developer/test tools (installed by default). +echo --no-timestamps Skip stable-ts timestamp dependencies (installed by default). +echo --no-smoke-checks Skip import/CLI validation after install. +echo --venv DIR Use a different virtual environment directory. +echo --dry-run Print the planned actions without installing. +echo -h, --help Show this help. +echo. +echo Environment: +echo PYTHON Optional path to python.exe, Python 3.10-3.12 required. +echo. +echo Examples: +echo install.bat +echo install.bat --cuda +echo install.bat --cpu --model-dir D:\models\VoxCPM2 +exit /b 0 + +:parse_args +if "%~1"=="" goto args_done +if /I "%~1"=="-h" goto usage +if /I "%~1"=="--help" goto usage +if /I "%~1"=="--cuda" ( + set "TORCH_BACKEND=cuda" + shift + goto parse_args +) +if /I "%~1"=="--cpu" ( + set "TORCH_BACKEND=cpu" + shift + goto parse_args +) +if /I "%~1"=="--download-model" ( + set "DOWNLOAD_MODEL=1" + shift + goto parse_args +) +if /I "%~1"=="--download-parakeet" ( + set "DOWNLOAD_PARAKEET_MODEL=1" + shift + goto parse_args +) +if /I "%~1"=="--skip-parakeet" ( + set "DOWNLOAD_PARAKEET_MODEL=0" + shift + goto parse_args +) +if /I "%~1"=="--skip-models" ( + set "DOWNLOAD_MODEL=0" + set "DOWNLOAD_MS_MODELS=0" + set "DOWNLOAD_PARAKEET_MODEL=0" + set "DOWNLOAD_TIMESTAMP_MODEL=0" + shift + goto parse_args +) +if /I "%~1"=="--skip-modelscope" ( + set "DOWNLOAD_MS_MODELS=0" + shift + goto parse_args +) +if /I "%~1"=="--skip-timestamp-model" ( + set "DOWNLOAD_TIMESTAMP_MODEL=0" + shift + goto parse_args +) +if /I "%~1"=="--no-dev" ( + set "INSTALL_DEV=0" + shift + goto parse_args +) +if /I "%~1"=="--no-timestamps" ( + set "INSTALL_TIMESTAMPS=0" + shift + goto parse_args +) +if /I "%~1"=="--dry-run" ( + set "DRY_RUN=1" + shift + goto parse_args +) +if /I "%~1"=="--no-smoke-checks" ( + set "RUN_SMOKE_CHECKS=0" + shift + goto parse_args +) +if /I "%~1"=="--venv" goto parse_venv +if /I "%~1"=="--pytorch-index-url" goto parse_pytorch_index +if /I "%~1"=="--model-id" goto parse_model_id +if /I "%~1"=="--model-dir" goto parse_model_dir +if /I "%~1"=="--parakeet-model-dir" goto parse_parakeet_model_dir + +echo Unknown option: %~1 +echo Run install.bat --help for usage. +exit /b 2 + +:parse_venv +shift +if "%~1"=="" goto arg_error +set "VENV_DIR=%~1" +shift +goto parse_args + +:parse_pytorch_index +shift +if "%~1"=="" goto arg_error +set "PYTORCH_INDEX_URL=%~1" +set "TORCH_BACKEND=custom" +shift +goto parse_args + +:parse_model_id +shift +if "%~1"=="" goto arg_error +set "MODEL_ID=%~1" +shift +goto parse_args + +:parse_model_dir +shift +if "%~1"=="" goto arg_error +set "MODEL_DIR=%~1" +shift +goto parse_args + +:parse_parakeet_model_dir +shift +if "%~1"=="" goto arg_error +set "PARAKEET_MODEL_DIR=%~1" +shift +goto parse_args + +:arg_error +echo Missing value for the previous option. +echo Run install.bat --help for usage. +exit /b 2 + +:args_done +if "%INSTALL_TIMESTAMPS%"=="0" set "DOWNLOAD_TIMESTAMP_MODEL=0" +if /I "%TORCH_BACKEND%"=="auto" ( + call :detect_torch_backend +) +if /I "%TORCH_BACKEND%"=="cuda" if not defined PYTORCH_INDEX_URL set "PYTORCH_INDEX_URL=https://download.pytorch.org/whl/cu121" +if /I "%TORCH_BACKEND%"=="cpu" if not defined PYTORCH_INDEX_URL set "PYTORCH_INDEX_URL=https://download.pytorch.org/whl/cpu" +if /I "%DOWNLOAD_PARAKEET_MODEL%"=="auto" set "DOWNLOAD_PARAKEET_MODEL=1" + +set "PROJECT_SPEC=." +if "%INSTALL_TIMESTAMPS%"=="1" if "%INSTALL_DEV%"=="1" set "PROJECT_SPEC=.[timestamps,dev]" +if "%INSTALL_TIMESTAMPS%"=="1" if "%INSTALL_DEV%"=="0" set "PROJECT_SPEC=.[timestamps]" +if "%INSTALL_TIMESTAMPS%"=="0" if "%INSTALL_DEV%"=="1" set "PROJECT_SPEC=.[dev]" + +echo. +echo VoxCPM setup +echo Root: %CD% +echo Venv: %VENV_DIR% +echo Project: %PROJECT_SPEC% +echo Torch: %TORCH_BACKEND% +if defined PYTORCH_INDEX_URL echo Torch index: %PYTORCH_INDEX_URL% +if "%DOWNLOAD_MODEL%"=="1" echo HF model: %MODEL_ID% -^> %MODEL_DIR% +if "%DOWNLOAD_PARAKEET_MODEL%"=="1" echo Parakeet: %PARAKEET_MODEL_ID% -^> %PARAKEET_MODEL_DIR% +if "%DOWNLOAD_MS_MODELS%"=="1" echo MS models: local denoiser + ASR models +if "%DOWNLOAD_TIMESTAMP_MODEL%"=="1" echo TS model: stable-ts Whisper base +if "%DRY_RUN%"=="1" echo Mode: dry run + +call :find_python +if errorlevel 1 goto fail + +echo Python: !PYTHON_CMD! + +if not exist "%VENV_DIR%\Scripts\python.exe" ( + echo. + echo ^> !PYTHON_CMD! -m venv "%VENV_DIR%" + if not "%DRY_RUN%"=="1" ( + !PYTHON_CMD! -m venv "%VENV_DIR%" + if errorlevel 1 goto fail + ) +) else ( + echo. + echo Reusing existing virtual environment: %VENV_DIR% +) + +if "%DRY_RUN%"=="1" ( + echo. + echo ^> call "%VENV_DIR%\Scripts\activate.bat" +) else ( + call "%VENV_DIR%\Scripts\activate.bat" + if errorlevel 1 goto fail +) + +call :run python -m pip install --upgrade pip +if errorlevel 1 goto fail + +call :run python -m pip install --upgrade uv +if errorlevel 1 goto fail + +if exist "%VENV_DIR%\Scripts\uv.exe" ( + set "UV_CMD=%VENV_DIR%\Scripts\uv.exe" +) else ( + set "UV_CMD=uv" +) + +if defined PYTORCH_INDEX_URL ( + call :run "!UV_CMD!" pip install --upgrade torch torchaudio --index-url "%PYTORCH_INDEX_URL%" + if errorlevel 1 ( + echo. + echo PyTorch wheel install failed; retrying with pip. + call :run python -m pip install --upgrade torch torchaudio --index-url "%PYTORCH_INDEX_URL%" + if errorlevel 1 goto fail + ) +) + +call :run "!UV_CMD!" pip install -e "%PROJECT_SPEC%" +if errorlevel 1 ( + echo. + echo uv install failed; retrying with pip. + call :run python -m pip install -e "%PROJECT_SPEC%" + if errorlevel 1 goto fail +) + +if not "%DRY_RUN%"=="1" ( + if not exist "models" mkdir "models" + if not exist "lora" mkdir "lora" + if not exist "outputs" mkdir "outputs" + if not exist "checkpoints" mkdir "checkpoints" +) else ( + echo. + echo Would create runtime directories: models, lora, outputs, checkpoints +) + +if not "%DOWNLOAD_MODEL%"=="1" goto skip_hf_download +echo. +echo ^> Downloading %MODEL_ID% to %MODEL_DIR% +if "%DRY_RUN%"=="1" goto skip_hf_download +python -c "from huggingface_hub import snapshot_download; snapshot_download(repo_id='%MODEL_ID%', local_dir=r'%MODEL_DIR%')" +if errorlevel 1 goto fail + +:skip_hf_download +if not "%DOWNLOAD_PARAKEET_MODEL%"=="1" goto skip_parakeet_download +echo. +echo ^> Downloading NVIDIA Parakeet ASR to %PARAKEET_MODEL_DIR% +if "%DRY_RUN%"=="1" goto skip_parakeet_download +python -c "from huggingface_hub import snapshot_download; snapshot_download(repo_id='%PARAKEET_MODEL_ID%', local_dir=r'%PARAKEET_MODEL_DIR%')" +if errorlevel 1 goto fail + +:skip_parakeet_download +if not "%DOWNLOAD_MS_MODELS%"=="1" goto skip_modelscope_downloads +echo. +echo ^> Downloading ModelScope denoiser to %ZIPENHANCER_MODEL_DIR% +if "%DRY_RUN%"=="1" goto skip_modelscope_asr_download +python -c "from modelscope import snapshot_download; snapshot_download('iic/speech_zipenhancer_ans_multiloss_16k_base', local_dir=r'%ZIPENHANCER_MODEL_DIR%')" +if errorlevel 1 goto fail + +:skip_modelscope_asr_download +echo. +echo ^> Downloading ModelScope ASR to %ASR_MODEL_DIR% +if "%DRY_RUN%"=="1" goto skip_modelscope_downloads +python -c "from modelscope import snapshot_download; snapshot_download('iic/SenseVoiceSmall', local_dir=r'%ASR_MODEL_DIR%')" +if errorlevel 1 goto fail + +:skip_modelscope_downloads +if not "%DOWNLOAD_TIMESTAMP_MODEL%"=="1" goto skip_timestamp_download +echo. +echo ^> Downloading stable-ts Whisper base model +if "%DRY_RUN%"=="1" goto skip_timestamp_download +python -c "import stable_whisper; stable_whisper.load_model('base')" +if errorlevel 1 goto fail + +:skip_timestamp_download + +if not "%RUN_SMOKE_CHECKS%"=="1" goto skip_smoke_checks +call :run python -m pip show voxcpm torch torchaudio gradio modelscope huggingface-hub +if errorlevel 1 goto fail +call :run python -c "import torch, torchaudio, gradio, voxcpm, soundfile, librosa, transformers, datasets, huggingface_hub, modelscope, safetensors, argbind, yaml, funasr, tensorboardX" +if errorlevel 1 goto fail +if not "%DOWNLOAD_PARAKEET_MODEL%"=="1" goto skip_parakeet_smoke +call :run python -c "from transformers import AutoModelForTDT, AutoProcessor; AutoProcessor.from_pretrained(r'%PARAKEET_MODEL_DIR%', local_files_only=True)" +if errorlevel 1 goto fail + +:skip_parakeet_smoke +if not "%INSTALL_TIMESTAMPS%"=="1" goto skip_timestamp_smoke +call :run python -c "import stable_whisper" +if errorlevel 1 goto fail + +:skip_timestamp_smoke +if not "%INSTALL_DEV%"=="1" goto skip_dev_smoke +call :run python -m pytest --version +if errorlevel 1 goto fail + +:skip_dev_smoke +call :run voxcpm --help +if errorlevel 1 goto fail +if /I not "%TORCH_BACKEND%"=="cuda" goto skip_smoke_checks +call :verify_cuda +if errorlevel 1 goto fail + +:skip_smoke_checks + +echo. +echo Install complete. +echo. +set "RUNTIME_DEVICE_ARG=" +if /I "%TORCH_BACKEND%"=="cuda" set "RUNTIME_DEVICE_ARG= --device cuda" +if /I "%TORCH_BACKEND%"=="cpu" set "RUNTIME_DEVICE_ARG= --device cpu" +echo Start commands: +echo %VENV_DIR%\Scripts\activate.bat +echo python app.py --model-id "%MODEL_DIR%" --port 8808%RUNTIME_DEVICE_ARG% --asr-backend auto +echo voxcpm --help +echo voxcpm design --model-path "%MODEL_DIR%"%RUNTIME_DEVICE_ARG% --text "Hello from VoxCPM2." --output outputs\demo.wav +echo python lora_ft_webui.py +echo. +echo Notes: +echo Web demo, CLI, and LoRA fine-tuning UI are installed. +if "%INSTALL_TIMESTAMPS%"=="1" echo Timestamp dependencies are installed. +if "%DOWNLOAD_MODEL%"=="1" echo Default Hugging Face model is installed at %MODEL_DIR%. +if "%DOWNLOAD_MODEL%"=="0" echo Hugging Face model pre-download was skipped. +if "%DOWNLOAD_PARAKEET_MODEL%"=="1" echo NVIDIA Parakeet ASR is installed at %PARAKEET_MODEL_DIR%. +if "%DOWNLOAD_PARAKEET_MODEL%"=="0" echo NVIDIA Parakeet ASR pre-download was skipped. +if "%DOWNLOAD_MS_MODELS%"=="1" echo Local ModelScope denoiser and ASR models are installed under models. +if "%DOWNLOAD_TIMESTAMP_MODEL%"=="1" echo stable-ts Whisper base model was cached. +echo CUDA is selected automatically when an NVIDIA GPU is detected; use --cpu to force CPU wheels. +exit /b 0 + +:detect_torch_backend +where nvidia-smi >nul 2>nul +if not errorlevel 1 ( + set "TORCH_BACKEND=cuda" + exit /b 0 +) + +powershell -NoProfile -ExecutionPolicy Bypass -Command "$gpus = Get-CimInstance Win32_VideoController; foreach ($gpu in $gpus) { if ($gpu.Name -match 'NVIDIA') { exit 0 } }; exit 1" >nul 2>nul +if not errorlevel 1 ( + set "TORCH_BACKEND=cuda" + exit /b 0 +) + +wmic path win32_VideoController get name 2>nul | findstr /I "NVIDIA" >nul 2>nul +if not errorlevel 1 ( + set "TORCH_BACKEND=cuda" + exit /b 0 +) + +set "TORCH_BACKEND=cpu" +exit /b 0 + +:find_python +if defined PYTHON ( + "%PYTHON%" -c "import sys; raise SystemExit(0 if (3, 10) <= sys.version_info[:2] < (3, 13) else 1)" >nul 2>nul + if not errorlevel 1 ( + set "PYTHON_CMD="%PYTHON%"" + exit /b 0 + ) +) + +py -3.12 -c "import sys; raise SystemExit(0 if (3, 10) <= sys.version_info[:2] < (3, 13) else 1)" >nul 2>nul +if not errorlevel 1 ( + set "PYTHON_CMD=py -3.12" + exit /b 0 +) + +py -3.11 -c "import sys; raise SystemExit(0 if (3, 10) <= sys.version_info[:2] < (3, 13) else 1)" >nul 2>nul +if not errorlevel 1 ( + set "PYTHON_CMD=py -3.11" + exit /b 0 +) + +py -3.10 -c "import sys; raise SystemExit(0 if (3, 10) <= sys.version_info[:2] < (3, 13) else 1)" >nul 2>nul +if not errorlevel 1 ( + set "PYTHON_CMD=py -3.10" + exit /b 0 +) + +python -c "import sys; raise SystemExit(0 if (3, 10) <= sys.version_info[:2] < (3, 13) else 1)" >nul 2>nul +if not errorlevel 1 ( + set "PYTHON_CMD=python" + exit /b 0 +) + +echo. +echo Python 3.10, 3.11, or 3.12 was not found. +echo Install Python from https://www.python.org/downloads/windows/ and rerun install.bat. +exit /b 1 + +:verify_cuda +echo. +echo ^> python -c "import torch; raise SystemExit(0 if torch.cuda.is_available() else 1)" +if "%DRY_RUN%"=="1" exit /b 0 +python -c "import torch; raise SystemExit(0 if torch.cuda.is_available() else 1)" +if errorlevel 1 ( + echo CUDA torch wheels were installed, but torch CUDA availability check returned false. + echo Check the NVIDIA driver, or rerun install.bat --cpu for CPU-only setup. +) +exit /b %ERRORLEVEL% + +:run +echo. +echo ^> %* +if "%DRY_RUN%"=="1" exit /b 0 +%* +exit /b %ERRORLEVEL% + +:fail +echo. +echo Installation failed. See the error above. +exit /b 1 diff --git a/lora_ft_webui.py b/lora_ft_webui.py index 3d91c3d..8752759 100644 --- a/lora_ft_webui.py +++ b/lora_ft_webui.py @@ -1324,4 +1324,4 @@ def change_language(lang): if __name__ == "__main__": # Ensure lora directory exists os.makedirs("lora", exist_ok=True) - app.queue().launch(server_name="0.0.0.0", server_port=7860) + app.queue().launch(server_name="0.0.0.0", server_port=7860, inbrowser=True) diff --git a/pyproject.toml b/pyproject.toml index 95659fa..3a9c1dc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,7 +28,7 @@ dependencies = [ "torch>=2.5.0", "torchaudio>=2.5.0", "torchcodec", - "transformers>=4.36.2", + "transformers>=5.12.0", "einops", "gradio>=6,<7", "inflect", diff --git a/scripts/test_pick_runtime_dtype.py b/scripts/test_pick_runtime_dtype.py index 160aba3..d5fce19 100644 --- a/scripts/test_pick_runtime_dtype.py +++ b/scripts/test_pick_runtime_dtype.py @@ -41,69 +41,78 @@ def expect_raises(fn, exc_type, label): return False -results = [] - -print("=== override set sanity ===") -results.append(expect("half" not in _VALID_DTYPE_OVERRIDES, True, "half removed from _VALID_DTYPE_OVERRIDES")) -results.append(expect("half" not in _LOW_PRECISION_DTYPES, True, "half removed from _LOW_PRECISION_DTYPES")) - -print("\n=== every accepted override parses through get_dtype ===") -for dt in sorted(_VALID_DTYPE_OVERRIDES): - try: - torch_dtype = get_dtype(dt) - print(f"[OK ] get_dtype({dt!r}) -> {torch_dtype}") - results.append(True) - except Exception as e: - print(f"[FAIL] get_dtype({dt!r}) raised: {e}") - results.append(False) - -print("\n=== pick_runtime_dtype: non-mps is a no-op ===") -results.append(expect(pick_runtime_dtype("cuda", "bfloat16"), "bfloat16", "cuda/bf16 untouched")) -results.append(expect(pick_runtime_dtype("cpu", "float16"), "float16", "cpu/fp16 untouched")) -results.append(expect(pick_runtime_dtype("cuda", "float32"), "float32", "cuda/fp32 untouched")) - -print("\n=== pick_runtime_dtype: mps forces fp32 for low-precision ===") -os.environ.pop("VOXCPM_MPS_DTYPE", None) -results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "float32", "mps/bf16 -> fp32")) -results.append(expect(pick_runtime_dtype("mps", "bf16"), "float32", "mps/bf16-alias -> fp32")) -results.append(expect(pick_runtime_dtype("mps", "float16"), "float32", "mps/fp16 -> fp32")) -results.append(expect(pick_runtime_dtype("mps", "fp16"), "float32", "mps/fp16-alias -> fp32")) -results.append(expect(pick_runtime_dtype("mps", "float32"), "float32", "mps/fp32 stays")) -results.append(expect(pick_runtime_dtype("mps", "fp32"), "fp32", "mps/fp32-alias stays")) - -print("\n=== pick_runtime_dtype: VOXCPM_MPS_DTYPE override ===") -os.environ["VOXCPM_MPS_DTYPE"] = "bfloat16" -results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "bfloat16", "override bf16 honored")) - -os.environ["VOXCPM_MPS_DTYPE"] = "FP16" -results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "fp16", "override is case-insensitive")) - -os.environ["VOXCPM_MPS_DTYPE"] = " float32 " -results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "float32", "override is whitespace-trimmed")) - -print("\n=== pick_runtime_dtype: 'half' is no longer a valid override ===") -os.environ["VOXCPM_MPS_DTYPE"] = "half" -results.append( - expect_raises( - lambda: pick_runtime_dtype("mps", "bfloat16"), - ValueError, - "override=half now rejected (was the bug)", +def run_checks(): + results = [] + + print("=== override set sanity ===") + results.append(expect("half" not in _VALID_DTYPE_OVERRIDES, True, "half removed from _VALID_DTYPE_OVERRIDES")) + results.append(expect("half" not in _LOW_PRECISION_DTYPES, True, "half removed from _LOW_PRECISION_DTYPES")) + + print("\n=== every accepted override parses through get_dtype ===") + for dt in sorted(_VALID_DTYPE_OVERRIDES): + try: + torch_dtype = get_dtype(dt) + print(f"[OK ] get_dtype({dt!r}) -> {torch_dtype}") + results.append(True) + except Exception as e: + print(f"[FAIL] get_dtype({dt!r}) raised: {e}") + results.append(False) + + print("\n=== pick_runtime_dtype: non-mps is a no-op ===") + results.append(expect(pick_runtime_dtype("cuda", "bfloat16"), "bfloat16", "cuda/bf16 untouched")) + results.append(expect(pick_runtime_dtype("cpu", "float16"), "float16", "cpu/fp16 untouched")) + results.append(expect(pick_runtime_dtype("cuda", "float32"), "float32", "cuda/fp32 untouched")) + + print("\n=== pick_runtime_dtype: mps forces fp32 for low-precision ===") + os.environ.pop("VOXCPM_MPS_DTYPE", None) + results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "float32", "mps/bf16 -> fp32")) + results.append(expect(pick_runtime_dtype("mps", "bf16"), "float32", "mps/bf16-alias -> fp32")) + results.append(expect(pick_runtime_dtype("mps", "float16"), "float32", "mps/fp16 -> fp32")) + results.append(expect(pick_runtime_dtype("mps", "fp16"), "float32", "mps/fp16-alias -> fp32")) + results.append(expect(pick_runtime_dtype("mps", "float32"), "float32", "mps/fp32 stays")) + results.append(expect(pick_runtime_dtype("mps", "fp32"), "fp32", "mps/fp32-alias stays")) + + print("\n=== pick_runtime_dtype: VOXCPM_MPS_DTYPE override ===") + os.environ["VOXCPM_MPS_DTYPE"] = "bfloat16" + results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "bfloat16", "override bf16 honored")) + + os.environ["VOXCPM_MPS_DTYPE"] = "FP16" + results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "fp16", "override is case-insensitive")) + + os.environ["VOXCPM_MPS_DTYPE"] = " float32 " + results.append(expect(pick_runtime_dtype("mps", "bfloat16"), "float32", "override is whitespace-trimmed")) + + print("\n=== pick_runtime_dtype: 'half' is no longer a valid override ===") + os.environ["VOXCPM_MPS_DTYPE"] = "half" + results.append( + expect_raises( + lambda: pick_runtime_dtype("mps", "bfloat16"), + ValueError, + "override=half now rejected (was the bug)", + ) ) -) - -os.environ["VOXCPM_MPS_DTYPE"] = "garbage" -results.append( - expect_raises( - lambda: pick_runtime_dtype("mps", "bfloat16"), - ValueError, - "override=garbage still rejected", + + os.environ["VOXCPM_MPS_DTYPE"] = "garbage" + results.append( + expect_raises( + lambda: pick_runtime_dtype("mps", "bfloat16"), + ValueError, + "override=garbage still rejected", + ) ) -) -os.environ.pop("VOXCPM_MPS_DTYPE", None) + os.environ.pop("VOXCPM_MPS_DTYPE", None) + + print("\n=== summary ===") + passed = sum(results) + total = len(results) + print(f"{passed}/{total} passed") + return passed == total + + +def test_pick_runtime_dtype_script_checks(): + assert run_checks() + -print("\n=== summary ===") -passed = sum(results) -total = len(results) -print(f"{passed}/{total} passed") -sys.exit(0 if passed == total else 1) +if __name__ == "__main__": + sys.exit(0 if run_checks() else 1) diff --git a/src/voxcpm/core.py b/src/voxcpm/core.py index 1a1d839..9b66330 100644 --- a/src/voxcpm/core.py +++ b/src/voxcpm/core.py @@ -4,7 +4,7 @@ import json import tempfile import numpy as np -from typing import Generator, Optional +from typing import Callable, Generator, Optional from huggingface_hub import snapshot_download from .model.voxcpm import VoxCPMModel, LoRAConfig from .model.voxcpm2 import VoxCPM2Model @@ -29,8 +29,8 @@ def __init__( (weights, configs, etc.). Typically the directory returned by a prior download step. zipenhancer_model_path: ModelScope acoustic noise suppression model - id or local path. If None, denoiser will not be initialized. - enable_denoiser: Whether to initialize the denoiser pipeline. + id or local path. If None, denoiser will not be available. + enable_denoiser: Whether denoising may be used when requested. optimize: Whether to optimize the model with torch.compile. True by default, but can be disabled for debugging. device: Runtime device. If set to ``None`` or ``"auto"``, VoxCPM will choose automatically (preferring CUDA, then MPS, then CPU). @@ -90,12 +90,7 @@ def __init__( self.text_normalizer = None self.denoiser = None - if enable_denoiser and zipenhancer_model_path is not None: - from .zipenhancer import ZipEnhancer - - self.denoiser = ZipEnhancer(zipenhancer_model_path) - else: - self.denoiser = None + self._denoiser_model_path = zipenhancer_model_path if enable_denoiser else None if optimize: print("Warm up VoxCPMModel...", file=sys.stderr) self.tts_model.generate( @@ -121,7 +116,7 @@ def from_pretrained( Args: hf_model_id: Explicit Hugging Face repository id (e.g. "org/repo") or local path. - load_denoiser: Whether to initialize the denoiser pipeline. + load_denoiser: Whether denoising may be used when requested. optimize: Whether to optimize the model with torch.compile. True by default, but can be disabled for debugging. zipenhancer_model_id: Denoiser model id or path for ModelScope acoustic noise suppression. @@ -180,6 +175,20 @@ def generate(self, *args, **kwargs) -> np.ndarray: def generate_streaming(self, *args, **kwargs) -> Generator[np.ndarray, None, None]: return self._generate(*args, streaming=True, **kwargs) + def _get_or_load_denoiser(self): + if self._denoiser_model_path is None: + return None + if self.denoiser is None: + try: + from .zipenhancer import ZipEnhancer + except ImportError as exc: + raise RuntimeError( + "ZipEnhancer denoising was requested, but its dependencies are not available. " + "Install the denoising dependencies or run with denoise disabled." + ) from exc + self.denoiser = ZipEnhancer(self._denoiser_model_path) + return self.denoiser + def _generate( self, text: str, @@ -197,6 +206,7 @@ def _generate( retry_badcase_ratio_threshold: float = 6.0, streaming: bool = False, seed: Optional[int] = None, + progress_callback: Optional[Callable[[int, int], None]] = None, ) -> Generator[np.ndarray, None, None]: """Synthesize speech for the given text and return a single waveform. @@ -251,16 +261,18 @@ def _generate( actual_prompt_path = prompt_wav_path actual_ref_path = reference_wav_path - if denoise and self.denoiser is not None: + denoiser = self._get_or_load_denoiser() if denoise else None + + if denoiser is not None: if prompt_wav_path is not None: with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as tmp: temp_files.append(tmp.name) - self.denoiser.enhance(prompt_wav_path, output_path=temp_files[-1]) + denoiser.enhance(prompt_wav_path, output_path=temp_files[-1]) actual_prompt_path = temp_files[-1] if reference_wav_path is not None: with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as tmp: temp_files.append(tmp.name) - self.denoiser.enhance(reference_wav_path, output_path=temp_files[-1]) + denoiser.enhance(reference_wav_path, output_path=temp_files[-1]) actual_ref_path = temp_files[-1] if actual_prompt_path is not None or actual_ref_path is not None: @@ -297,6 +309,7 @@ def _generate( retry_badcase_ratio_threshold=retry_badcase_ratio_threshold, streaming=streaming, seed=seed, + progress_callback=progress_callback, ) if streaming: diff --git a/src/voxcpm/model/voxcpm.py b/src/voxcpm/model/voxcpm.py index 20fc15b..948e34f 100644 --- a/src/voxcpm/model/voxcpm.py +++ b/src/voxcpm/model/voxcpm.py @@ -20,7 +20,7 @@ import os import sys -from typing import Tuple, Union, Generator, List, Optional +from typing import Callable, Tuple, Union, Generator, List, Optional import torch import torch.nn as nn @@ -371,6 +371,7 @@ def _generate( retry_badcase_ratio_threshold: float = 6.0, # setting acceptable ratio of audio length to text length (for badcase detection) streaming: bool = False, seed: Optional[int] = None, + progress_callback: Optional[Callable[[int, int], None]] = None, ) -> Generator[torch.Tensor, None, None]: if retry_badcase and streaming: warnings.warn("Retry on bad cases is not supported in streaming mode, setting retry_badcase=False.") @@ -474,6 +475,7 @@ def _generate( inference_timesteps=inference_timesteps, cfg_value=cfg_value, streaming=streaming, + progress_callback=progress_callback, ) if streaming: patch_len = self.patch_size * self.chunk_size @@ -616,6 +618,7 @@ def _generate_with_prompt_cache( streaming: bool = False, streaming_prefix_len: int = 3, seed: Optional[int] = None, + progress_callback: Optional[Callable[[int, int], None]] = None, ) -> Generator[Tuple[torch.Tensor, torch.Tensor, Union[torch.Tensor, List[torch.Tensor]]], None, None]: """ Generate audio using pre-built prompt cache. @@ -710,6 +713,7 @@ def _generate_with_prompt_cache( cfg_value=cfg_value, streaming=streaming, streaming_prefix_len=streaming_prefix_len, + progress_callback=progress_callback, ) if streaming: patch_len = self.patch_size * self.chunk_size @@ -763,6 +767,7 @@ def _inference( cfg_value: float = 2.0, streaming: bool = False, streaming_prefix_len: int = 3, + progress_callback: Optional[Callable[[int, int], None]] = None, ) -> Generator[Tuple[torch.Tensor, Union[torch.Tensor, List[torch.Tensor]]], None, None]: """Core inference method for audio generation. @@ -831,6 +836,8 @@ def _inference( residual_hidden = residual_enc_outputs[:, -1, :] for i in tqdm(range(max_len)): + if progress_callback is not None: + progress_callback(i, max_len) dit_hidden_1 = self.lm_to_dit_proj(lm_hidden) # [b, h_dit] dit_hidden_2 = self.res_to_dit_proj(residual_hidden) # [b, h_dit] dit_hidden = dit_hidden_1 + dit_hidden_2 # [b, h_dit] diff --git a/src/voxcpm/model/voxcpm2.py b/src/voxcpm/model/voxcpm2.py index 174dea3..8e60e74 100644 --- a/src/voxcpm/model/voxcpm2.py +++ b/src/voxcpm/model/voxcpm2.py @@ -20,7 +20,7 @@ import os import sys -from typing import Tuple, Union, Generator, List, Optional +from typing import Callable, Tuple, Union, Generator, List, Optional import torch import torch.nn as nn @@ -482,6 +482,7 @@ def _generate( streaming: bool = False, streaming_prefix_len: int = 4, seed: Optional[int] = None, + progress_callback: Optional[Callable[[int, int], None]] = None, ) -> Generator[torch.Tensor, None, None]: if retry_badcase and streaming: warnings.warn("Retry on bad cases is not supported in streaming mode, setting retry_badcase=False.") @@ -656,6 +657,7 @@ def _generate( cfg_value=cfg_value, streaming=streaming, streaming_prefix_len=streaming_prefix_len, + progress_callback=progress_callback, ) if streaming: with self.audio_vae.streaming_decode() as vae_dec: @@ -808,6 +810,7 @@ def _generate_with_prompt_cache( streaming: bool = False, streaming_prefix_len: int = 4, seed: Optional[int] = None, + progress_callback: Optional[Callable[[int, int], None]] = None, ) -> Generator[Tuple[torch.Tensor, torch.Tensor, Union[torch.Tensor, List[torch.Tensor]]], None, None]: """ Generate audio using pre-built prompt cache. @@ -952,6 +955,7 @@ def _generate_with_prompt_cache( cfg_value=cfg_value, streaming=streaming, streaming_prefix_len=streaming_prefix_len, + progress_callback=progress_callback, ) if streaming: with self.audio_vae.streaming_decode() as vae_dec: @@ -1007,6 +1011,7 @@ def _inference( cfg_value: float = 2.0, streaming: bool = False, streaming_prefix_len: int = 4, + progress_callback: Optional[Callable[[int, int], None]] = None, ) -> Generator[Tuple[torch.Tensor, Union[torch.Tensor, List[torch.Tensor]]], None, None]: """Core inference method for audio generation. @@ -1081,6 +1086,8 @@ def _inference( residual_hidden = residual_enc_outputs[:, -1, :] for i in tqdm(range(max_len)): + if progress_callback is not None: + progress_callback(i, max_len) dit_hidden_1 = self.lm_to_dit_proj(lm_hidden) # [b, h_dit] dit_hidden_2 = self.res_to_dit_proj(residual_hidden) # [b, h_dit] dit_hidden = torch.cat((dit_hidden_1, dit_hidden_2), dim=-1) diff --git a/tests/test_app.py b/tests/test_app.py new file mode 100644 index 0000000..5fb1850 --- /dev/null +++ b/tests/test_app.py @@ -0,0 +1,295 @@ +from __future__ import annotations + +import numpy as np +import pytest +import soundfile as sf +import threading +import time + +import app + + +class PathLikeObject: + path = "object.wav" + + +def test_coerce_audio_filepath_accepts_gradio_file_shapes(): + assert app._coerce_audio_filepath(None) is None + assert app._coerce_audio_filepath("") is None + assert app._coerce_audio_filepath("plain.wav") == "plain.wav" + assert app._coerce_audio_filepath({"path": "dict.wav"}) == "dict.wav" + assert app._coerce_audio_filepath(PathLikeObject()) == "object.wav" + + +def test_extract_asr_text_removes_sensevoice_tags(): + result = [{"text": "<|zh|><|NEUTRAL|><|Speech|><|withitn|>你好,世界"}] + + assert app._extract_asr_text(result) == "你好,世界" + + +def test_extract_parakeet_asr_text_accepts_batch_decode_output(): + assert app._extract_parakeet_asr_text([" hej ", "", " världen"]) == "hej världen" + + +def test_normalize_asr_backend_rejects_unknown_backend(): + with pytest.raises(ValueError, match="Unknown ASR backend"): + app._normalize_asr_backend("whisper") + + +def test_prepare_asr_audio_keeps_16khz_mono_wav(tmp_path): + wav_path = tmp_path / "mono.wav" + sf.write(wav_path, np.zeros(160, dtype=np.float32), 16000) + + prepared_path, temp_path = app._prepare_asr_audio(str(wav_path)) + + assert prepared_path == str(wav_path) + assert temp_path is None + + +def test_prepare_asr_audio_converts_to_16khz_mono_wav(tmp_path): + wav_path = tmp_path / "stereo_8k.wav" + audio = np.zeros((80, 2), dtype=np.float32) + sf.write(wav_path, audio, 8000) + + prepared_path, temp_path = app._prepare_asr_audio(str(wav_path)) + + try: + info = sf.info(prepared_path) + assert temp_path == prepared_path + assert info.samplerate == 16000 + assert info.channels == 1 + finally: + if temp_path: + app.os.unlink(temp_path) + + +def test_resolve_generation_inputs_auto_transcribes_blank_ultimate_prompt(): + class FakeDemo: + calls = [] + + def prompt_wav_recognition(self, path, progress_callback=None): + self.calls.append(path) + return " auto transcript " + + demo = FakeDemo() + + audio_path, prompt_text, control = app._resolve_generation_inputs( + demo, + {"path": "ref.wav"}, + True, + "", + "warm voice", + ) + + assert audio_path == "ref.wav" + assert prompt_text == "auto transcript" + assert control == "" + assert demo.calls == ["ref.wav"] + + +def test_resolve_generation_inputs_requires_audio_for_ultimate_mode(): + with pytest.raises(app.gr.Error, match="Upload reference audio"): + app._resolve_generation_inputs(object(), None, True, "", "") + + +def test_auto_asr_backend_prefers_local_parakeet_on_cuda(): + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.asr_backend = "auto" + demo.device = "cuda" + demo.parakeet_model_id = "models/nvidia__parakeet-tdt-0.6b-v3" + + assert demo._should_use_parakeet_asr() is True + assert demo._resolved_asr_backend_name() == "parakeet" + + +def test_auto_asr_backend_uses_sensevoice_without_local_parakeet(): + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.asr_backend = "auto" + demo.device = "cuda" + demo.parakeet_model_id = None + + assert demo._should_use_parakeet_asr() is False + assert demo._resolved_asr_backend_name() == "sensevoice" + + +def test_sensevoice_asr_backend_disables_local_parakeet(): + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.asr_backend = "sensevoice" + demo.device = "cuda" + demo.parakeet_model_id = "models/nvidia__parakeet-tdt-0.6b-v3" + + assert demo._should_use_parakeet_asr() is False + + +def test_parakeet_asr_backend_uses_local_parakeet(): + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.asr_backend = "parakeet" + demo.device = "cuda" + demo.parakeet_model_id = "models/nvidia__parakeet-tdt-0.6b-v3" + + assert demo._should_use_parakeet_asr() is True + assert demo._resolved_asr_backend_name() == "parakeet" + + +def test_preload_models_loads_tts_denoiser_parakeet_and_sensevoice_fallback_on_cuda_auto(): + class FakeCoreModel: + def _get_or_load_denoiser(self): + calls.append("denoiser") + + calls = [] + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.asr_backend = "auto" + demo.device = "cuda" + demo.parakeet_model_id = "models/nvidia__parakeet-tdt-0.6b-v3" + demo.get_or_load_voxcpm = lambda: calls.append("tts") or FakeCoreModel() + demo.get_or_load_parakeet_asr_model = lambda: calls.append("parakeet") + demo.get_or_load_asr_model = lambda: calls.append("sensevoice") + + demo.preload_models() + + assert calls == ["tts", "denoiser", "parakeet", "sensevoice"] + + +def test_get_or_load_asr_model_serializes_concurrent_loads(monkeypatch): + load_count = 0 + load_count_lock = threading.Lock() + + class FakeAutoModel: + def __init__(self, **kwargs): + nonlocal load_count + with load_count_lock: + load_count += 1 + time.sleep(0.05) + + monkeypatch.setattr(app, "AutoModel", FakeAutoModel) + + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.asr_model = None + demo.asr_model_id = "fake/asr" + demo.asr_device = "cpu" + + results = [] + errors = [] + + def load_model(): + try: + results.append(demo.get_or_load_asr_model()) + except Exception as exc: + errors.append(exc) + + threads = [threading.Thread(target=load_model) for _ in range(8)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert errors == [] + assert load_count == 1 + assert len(results) == len(threads) + assert all(result is results[0] for result in results) + + +def test_run_demo_preloads_models_before_launching_web_ui(monkeypatch): + events = [] + launch_kwargs = {} + + class FakeDemo: + def __init__(self, model_id, device, asr_backend): + events.append(("demo", model_id, device, asr_backend)) + + def preload_models(self, **kwargs): + events.append(("preload", kwargs)) + + class FakeQueuedInterface: + def launch(self, **kwargs): + events.append("launch") + launch_kwargs.update(kwargs) + + class FakeInterface: + def queue(self, **kwargs): + events.append(("queue", kwargs)) + return FakeQueuedInterface() + + monkeypatch.setattr(app, "VoxCPMDemo", FakeDemo) + monkeypatch.setattr(app, "create_demo_interface", lambda demo: FakeInterface()) + + app.run_demo() + + assert events == [ + ("demo", "openbmb/VoxCPM2", "auto", "auto"), + ("preload", {"preload_asr": True, "preload_tts": True, "preload_denoiser": True}), + ("queue", {"max_size": 10, "default_concurrency_limit": 1}), + "launch", + ] + assert launch_kwargs["server_name"] == "127.0.0.1" + assert launch_kwargs["server_port"] == 8808 + assert launch_kwargs["inbrowser"] is True + + +def test_prompt_wav_recognition_reports_progress_and_uses_parakeet(monkeypatch): + progress_events = [] + calls = [] + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.asr_backend = "parakeet" + demo.device = "cuda" + demo.parakeet_model_id = "models/nvidia__parakeet-tdt-0.6b-v3" + + monkeypatch.setattr(app, "_prepare_asr_audio", lambda path: ("prepared.wav", None)) + + def fake_parakeet(path, progress_callback=None): + calls.append(("parakeet", path)) + app._emit_progress(progress_callback, 0.55, "Transcribing reference audio with Parakeet, 55%") + return "transcript" + + demo._recognize_with_parakeet = fake_parakeet + demo._recognize_with_sensevoice = lambda path, progress_callback=None: calls.append(("sensevoice", path)) or "" + + text = demo.prompt_wav_recognition( + "ref.wav", progress_callback=lambda value, label: progress_events.append((value, label)) + ) + + assert text == "transcript" + assert calls == [("parakeet", "prepared.wav")] + assert progress_events[0][0] == 0.05 + assert "Parakeet" in progress_events[-1][1] + + +def test_generate_tts_audio_normalizes_gradio_filedata_path(): + class FakeTTS: + sample_rate = 24000 + last_successful_seed = 456 + + class FakeModel: + tts_model = FakeTTS() + + def __init__(self): + self.kwargs = None + + def generate(self, **kwargs): + self.kwargs = kwargs + return np.array([0.0], dtype=np.float32) + + fake_model = FakeModel() + demo = app.VoxCPMDemo.__new__(app.VoxCPMDemo) + demo.get_or_load_voxcpm = lambda: fake_model + + progress_callback = lambda step, total: None + + sr, wav, seed = app.VoxCPMDemo.generate_tts_audio( + demo, + text_input="Hello", + reference_wav_path_input={"path": "ref.wav"}, + prompt_text="reference transcript", + do_normalize=False, + denoise=False, + seed=123, + progress_callback=progress_callback, + ) + + assert sr == 24000 + assert seed == 456 + np.testing.assert_array_equal(wav, np.array([0.0], dtype=np.float32)) + assert fake_model.kwargs["reference_wav_path"] == "ref.wav" + assert fake_model.kwargs["prompt_wav_path"] == "ref.wav" + assert fake_model.kwargs["prompt_text"] == "reference transcript" + assert fake_model.kwargs["progress_callback"] is progress_callback diff --git a/tests/test_core_denoiser.py b/tests/test_core_denoiser.py new file mode 100644 index 0000000..fc0d55e --- /dev/null +++ b/tests/test_core_denoiser.py @@ -0,0 +1,172 @@ +from __future__ import annotations + +import importlib.util +import json +import sys +import types +from pathlib import Path + +import numpy as np + +ROOT = Path(__file__).resolve().parents[1] +CORE_PATH = ROOT / "src" / "voxcpm" / "core.py" + + +class FakeTensor: + def squeeze(self, dim): + return self + + def cpu(self): + return self + + def numpy(self): + return np.array([0.0], dtype=np.float32) + + +class DummyVoxCPM2Model: + sample_rate = 16000 + + @classmethod + def from_local(cls, *args, **kwargs): + return cls() + + def __init__(self): + self.prompt_cache_calls = [] + self.generate_calls = [] + + def generate(self, **kwargs): + self.generate_calls.append(kwargs) + + def build_prompt_cache(self, **kwargs): + self.prompt_cache_calls.append(kwargs) + return {"prompt_cache": True} + + def _generate_with_prompt_cache(self, **kwargs): + yield FakeTensor(), None, None + + +class DummyVoxCPMModel(DummyVoxCPM2Model): + pass + + +class DummyLoRAConfig: + def __init__(self, **kwargs): + self.kwargs = kwargs + + +def _next_and_close(generator): + try: + return next(generator) + finally: + close = getattr(generator, "close", None) + if close is not None: + close() + + +def load_core_with_stubs(monkeypatch): + for module_name in [ + "voxcpm", + "voxcpm.core", + "voxcpm.model", + "voxcpm.model.utils", + "voxcpm.model.voxcpm", + "voxcpm.model.voxcpm2", + "voxcpm.zipenhancer", + "huggingface_hub", + ]: + monkeypatch.delitem(sys.modules, module_name, raising=False) + + pkg = types.ModuleType("voxcpm") + pkg.__path__ = [str(ROOT / "src" / "voxcpm")] + monkeypatch.setitem(sys.modules, "voxcpm", pkg) + + model_pkg = types.ModuleType("voxcpm.model") + model_pkg.__path__ = [str(ROOT / "src" / "voxcpm" / "model")] + monkeypatch.setitem(sys.modules, "voxcpm.model", model_pkg) + + utils_stub = types.ModuleType("voxcpm.model.utils") + utils_stub.next_and_close = _next_and_close + monkeypatch.setitem(sys.modules, "voxcpm.model.utils", utils_stub) + + v1_stub = types.ModuleType("voxcpm.model.voxcpm") + v1_stub.VoxCPMModel = DummyVoxCPMModel + v1_stub.LoRAConfig = DummyLoRAConfig + monkeypatch.setitem(sys.modules, "voxcpm.model.voxcpm", v1_stub) + + v2_stub = types.ModuleType("voxcpm.model.voxcpm2") + v2_stub.VoxCPM2Model = DummyVoxCPM2Model + monkeypatch.setitem(sys.modules, "voxcpm.model.voxcpm2", v2_stub) + + hub_stub = types.ModuleType("huggingface_hub") + hub_stub.snapshot_download = lambda **kwargs: kwargs["repo_id"] + monkeypatch.setitem(sys.modules, "huggingface_hub", hub_stub) + + spec = importlib.util.spec_from_file_location("voxcpm.core", CORE_PATH) + core = importlib.util.module_from_spec(spec) + monkeypatch.setitem(sys.modules, "voxcpm.core", core) + assert spec.loader is not None + spec.loader.exec_module(core) + return core + + +def make_model_dir(tmp_path): + model_dir = tmp_path / "model" + model_dir.mkdir() + (model_dir / "config.json").write_text(json.dumps({"architecture": "voxcpm2"}), encoding="utf-8") + return model_dir + + +def test_denoiser_is_not_loaded_during_model_init(monkeypatch, tmp_path): + core = load_core_with_stubs(monkeypatch) + model_dir = make_model_dir(tmp_path) + + model = core.VoxCPM( + str(model_dir), + zipenhancer_model_path="zip-model", + enable_denoiser=True, + optimize=False, + ) + + assert model.denoiser is None + assert model._denoiser_model_path == "zip-model" + assert "voxcpm.zipenhancer" not in sys.modules + + +def test_denoiser_is_loaded_when_generation_requests_denoise(monkeypatch, tmp_path): + core = load_core_with_stubs(monkeypatch) + model_dir = make_model_dir(tmp_path) + ref_audio = tmp_path / "reference.wav" + ref_audio.write_bytes(b"RIFF") + + init_calls = [] + enhance_calls = [] + + zipenhancer_stub = types.ModuleType("voxcpm.zipenhancer") + + class FakeZipEnhancer: + def __init__(self, model_path): + init_calls.append(model_path) + + def enhance(self, input_path, output_path=None, normalize_loudness=True): + enhance_calls.append((input_path, output_path, normalize_loudness)) + Path(output_path).write_bytes(b"RIFF") + return output_path + + zipenhancer_stub.ZipEnhancer = FakeZipEnhancer + monkeypatch.setitem(sys.modules, "voxcpm.zipenhancer", zipenhancer_stub) + + model = core.VoxCPM( + str(model_dir), + zipenhancer_model_path="zip-model", + enable_denoiser=True, + optimize=False, + ) + + wav = model.generate("hello", reference_wav_path=str(ref_audio), denoise=True) + + assert init_calls == ["zip-model"] + assert len(enhance_calls) == 1 + assert enhance_calls[0][0] == str(ref_audio) + assert enhance_calls[0][1] != str(ref_audio) + assert model.tts_model.prompt_cache_calls[0]["reference_wav_path"] == enhance_calls[0][1] + np.testing.assert_array_equal(wav, np.array([0.0], dtype=np.float32))