diff --git a/README.md b/README.md index 37f58b1..6e03170 100644 --- a/README.md +++ b/README.md @@ -108,6 +108,7 @@ wav = model.generate( text="VoxCPM2 is the current recommended release for realistic multilingual speech synthesis.", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("demo.wav", wav, model.tts_model.sample_rate) print("saved: demo.wav") @@ -131,6 +132,7 @@ wav = model.generate( text="VoxCPM2 is the current recommended release for realistic multilingual speech synthesis.", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("demo.wav", wav, model.tts_model.sample_rate) ``` @@ -144,6 +146,7 @@ wav = model.generate( text="(A young woman, gentle and sweet voice)Hello, welcome to VoxCPM2!", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("voice_design.wav", wav, model.tts_model.sample_rate) ``` @@ -164,6 +167,7 @@ wav = model.generate( reference_wav_path="path/to/voice.wav", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("controllable_clone.wav", wav, model.tts_model.sample_rate) ``` @@ -210,6 +214,7 @@ voxcpm design \ voxcpm design \ --text "VoxCPM2 brings studio-quality multilingual speech synthesis." \ --control "Young female voice, warm and gentle, slightly smiling" \ + --seed 42 \ --output out.wav # Voice cloning (reference audio) diff --git a/README_zh.md b/README_zh.md index 835e143..ea86dc7 100644 --- a/README_zh.md +++ b/README_zh.md @@ -10,7 +10,7 @@ Documentation Hugging Face ModelScope - + DemoPage

@@ -110,6 +110,7 @@ wav = model.generate( text="VoxCPM2 是目前推荐使用的多语言语音合成版本。", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("demo.wav", wav, model.tts_model.sample_rate) print("已保存: demo.wav") @@ -133,6 +134,7 @@ wav = model.generate( text="VoxCPM2 是目前推荐使用的多语言语音合成版本。", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("demo.wav", wav, model.tts_model.sample_rate) ``` @@ -146,6 +148,7 @@ wav = model.generate( text="(年轻女性,声音温柔甜美)你好,欢迎使用VoxCPM2!", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("voice_design.wav", wav, model.tts_model.sample_rate) ``` @@ -166,6 +169,7 @@ wav = model.generate( reference_wav_path="path/to/voice.wav", cfg_value=2.0, inference_timesteps=10, + seed=42, ) sf.write("controllable_clone.wav", wav, model.tts_model.sample_rate) ``` @@ -212,6 +216,7 @@ voxcpm design \ voxcpm design \ --text "VoxCPM2带来全新语音合成体验。" \ --control "年轻女声,温暖温柔,略带微笑" \ + --seed 42 \ --output out.wav # 声音克隆(参考音频) diff --git a/app.py b/app.py index 95eac94..99a3c5c 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 @@ -44,8 +45,8 @@ _EXAMPLES_FOOTER_EN = ( "**Example 1 — Gentle & Melancholic Girl** \n" '`Control Instruction`: *"A young girl with a soft, sweet voice. ' 'Speaks slowly with a melancholic, slightly tsundere tone."* \n' - '`Target Text`: *"I never asked you to stay… It\'s not like I care or anything. ' - 'But… why does it still hurt so much now that you\'re gone?"* \n\n' + "`Target Text`: *\"I never asked you to stay… It's not like I care or anything. " + "But… why does it still hurt so much now that you're gone?\"* \n\n" "**Example 2 — Laid-Back Surfer Dude** \n" '`Control Instruction`: *"Relaxed young male voice, slightly nasal, ' 'lazy drawl, very casual and chill."* \n' @@ -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, }, @@ -142,7 +147,7 @@ _I18N_TRANSLATIONS = { "examples_footer": _EXAMPLES_FOOTER_ZH, }, "zh-Hans": None, # alias, filled below - "zh": None, # alias, filled below + "zh": None, # alias, filled below } _I18N_TRANSLATIONS["zh-Hans"] = _I18N_TRANSLATIONS["zh-CN"] _I18N_TRANSLATIONS["zh"] = _I18N_TRANSLATIONS["zh-CN"] @@ -155,8 +160,7 @@ for _d in _I18N_TRANSLATIONS.values(): I18N = gr.I18n(**_I18N_TRANSLATIONS) DEFAULT_TARGET_TEXT = ( - "VoxCPM2 is a creative multilingual TTS model from ModelBest, " - "designed to generate highly realistic speech." + "VoxCPM2 is a creative multilingual TTS model from ModelBest, " "designed to generate highly realistic speech." ) _CUSTOM_CSS = """ @@ -219,6 +223,7 @@ _APP_THEME = gr.themes.Soft( # ---------- Model ---------- + class VoxCPMDemo: def __init__(self, model_id: str = "openbmb/VoxCPM2", device: str = "auto") -> None: self.device = resolve_runtime_device(device, "cuda") @@ -247,9 +252,7 @@ class VoxCPMDemo: 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}" - ) + 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, @@ -279,6 +282,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 +291,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 +308,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,16 +341,32 @@ 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 ---------- + 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 +377,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 +391,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 +490,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 +521,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,14 +544,16 @@ 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", ) return interface + def run_demo( server_name: str = "0.0.0.0", server_port: int = 8808, @@ -523,9 +575,12 @@ def run_demo( if __name__ == "__main__": import argparse + parser = argparse.ArgumentParser() parser.add_argument( - "--model-id", type=str, default="openbmb/VoxCPM2", + "--model-id", + type=str, + default="openbmb/VoxCPM2", help="Local path or HuggingFace repo ID (default: openbmb/VoxCPM2)", ) parser.add_argument("--port", type=int, default=8808, help="Server port") diff --git a/app_old.py b/app_old.py index d46c2e1..c6ddbeb 100644 --- a/app_old.py +++ b/app_old.py @@ -6,6 +6,7 @@ import gradio as gr from typing import Optional, Tuple from funasr import AutoModel from pathlib import Path + os.environ["TOKENIZERS_PARALLELISM"] = "false" if os.environ.get("HF_REPO_ID", "").strip() == "": os.environ["HF_REPO_ID"] = "openbmb/VoxCPM1.5" @@ -23,7 +24,7 @@ class VoxCPMDemo: self.asr_model: Optional[AutoModel] = AutoModel( model=self.asr_model_id, disable_update=True, - log_level='DEBUG', + log_level="DEBUG", device="cuda:0" if self.device == "cuda" else "cpu", ) @@ -48,6 +49,7 @@ class VoxCPMDemo: if not os.path.isdir(target_dir): try: from huggingface_hub import snapshot_download # type: ignore + os.makedirs(target_dir, exist_ok=True) print(f"Downloading model from HF repo '{repo_id}' to '{target_dir}' ...", file=sys.stderr) snapshot_download(repo_id=repo_id, local_dir=target_dir, local_dir_use_symlinks=False) @@ -72,7 +74,7 @@ class VoxCPMDemo: if prompt_wav is None: return "" res = self.asr_model.generate(input=prompt_wav, language="auto", use_itn=True) - text = res[0]["text"].split('|>')[-1] + text = res[0]["text"].split("|>")[-1] return text def generate_tts_audio( @@ -149,11 +151,13 @@ _CUSTOM_CSS = """ def create_demo_interface(demo: VoxCPMDemo): """Build the Gradio UI for VoxCPM demo.""" - gr.set_static_paths(paths=[Path.cwd().absolute()/"assets"]) + gr.set_static_paths(paths=[Path.cwd().absolute() / "assets"]) with gr.Blocks() as interface: # Header logo - gr.HTML('
VoxCPM Logo
') + gr.HTML( + '
VoxCPM Logo
' + ) # Quick Start with gr.Accordion("📋 Quick Start Guide |快速入门", open=False, elem_id="acc_quick"): @@ -201,7 +205,7 @@ def create_demo_interface(demo: VoxCPMDemo): with gr.Row(): with gr.Column(): prompt_wav = gr.Audio( - sources=["upload", 'microphone'], + sources=["upload", "microphone"], type="filepath", label="Prompt Speech (Optional, or let VoxCPM improvise)", value="./examples/example.wav", @@ -210,13 +214,13 @@ def create_demo_interface(demo: VoxCPMDemo): value=False, label="Prompt Speech Enhancement", elem_id="chk_denoise", - info="We use ZipEnhancer model to denoise the prompt audio." + info="We use ZipEnhancer model to denoise the prompt audio.", ) with gr.Row(): prompt_text = gr.Textbox( value="Just by listening a few minutes a day, you'll be able to eliminate negative thoughts by conditioning your mind to be more positive.", label="Prompt Text", - placeholder="Please enter the prompt text. Automatic recognition is supported, and you can correct the results yourself..." + placeholder="Please enter the prompt text. Automatic recognition is supported, and you can correct the results yourself...", ) run_btn = gr.Button("Generate Speech", variant="primary") @@ -227,7 +231,7 @@ def create_demo_interface(demo: VoxCPMDemo): value=2.0, step=0.1, label="CFG Value (Guidance Scale)", - info="Higher values increase adherence to prompt, lower values allow more creativity" + info="Higher values increase adherence to prompt, lower values allow more creativity", ) inference_timesteps = gr.Slider( minimum=4, @@ -235,7 +239,7 @@ def create_demo_interface(demo: VoxCPMDemo): value=10, step=1, label="Inference Timesteps", - info="Number of inference timesteps for generation (higher values may improve quality but slower)" + info="Number of inference timesteps for generation (higher values may improve quality but slower)", ) with gr.Row(): text = gr.Textbox( @@ -247,14 +251,22 @@ def create_demo_interface(demo: VoxCPMDemo): value=False, label="Text Normalization", elem_id="chk_normalize", - info="We use wetext library to normalize the input text." + info="We use wetext library to normalize the input text.", ) audio_output = gr.Audio(label="Output Audio") # Wiring run_btn.click( fn=demo.generate_tts_audio, - inputs=[text, prompt_wav, prompt_text, cfg_value, inference_timesteps, DoNormalizeText, DoDenoisePromptAudio], + inputs=[ + text, + prompt_wav, + prompt_text, + cfg_value, + inference_timesteps, + DoNormalizeText, + DoDenoisePromptAudio, + ], outputs=[audio_output], show_progress=True, api_name="generate", @@ -277,4 +289,4 @@ def run_demo(server_name: str = "localhost", server_port: int = 7860, show_error if __name__ == "__main__": - run_demo() \ No newline at end of file + run_demo() diff --git a/lora_ft_webui.py b/lora_ft_webui.py index e4a6822..3d91c3d 100644 --- a/lora_ft_webui.py +++ b/lora_ft_webui.py @@ -308,7 +308,8 @@ def run_inference(text, prompt_wav, prompt_text, lora_selection, cfg_scale, step if new_r is not None and current_r is not None and new_r != current_r: print(f"LoRA rank mismatch (model r={current_r}, checkpoint r={new_r}), reloading...", file=sys.stderr) reload_base = ( - new_base_model if new_base_model and os.path.exists(new_base_model) + new_base_model + if new_base_model and os.path.exists(new_base_model) else (pretrained_path if pretrained_path and pretrained_path.strip() else default_pretrained_path) ) try: @@ -982,9 +983,7 @@ with gr.Blocks(title="VoxCPM LoRA WebUI", theme=gr.themes.Soft(), css=custom_css gr.Markdown("#### 分发选项 (Distribution)") with gr.Row(): - hf_model_id = gr.Textbox( - label="HuggingFace Model ID (e.g., openbmb/VoxCPM2)", value="" - ) + hf_model_id = gr.Textbox(label="HuggingFace Model ID (e.g., openbmb/VoxCPM2)", value="") distribute = gr.Checkbox(label="分发模式 (distribute)", value=False) with gr.Column(scale=2, elem_classes="form-section"): diff --git a/scripts/test_pick_runtime_dtype.py b/scripts/test_pick_runtime_dtype.py index 5bff4eb..160aba3 100644 --- a/scripts/test_pick_runtime_dtype.py +++ b/scripts/test_pick_runtime_dtype.py @@ -3,6 +3,7 @@ Loads src/voxcpm/model/utils.py directly to avoid the heavy voxcpm package init. Run with: `python scripts/test_pick_runtime_dtype.py`. """ + import importlib.util import os import pathlib diff --git a/scripts/test_voxcpm_ft_infer.py b/scripts/test_voxcpm_ft_infer.py index 1a5a769..fb2e931 100644 --- a/scripts/test_voxcpm_ft_infer.py +++ b/scripts/test_voxcpm_ft_infer.py @@ -19,6 +19,7 @@ With voice cloning: --text "Hello, this is voice cloning result." \ --prompt_audio path/to/ref.wav \ --prompt_text "Reference audio transcript" \ + --seed 42 \ --output ft_clone.wav """ @@ -86,6 +87,12 @@ def parse_args(): action="store_true", help="Enable text normalization", ) + parser.add_argument( + "--seed", + type=int, + default=None, + help="Random seed for generation (default: None)", + ) return parser.parse_args() @@ -109,6 +116,9 @@ def main(): print(f"[FT Inference] Using reference audio: {prompt_wav_path}", file=sys.stderr) print(f"[FT Inference] Reference text: {prompt_text}", file=sys.stderr) + if args.seed is not None: + print(f"[FT Inference] Using seed: {args.seed}", file=sys.stderr) + audio_np = model.generate( text=args.text, prompt_wav_path=prompt_wav_path, @@ -118,6 +128,7 @@ def main(): max_len=args.max_len, normalize=args.normalize, denoise=False, + seed=args.seed, ) # Save audio diff --git a/scripts/test_voxcpm_lora_infer.py b/scripts/test_voxcpm_lora_infer.py index 2e08bbe..9af5347 100644 --- a/scripts/test_voxcpm_lora_infer.py +++ b/scripts/test_voxcpm_lora_infer.py @@ -16,6 +16,7 @@ With voice cloning: --text "This is voice cloning result." \ --prompt_audio path/to/ref.wav \ --prompt_text "Reference audio transcript" \ + --seed 42 \ --output lora_clone.wav Note: The script reads base_model path and lora_config from lora_config.json @@ -94,6 +95,12 @@ def parse_args(): action="store_true", help="Enable text normalization", ) + parser.add_argument( + "--seed", + type=int, + default=None, + help="Random seed for generation (default: None)", + ) return parser.parse_args() @@ -130,6 +137,8 @@ def main(): print( f" LoRA config: r={lora_cfg.r}, alpha={lora_cfg.alpha}" if lora_cfg else " LoRA config: None", file=sys.stderr ) + if args.seed is not None: + print(f" Seed: {args.seed}", file=sys.stderr) # 3. Load model with LoRA (no denoiser) print(f"\n[1/2] Loading model with LoRA: {pretrained_path}", file=sys.stderr) @@ -161,6 +170,7 @@ def main(): max_len=args.max_len, normalize=args.normalize, denoise=False, + seed=args.seed, ) lora_output = out_path.with_stem(out_path.stem + "_with_lora") sf.write(str(lora_output), audio_np, model.tts_model.sample_rate) @@ -181,6 +191,7 @@ def main(): max_len=args.max_len, normalize=args.normalize, denoise=False, + seed=args.seed, ) disabled_output = out_path.with_stem(out_path.stem + "_lora_disabled") sf.write(str(disabled_output), audio_np, model.tts_model.sample_rate) @@ -201,6 +212,7 @@ def main(): max_len=args.max_len, normalize=args.normalize, denoise=False, + seed=args.seed, ) reenabled_output = out_path.with_stem(out_path.stem + "_lora_reenabled") sf.write(str(reenabled_output), audio_np, model.tts_model.sample_rate) @@ -221,6 +233,7 @@ def main(): max_len=args.max_len, normalize=args.normalize, denoise=False, + seed=args.seed, ) reset_output = out_path.with_stem(out_path.stem + "_lora_reset") sf.write(str(reset_output), audio_np, model.tts_model.sample_rate) @@ -242,6 +255,7 @@ def main(): max_len=args.max_len, normalize=args.normalize, denoise=False, + seed=args.seed, ) reload_output = out_path.with_stem(out_path.stem + "_lora_reloaded") sf.write(str(reload_output), audio_np, model.tts_model.sample_rate) diff --git a/scripts/train_voxcpm_finetune.py b/scripts/train_voxcpm_finetune.py index c3da4dc..56e2c6d 100644 --- a/scripts/train_voxcpm_finetune.py +++ b/scripts/train_voxcpm_finetune.py @@ -599,7 +599,9 @@ def generate_sample_audio( ) with torch.no_grad(): with autocast_ctx: - generated = unwrapped_model.generate(target_text=text, inference_timesteps=10, cfg_value=2.0) + generated = unwrapped_model.generate( + target_text=text, inference_timesteps=10, cfg_value=2.0, seed=42 + ) # Restore training setup # unwrapped_model.to(torch.float32) diff --git a/src/voxcpm/cli.py b/src/voxcpm/cli.py index 794659b..3074cb7 100644 --- a/src/voxcpm/cli.py +++ b/src/voxcpm/cli.py @@ -248,6 +248,7 @@ def _run_single(args, parser, *, text: str, output: str, prompt_text: str | None inference_timesteps=args.inference_timesteps, normalize=args.normalize, denoise=args.denoise and (args.prompt_audio is not None or args.reference_audio is not None), + seed=args.seed, ) import soundfile as sf @@ -332,6 +333,7 @@ def cmd_batch(args, parser): inference_timesteps=args.inference_timesteps, normalize=args.normalize, denoise=args.denoise and (prompt_audio_path is not None or reference_audio_path is not None), + seed=args.seed, ) output_file = output_dir / f"output_{i:03d}.wav" @@ -413,6 +415,12 @@ def _add_common_generation_args(parser): help="Inference steps (int, recommended 4–30, default: 10)", ) parser.add_argument("--normalize", action="store_true", help="Enable text normalization") + parser.add_argument( + "--seed", + type=int, + default=None, + help="Random seed for generation (default: None)", + ) def _add_prompt_reference_args(parser): @@ -587,6 +595,12 @@ Examples: help="Inference steps (int, recommended 4–30, default: 10)", ) batch_parser.add_argument("--normalize", action="store_true", help="Enable text normalization") + batch_parser.add_argument( + "--seed", + type=int, + default=None, + help="Random seed for generation (default: None)", + ) _add_model_args(batch_parser) _add_lora_args(batch_parser) _add_timestamp_args(batch_parser, include_output=False) diff --git a/src/voxcpm/core.py b/src/voxcpm/core.py index 8487298..1a1d839 100644 --- a/src/voxcpm/core.py +++ b/src/voxcpm/core.py @@ -196,6 +196,7 @@ class VoxCPM: retry_badcase_max_times: int = 3, retry_badcase_ratio_threshold: float = 6.0, streaming: bool = False, + seed: Optional[int] = None, ) -> Generator[np.ndarray, None, None]: """Synthesize speech for the given text and return a single waveform. @@ -218,6 +219,7 @@ class VoxCPM: retry_badcase_max_times: Maximum number of times to retry badcase. retry_badcase_ratio_threshold: Threshold for audio-to-text ratio. streaming: Whether to return a generator of audio chunks. + seed: Optional random seed for reproducibility. Returns: Generator of numpy.ndarray: 1D waveform array (float32) on CPU. Yields audio chunks for each generation step if ``streaming=True``, @@ -294,6 +296,7 @@ class VoxCPM: retry_badcase_max_times=retry_badcase_max_times, retry_badcase_ratio_threshold=retry_badcase_ratio_threshold, streaming=streaming, + seed=seed, ) if streaming: diff --git a/src/voxcpm/model/utils.py b/src/voxcpm/model/utils.py index 940fc74..f6d7463 100644 --- a/src/voxcpm/model/utils.py +++ b/src/voxcpm/model/utils.py @@ -5,9 +5,12 @@ from transformers import PreTrainedTokenizer _LOW_PRECISION_DTYPES = {"bfloat16", "bf16", "float16", "fp16"} _VALID_DTYPE_OVERRIDES = { - "bfloat16", "bf16", - "float16", "fp16", - "float32", "fp32", + "bfloat16", + "bf16", + "float16", + "fp16", + "float32", + "fp32", } @@ -21,6 +24,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. @@ -160,10 +176,7 @@ def pick_runtime_dtype(device: str, configured_dtype: str) -> str: override = os.environ.get("VOXCPM_MPS_DTYPE", "").strip().lower() if override: if override not in _VALID_DTYPE_OVERRIDES: - raise ValueError( - f"VOXCPM_MPS_DTYPE='{override}' is not one of " - f"{sorted(_VALID_DTYPE_OVERRIDES)}" - ) + raise ValueError(f"VOXCPM_MPS_DTYPE='{override}' is not one of " f"{sorted(_VALID_DTYPE_OVERRIDES)}") return override if (configured_dtype or "").lower() in _LOW_PRECISION_DTYPES: @@ -211,15 +224,13 @@ def resolve_runtime_device(device: Optional[str], configured_device: str = "cuda if explicit.startswith("cuda"): if not torch.cuda.is_available(): raise ValueError( - f"Requested device '{device}', but CUDA is not available. " - "Use device='auto' for automatic fallback." + f"Requested device '{device}', but CUDA is not available. " "Use device='auto' for automatic fallback." ) return explicit if explicit == "mps": if not _has_mps(): raise ValueError( - "Requested device 'mps', but MPS is not available. " - "Use device='auto' for automatic fallback." + "Requested device 'mps', but MPS is not available. " "Use device='auto' for automatic fallback." ) return "mps" if explicit == "cpu": diff --git a/src/voxcpm/model/voxcpm.py b/src/voxcpm/model/voxcpm.py index 445618b..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) @@ -367,6 +370,7 @@ class VoxCPMModel(nn.Module): retry_badcase_max_times: int = 3, 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, ) -> 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.") @@ -452,7 +456,12 @@ class VoxCPMModel(nn.Module): target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -471,6 +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 = last_attempt_seed yield decode_audio break else: @@ -482,6 +492,7 @@ class VoxCPMModel(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + current_seed += 1 continue else: break @@ -489,6 +500,7 @@ class VoxCPMModel(nn.Module): 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 @@ -603,6 +615,7 @@ class VoxCPMModel(nn.Module): retry_badcase_ratio_threshold: float = 6.0, streaming: bool = False, streaming_prefix_len: int = 3, + seed: Optional[int] = None, ) -> Generator[Tuple[torch.Tensor, torch.Tensor, Union[torch.Tensor, List[torch.Tensor]]], None, None]: """ Generate audio using pre-built prompt cache. @@ -678,7 +691,12 @@ class VoxCPMModel(nn.Module): # run inference target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -698,6 +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 = last_attempt_seed yield (decode_audio, target_text_token, pred_audio_feat) break else: @@ -709,12 +728,14 @@ class VoxCPMModel(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + current_seed += 1 continue else: 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 90495f1..174dea3 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, @@ -55,7 +57,9 @@ from .utils import ( # A simple function to trim audio silence using VAD, not used default -def _trim_audio_silence_vad(audio: torch.Tensor, sample_rate: int, max_silence_ms: float = 200.0, top_db: float = 35.0) -> torch.Tensor: +def _trim_audio_silence_vad( + audio: torch.Tensor, sample_rate: int, max_silence_ms: float = 200.0, top_db: float = 35.0 +) -> torch.Tensor: if audio.numel() == 0: return audio y = audio.squeeze(0).numpy() @@ -184,6 +188,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) @@ -476,6 +481,7 @@ class VoxCPM2Model(nn.Module): trim_silence_vad: bool = False, streaming: bool = False, streaming_prefix_len: int = 4, + seed: Optional[int] = 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.") @@ -633,7 +639,12 @@ class VoxCPM2Model(nn.Module): target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -651,6 +662,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 = last_attempt_seed yield decode_audio break else: @@ -662,6 +674,7 @@ class VoxCPM2Model(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + current_seed += 1 continue else: break @@ -669,10 +682,11 @@ class VoxCPM2Model(nn.Module): 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: - decode_audio = decode_audio[..., decode_patch_len * context_len:].squeeze(1).cpu() + decode_audio = decode_audio[..., decode_patch_len * context_len :].squeeze(1).cpu() else: decode_audio = decode_audio.squeeze(1).cpu() yield decode_audio @@ -793,6 +807,7 @@ class VoxCPM2Model(nn.Module): retry_badcase_ratio_threshold: float = 6.0, streaming: bool = False, streaming_prefix_len: int = 4, + seed: Optional[int] = None, ) -> Generator[Tuple[torch.Tensor, torch.Tensor, Union[torch.Tensor, List[torch.Tensor]]], None, None]: """ Generate audio using pre-built prompt cache. @@ -920,7 +935,12 @@ class VoxCPM2Model(nn.Module): # run inference target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = materialize_generation_seed(seed) + last_attempt_seed = current_seed while retry_badcase_times < retry_badcase_max_times: + last_attempt_seed = current_seed + apply_generation_seed(last_attempt_seed) + inference_result = self._inference( text_token, text_mask, @@ -938,6 +958,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 = last_attempt_seed yield (decode_audio, target_text_token, pred_audio_feat) break else: @@ -949,16 +970,18 @@ class VoxCPM2Model(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + current_seed += 1 continue else: 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: - decode_audio = decode_audio[..., decode_patch_len * context_len:].squeeze(1).cpu() + decode_audio = decode_audio[..., decode_patch_len * context_len :].squeeze(1).cpu() else: decode_audio = decode_audio.squeeze(1).cpu() yield (decode_audio, target_text_token, pred_audio_feat) diff --git a/src/voxcpm/modules/audiovae/audio_vae_v2.py b/src/voxcpm/modules/audiovae/audio_vae_v2.py index 4cd851c..c2ad590 100644 --- a/src/voxcpm/modules/audiovae/audio_vae_v2.py +++ b/src/voxcpm/modules/audiovae/audio_vae_v2.py @@ -551,8 +551,7 @@ class StreamingVAEDecoder: if x.shape[-1] >= _p: states[_k] = x[:, :, -_p:].detach() else: - prev = states.get(_k, torch.zeros(x.shape[0], x.shape[1], _p, - device=x.device, dtype=x.dtype)) + prev = states.get(_k, torch.zeros(x.shape[0], x.shape[1], _p, device=x.device, dtype=x.dtype)) states[_k] = torch.cat([prev, x], dim=-1)[:, :, -_p:].detach() return nn.Conv1d.forward(_m, x_pad) diff --git a/src/voxcpm/training/packers.py b/src/voxcpm/training/packers.py index cfcc193..02e7712 100644 --- a/src/voxcpm/training/packers.py +++ b/src/voxcpm/training/packers.py @@ -350,45 +350,80 @@ class AudioFeatureProcessingPacker: # -- text token track -- # [103, 0×R, 104, text_ids, 101, 0×A, 102] - text_token_info = torch.cat([ - _tok([self.audio_prompt_start_id]), - torch.zeros(ref_len, dtype=torch.int32, device=device), - _tok([self.audio_prompt_end_id]), - text_token, - _tok([self.audio_start_id]), - torch.zeros(tgt_len, dtype=torch.int32, device=device), - _tok([self.audio_end_id]), - ]) + text_token_info = torch.cat( + [ + _tok([self.audio_prompt_start_id]), + torch.zeros(ref_len, dtype=torch.int32, device=device), + _tok([self.audio_prompt_end_id]), + text_token, + _tok([self.audio_start_id]), + torch.zeros(tgt_len, dtype=torch.int32, device=device), + _tok([self.audio_end_id]), + ] + ) # -- audio feature track -- zero_1 = torch.zeros((1,) + feat_shape, dtype=torch.float32, device=device) zero_txt = torch.zeros((txt_len,) + feat_shape, dtype=torch.float32, device=device) - audio_feat_info = torch.cat([ - zero_1, ref_feats, zero_1, # 103, ref, 104 - zero_txt, # text - zero_1, tgt_feats, zero_1, # 101, target, 102 - ], dim=0) + audio_feat_info = torch.cat( + [ + zero_1, + ref_feats, + zero_1, # 103, ref, 104 + zero_txt, # text + zero_1, + tgt_feats, + zero_1, # 101, target, 102 + ], + dim=0, + ) # -- masks -- - text_mask = torch.cat([ - torch.ones(1), torch.zeros(ref_len), torch.ones(1), - torch.ones(txt_len), - torch.ones(1), torch.zeros(tgt_len), torch.ones(1), - ]).to(torch.int32).to(device) + text_mask = ( + torch.cat( + [ + torch.ones(1), + torch.zeros(ref_len), + torch.ones(1), + torch.ones(txt_len), + torch.ones(1), + torch.zeros(tgt_len), + torch.ones(1), + ] + ) + .to(torch.int32) + .to(device) + ) - audio_mask = torch.cat([ - torch.zeros(1), torch.ones(ref_len), torch.zeros(1), - torch.zeros(txt_len), - torch.zeros(1), torch.ones(tgt_len), torch.zeros(1), - ]).to(torch.int32).to(device) + audio_mask = ( + torch.cat( + [ + torch.zeros(1), + torch.ones(ref_len), + torch.zeros(1), + torch.zeros(txt_len), + torch.zeros(1), + torch.ones(tgt_len), + torch.zeros(1), + ] + ) + .to(torch.int32) + .to(device) + ) - loss_mask = torch.cat([ - torch.zeros(1 + ref_len + 1), # ref part: no loss - torch.zeros(txt_len), # text: no loss - torch.zeros(1), # 101: no loss - torch.ones(tgt_len), # target audio: LOSS - torch.zeros(1), # 102: no loss - ]).to(torch.int32).to(device) + loss_mask = ( + torch.cat( + [ + torch.zeros(1 + ref_len + 1), # ref part: no loss + torch.zeros(txt_len), # text: no loss + torch.zeros(1), # 101: no loss + torch.ones(tgt_len), # target audio: LOSS + torch.zeros(1), # 102: no loss + ] + ) + .to(torch.int32) + .to(device) + ) total_len = 1 + ref_len + 1 + txt_len + 1 + tgt_len + 1 labels = torch.zeros(total_len, dtype=torch.int32, device=device) diff --git a/src/voxcpm/training/validate.py b/src/voxcpm/training/validate.py index bad32a0..125016a 100644 --- a/src/voxcpm/training/validate.py +++ b/src/voxcpm/training/validate.py @@ -44,10 +44,7 @@ def _check_audio_file(audio_path: str, sample_rate: int) -> Optional[str]: if info.frames == 0: return f"Audio file is empty: {audio_path}" if info.samplerate != sample_rate: - return ( - f"Sample rate mismatch in {audio_path}: " - f"expected {sample_rate} Hz, got {info.samplerate} Hz" - ) + return f"Sample rate mismatch in {audio_path}: " f"expected {sample_rate} Hz, got {info.samplerate} Hz" return None except ImportError: # soundfile not available; just check existence @@ -187,13 +184,9 @@ def validate_manifest( if duration is not None: result.audio_durations.append(duration) if duration < 0.3: - result.warnings.append( - f"Line {i + 1}: Very short audio ({duration:.2f}s)" - ) + result.warnings.append(f"Line {i + 1}: Very short audio ({duration:.2f}s)") elif duration > 30.0: - result.warnings.append( - f"Line {i + 1}: Very long audio ({duration:.1f}s), may cause OOM" - ) + result.warnings.append(f"Line {i + 1}: Very long audio ({duration:.1f}s), may cause OOM") else: result.errors.append(f"Line {i + 1}: Invalid audio path") has_error = True @@ -209,9 +202,7 @@ def validate_manifest( if os.path.isfile(ref_path): result.has_ref_audio += 1 else: - result.warnings.append( - f"Line {i + 1}: ref_audio file not found: {ref_path}" - ) + result.warnings.append(f"Line {i + 1}: ref_audio file not found: {ref_path}") if not has_error: result.valid_samples += 1 @@ -222,14 +213,10 @@ def validate_manifest( # Summarize truncated errors if missing_audio_count > 5: result.errors.append( - f"... and {missing_audio_count - 5} more missing audio files " - f"({missing_audio_count} total)" + f"... and {missing_audio_count - 5} more missing audio files " f"({missing_audio_count} total)" ) if empty_text_count > 5: - result.warnings.append( - f"... and {empty_text_count - 5} more empty text entries " - f"({empty_text_count} total)" - ) + result.warnings.append(f"... and {empty_text_count - 5} more empty text entries " f"({empty_text_count} total)") return result diff --git a/tests/test_cli.py b/tests/test_cli.py index 9402123..acba60a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -676,3 +676,61 @@ def test_detect_model_architecture_uses_local_configs(): assert cli.detect_model_architecture(v1_args) == "voxcpm" assert cli.detect_model_architecture(v2_args) == "voxcpm2" + + +def test_parser_accepts_seed(): + parser = cli._build_parser() + # Default seed should be None + args = parser.parse_args(["design", "--text", "hello", "--output", "out.wav"]) + assert args.seed is None + + # Custom seed should be parsed as int + args = parser.parse_args(["design", "--text", "hello", "--output", "out.wav", "--seed", "42"]) + assert args.seed == 42 + + +def test_design_subcommand_passes_seed(monkeypatch, tmp_path): + dummy_model = DummyModel() + monkeypatch.setattr(cli, "load_model", lambda args: dummy_model) + patch_soundfile_write(monkeypatch) + + run_main( + monkeypatch, + [ + "design", + "--text", + "hello", + "--seed", + "123", + "--output", + str(tmp_path / "out.wav"), + ], + ) + + assert dummy_model.calls[0]["seed"] == 123 + + +def test_batch_subcommand_passes_seed(monkeypatch, tmp_path): + dummy_model = DummyModel() + input_file = tmp_path / "texts.txt" + input_file.write_text("hello\nworld\n", encoding="utf-8") + + monkeypatch.setattr(cli, "load_model", lambda args: dummy_model) + patch_soundfile_write(monkeypatch) + + run_main( + monkeypatch, + [ + "batch", + "--input", + str(input_file), + "--output-dir", + str(tmp_path / "outs"), + "--seed", + "999", + ], + ) + + assert len(dummy_model.calls) == 2 + assert dummy_model.calls[0]["seed"] == 999 + assert dummy_model.calls[1]["seed"] == 999 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)] diff --git a/tests/test_validate.py b/tests/test_validate.py index 06e1b8a..039c5be 100644 --- a/tests/test_validate.py +++ b/tests/test_validate.py @@ -207,6 +207,7 @@ class TestValidateManifest: audio = tmp_dir / "audio_8k.wav" import numpy as np + samples = np.zeros(8000, dtype=np.float32) sf.write(str(audio), samples, 8000) @@ -240,6 +241,7 @@ class TestValidateManifest: def test_cli_validate_exit_code(self, tmp_dir): """validate subcommand must exit 1 on validation error (missing audio).""" import subprocess + manifest = tmp_dir / "bad.jsonl" _write_manifest(manifest, [{"text": "hi", "audio": "/nonexistent/x.wav"}])