diff --git a/app.py b/app.py index 95eac94..b3fa16b 100644 --- a/app.py +++ b/app.py @@ -2,6 +2,7 @@ import os import re import sys import logging +import random import numpy as np import gradio as gr from typing import Optional, Tuple @@ -115,6 +116,10 @@ _I18N_TRANSLATIONS = { "cfg_info": "Higher → closer to the prompt / reference; lower → more creative variation", "dit_steps_label": "LocDiT flow-matching steps", "dit_steps_info": "LocDiT flow-matching steps — more steps → maybe better audio quality, but slower", + "seed_label": "Seed", + "seed_info": "Seed used for reproducible generation. Updated with the actual successful seed after generation.", + "random_seed_label": "Random Seed", + "random_seed_info": "Generate a new seed before each inference run.", "usage_instructions": _USAGE_INSTRUCTIONS_EN, "examples_footer": _EXAMPLES_FOOTER_EN, }, @@ -279,6 +284,7 @@ class VoxCPMDemo: do_normalize: bool, denoise: bool, inference_timesteps: int = 10, + seed: Optional[int] = None, ) -> dict: generate_kwargs = dict( text=final_text, @@ -287,6 +293,7 @@ class VoxCPMDemo: inference_timesteps=inference_timesteps, normalize=do_normalize, denoise=denoise, + seed=seed, ) if prompt_text_clean and audio_path: generate_kwargs["prompt_wav_path"] = audio_path @@ -303,7 +310,8 @@ class VoxCPMDemo: do_normalize: bool = True, denoise: bool = True, inference_timesteps: int = 10, - ) -> Tuple[int, np.ndarray]: + seed: Optional[int] = None, + ) -> Tuple[int, np.ndarray, Optional[int]]: current_model = self.get_or_load_voxcpm() text = (text_input or "").strip() @@ -335,9 +343,11 @@ class VoxCPMDemo: do_normalize=do_normalize, denoise=denoise, inference_timesteps=inference_timesteps, + seed=seed, ) wav = current_model.generate(**generate_kwargs) - return (current_model.tts_model.sample_rate, wav) + last_successful_seed = getattr(current_model.tts_model, "last_successful_seed", seed) + return (current_model.tts_model.sample_rate, wav, last_successful_seed) # ---------- UI ---------- @@ -345,6 +355,19 @@ class VoxCPMDemo: def create_demo_interface(demo: VoxCPMDemo): gr.set_static_paths(paths=[Path.cwd().absolute() / "assets"]) + def _coerce_seed(seed_value) -> Optional[int]: + if seed_value is None or seed_value == "": + return None + return int(seed_value) + + def _prepare_seed(use_random_seed: bool, seed_value): + if use_random_seed: + return random.randint(0, 2**32 - 1) + return _coerce_seed(seed_value) + + def _on_random_seed_toggle(checked): + return gr.update(interactive=not checked) + def _generate( text: str, control_instruction: str, @@ -355,10 +378,12 @@ def create_demo_interface(demo: VoxCPMDemo): do_normalize: bool, denoise: bool, dit_steps: int, + seed_value, ): actual_prompt_text = prompt_text_value.strip() if use_prompt_text else "" actual_control = "" if use_prompt_text else control_instruction - sr, wav_np = demo.generate_tts_audio( + seed = _coerce_seed(seed_value) + sr, wav_np, last_successful_seed = demo.generate_tts_audio( text_input=text, control_instruction=actual_control, reference_wav_path_input=ref_wav, @@ -367,8 +392,9 @@ def create_demo_interface(demo: VoxCPMDemo): do_normalize=do_normalize, denoise=denoise, inference_timesteps=int(dit_steps), + seed=seed, ) - return (sr, wav_np) + return (sr, wav_np), last_successful_seed def _on_toggle_instant(checked): """Instant UI toggle — no ASR, no blocking.""" @@ -465,6 +491,20 @@ def create_demo_interface(demo: VoxCPMDemo): label=I18N("dit_steps_label"), info=I18N("dit_steps_info"), ) + with gr.Row(): + seed_value = gr.Number( + value=random.randint(0, 2**32 - 1), + precision=0, + label=I18N("seed_label"), + info=I18N("seed_info"), + interactive=False, + ) + random_seed = gr.Checkbox( + value=True, + label=I18N("random_seed_label"), + elem_classes=["switch-toggle"], + info=I18N("random_seed_info"), + ) run_btn = gr.Button(I18N("generate_btn"), variant="primary", size="lg") @@ -482,7 +522,18 @@ def create_demo_interface(demo: VoxCPMDemo): outputs=[prompt_text], ) + random_seed.change( + fn=_on_random_seed_toggle, + inputs=[random_seed], + outputs=[seed_value], + ) + run_btn.click( + fn=_prepare_seed, + inputs=[random_seed, seed_value], + outputs=[seed_value], + show_progress=False, + ).then( fn=_generate, inputs=[ text, @@ -494,8 +545,9 @@ def create_demo_interface(demo: VoxCPMDemo): DoNormalizeText, DoDenoisePromptAudio, dit_steps, + seed_value, ], - outputs=[audio_output], + outputs=[audio_output, seed_value], show_progress=True, api_name="generate", ) diff --git a/src/voxcpm/model/utils.py b/src/voxcpm/model/utils.py index 940fc74..203cbe6 100644 --- a/src/voxcpm/model/utils.py +++ b/src/voxcpm/model/utils.py @@ -21,6 +21,19 @@ def next_and_close(gen): gen.close() +def materialize_generation_seed(seed: Optional[int]) -> int: + """Return a concrete seed for a generation request.""" + if seed is not None: + return int(seed) + return int(torch.seed() & 0xFFFFFFFF) + + +def apply_generation_seed(seed: int) -> None: + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + def mask_multichar_chinese_tokens(tokenizer: PreTrainedTokenizer): """Create a tokenizer wrapper that converts multi-character Chinese tokens to single characters. diff --git a/src/voxcpm/model/voxcpm.py b/src/voxcpm/model/voxcpm.py index 0e28fdf..20fc15b 100644 --- a/src/voxcpm/model/voxcpm.py +++ b/src/voxcpm/model/voxcpm.py @@ -45,7 +45,9 @@ from ..modules.locdit import CfmConfig, UnifiedCFM, VoxCPMLocDiT from ..modules.locenc import VoxCPMLocEnc from ..modules.minicpm4 import MiniCPM4Config, MiniCPMModel from .utils import ( + apply_generation_seed, get_dtype, + materialize_generation_seed, mask_multichar_chinese_tokens, next_and_close, pick_runtime_dtype, @@ -140,6 +142,7 @@ class VoxCPMModel(nn.Module): self.text_tokenizer = mask_multichar_chinese_tokens(tokenizer) self.audio_start_token = 101 self.audio_end_token = 102 + self.last_successful_seed = None # Residual Acoustic LM residual_lm_config = config.lm_config.model_copy(deep=True) @@ -453,13 +456,12 @@ class VoxCPMModel(nn.Module): target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 - current_seed = seed + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: - if current_seed is not None: - torch.manual_seed(current_seed) - if torch.cuda.is_available(): - torch.cuda.manual_seed_all(current_seed) - + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -478,7 +480,7 @@ class VoxCPMModel(nn.Module): for latent_pred, _ in inference_result: decode_audio = self.audio_vae.decode(latent_pred.to(torch.float32)) decode_audio = decode_audio[..., -patch_len:].squeeze(1).cpu() - self.last_successful_seed = current_seed + self.last_successful_seed = last_attempt_seed yield decode_audio break else: @@ -490,16 +492,15 @@ class VoxCPMModel(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 - if current_seed is not None: - current_seed += 1 + current_seed += 1 continue else: - self.last_successful_seed = current_seed break else: break if not streaming: + self.last_successful_seed = last_attempt_seed decode_audio = self.audio_vae.decode(latent_pred.to(torch.float32)).squeeze(1).cpu() yield decode_audio @@ -690,13 +691,12 @@ class VoxCPMModel(nn.Module): # run inference target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 - current_seed = seed + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: - if current_seed is not None: - torch.manual_seed(current_seed) - if torch.cuda.is_available(): - torch.cuda.manual_seed_all(current_seed) - + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -716,7 +716,7 @@ class VoxCPMModel(nn.Module): for latent_pred, pred_audio_feat in inference_result: decode_audio = self.audio_vae.decode(latent_pred.to(torch.float32)) decode_audio = decode_audio[..., -patch_len:].squeeze(1).cpu() - self.last_successful_seed = current_seed + self.last_successful_seed = last_attempt_seed yield (decode_audio, target_text_token, pred_audio_feat) break else: @@ -728,15 +728,14 @@ class VoxCPMModel(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 - if current_seed is not None: - current_seed += 1 + current_seed += 1 continue else: - self.last_successful_seed = current_seed break else: break if not streaming: + self.last_successful_seed = last_attempt_seed decode_audio = self.audio_vae.decode(latent_pred.to(torch.float32)) patch_len = self.patch_size * self.chunk_size if audio_mask.sum().item() > 0: diff --git a/src/voxcpm/model/voxcpm2.py b/src/voxcpm/model/voxcpm2.py index 2fc2986..fd3f290 100644 --- a/src/voxcpm/model/voxcpm2.py +++ b/src/voxcpm/model/voxcpm2.py @@ -46,7 +46,9 @@ from ..modules.locdit import CfmConfig, UnifiedCFM, VoxCPMLocDiTV2 from ..modules.locenc import VoxCPMLocEnc from ..modules.minicpm4 import MiniCPM4Config, MiniCPMModel from .utils import ( + apply_generation_seed, get_dtype, + materialize_generation_seed, mask_multichar_chinese_tokens, next_and_close, pick_runtime_dtype, @@ -184,6 +186,7 @@ class VoxCPM2Model(nn.Module): self.audio_end_token = 102 self.ref_audio_start_token = 103 self.ref_audio_end_token = 104 + self.last_successful_seed = None # Residual Acoustic LM residual_lm_config = config.lm_config.model_copy(deep=True) @@ -634,13 +637,12 @@ class VoxCPM2Model(nn.Module): target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 - current_seed = seed + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: - if current_seed is not None: - torch.manual_seed(current_seed) - if torch.cuda.is_available(): - torch.cuda.manual_seed_all(current_seed) - + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -658,7 +660,7 @@ class VoxCPM2Model(nn.Module): for latent_pred, _, _ctx in inference_result: decode_audio = vae_dec.decode_chunk(latent_pred.to(torch.float32)) decode_audio = decode_audio.squeeze(1).cpu() - self.last_successful_seed = current_seed + self.last_successful_seed = last_attempt_seed yield decode_audio break else: @@ -670,16 +672,15 @@ class VoxCPM2Model(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 - if current_seed is not None: - current_seed += 1 + current_seed += 1 continue else: - self.last_successful_seed = current_seed break else: break if not streaming: + self.last_successful_seed = last_attempt_seed decode_audio = self.audio_vae.decode(latent_pred.to(torch.float32)) decode_patch_len = self.patch_size * self._decode_chunk_size if context_len > 0: @@ -932,13 +933,12 @@ class VoxCPM2Model(nn.Module): # run inference target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 - current_seed = seed + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: - if current_seed is not None: - torch.manual_seed(current_seed) - if torch.cuda.is_available(): - torch.cuda.manual_seed_all(current_seed) - + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -956,7 +956,7 @@ class VoxCPM2Model(nn.Module): for latent_pred, pred_audio_feat, _ctx in inference_result: decode_audio = vae_dec.decode_chunk(latent_pred.to(torch.float32)) decode_audio = decode_audio.squeeze(1).cpu() - self.last_successful_seed = current_seed + self.last_successful_seed = last_attempt_seed yield (decode_audio, target_text_token, pred_audio_feat) break else: @@ -968,15 +968,14 @@ class VoxCPM2Model(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 - if current_seed is not None: - current_seed += 1 + current_seed += 1 continue else: - self.last_successful_seed = current_seed break else: break if not streaming: + self.last_successful_seed = last_attempt_seed decode_audio = self.audio_vae.decode(latent_pred.to(torch.float32)) decode_patch_len = self.patch_size * self._decode_chunk_size if context_len > 0: diff --git a/tests/test_model_utils.py b/tests/test_model_utils.py index bb69ffc..cc33845 100644 --- a/tests/test_model_utils.py +++ b/tests/test_model_utils.py @@ -47,3 +47,25 @@ def test_resolve_runtime_device_rejects_unavailable_explicit_cuda(monkeypatch): with pytest.raises(ValueError, match="CUDA is not available"): utils.resolve_runtime_device("cuda:0", "cuda") + + +def test_materialize_generation_seed_preserves_explicit_seed(): + assert utils.materialize_generation_seed(42) == 42 + + +def test_materialize_generation_seed_creates_concrete_seed_for_none(monkeypatch): + monkeypatch.setattr(utils.torch, "seed", lambda: 0x123456789) + + assert utils.materialize_generation_seed(None) == 0x23456789 + + +def test_apply_generation_seed_sets_cpu_and_cuda_rng(monkeypatch): + calls = [] + + monkeypatch.setattr(utils.torch, "manual_seed", lambda seed: calls.append(("cpu", seed))) + monkeypatch.setattr(utils.torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(utils.torch.cuda, "manual_seed_all", lambda seed: calls.append(("cuda", seed))) + + utils.apply_generation_seed(123) + + assert calls == [("cpu", 123), ("cuda", 123)]