diff --git a/README.md b/README.md index dc1a544..c1c9931 100644 --- a/README.md +++ b/README.md @@ -111,6 +111,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") @@ -134,6 +135,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) ``` @@ -147,6 +149,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) ``` @@ -167,6 +170,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) ``` @@ -213,6 +217,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 ec781e4..c019235 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/scripts/test_voxcpm_ft_infer.py b/scripts/test_voxcpm_ft_infer.py index 1a5a769..41d0378 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 @@ -132,4 +143,4 @@ def main(): if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/scripts/test_voxcpm_lora_infer.py b/scripts/test_voxcpm_lora_infer.py index 2e08bbe..df66cfa 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) @@ -259,4 +273,4 @@ def main(): if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/scripts/train_voxcpm_finetune.py b/scripts/train_voxcpm_finetune.py index c3da4dc..590258d 100644 --- a/scripts/train_voxcpm_finetune.py +++ b/scripts/train_voxcpm_finetune.py @@ -599,7 +599,7 @@ 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 f6d40d6..ab4e558 100644 --- a/src/voxcpm/cli.py +++ b/src/voxcpm/cli.py @@ -261,6 +261,7 @@ def _run_single(args, parser, *, text: str, output: str, prompt_text: str | None 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 @@ -348,6 +349,7 @@ def cmd_batch(args, parser): 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" @@ -390,6 +392,12 @@ def _add_common_generation_args(parser): 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): @@ -548,6 +556,12 @@ Examples: 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) @@ -649,4 +663,4 @@ def main(): if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/src/voxcpm/core.py b/src/voxcpm/core.py index 81692c9..2ec7db4 100644 --- a/src/voxcpm/core.py +++ b/src/voxcpm/core.py @@ -193,6 +193,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. @@ -215,6 +216,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``, @@ -291,6 +293,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/voxcpm.py b/src/voxcpm/model/voxcpm.py index 445618b..0e28fdf 100644 --- a/src/voxcpm/model/voxcpm.py +++ b/src/voxcpm/model/voxcpm.py @@ -367,6 +367,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 +453,13 @@ class VoxCPMModel(nn.Module): target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = 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) + inference_result = self._inference( text_token, text_mask, @@ -471,6 +478,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 yield decode_audio break else: @@ -482,8 +490,11 @@ class VoxCPMModel(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + if current_seed is not None: + current_seed += 1 continue else: + self.last_successful_seed = current_seed break else: break @@ -603,6 +614,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 +690,13 @@ class VoxCPMModel(nn.Module): # run inference target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = 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) + 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 = current_seed yield (decode_audio, target_text_token, pred_audio_feat) break else: @@ -709,8 +728,11 @@ class VoxCPMModel(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + if current_seed is not None: + current_seed += 1 continue else: + self.last_successful_seed = current_seed break else: break diff --git a/src/voxcpm/model/voxcpm2.py b/src/voxcpm/model/voxcpm2.py index 90495f1..2fc2986 100644 --- a/src/voxcpm/model/voxcpm2.py +++ b/src/voxcpm/model/voxcpm2.py @@ -476,6 +476,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 +634,13 @@ class VoxCPM2Model(nn.Module): target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = 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) + inference_result = self._inference( text_token, text_mask, @@ -651,6 +658,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 yield decode_audio break else: @@ -662,8 +670,11 @@ class VoxCPM2Model(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + if current_seed is not None: + current_seed += 1 continue else: + self.last_successful_seed = current_seed break else: break @@ -793,6 +804,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 +932,13 @@ class VoxCPM2Model(nn.Module): # run inference target_text_length = len(self.text_tokenizer(target_text)) retry_badcase_times = 0 + current_seed = 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) + inference_result = self._inference( text_token, text_mask, @@ -938,6 +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 yield (decode_audio, target_text_token, pred_audio_feat) break else: @@ -949,8 +968,11 @@ class VoxCPM2Model(nn.Module): file=sys.stderr, ) retry_badcase_times += 1 + if current_seed is not None: + current_seed += 1 continue else: + self.last_successful_seed = current_seed break else: break diff --git a/tests/test_cli.py b/tests/test_cli.py index 2b568b1..da0352f 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -546,3 +546,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 \ No newline at end of file