feat: add seed support for reproducible generation in v1 and v2

- Exposed 'seed' parameter in VoxCPMModel and VoxCPM2Model generation methods.
- Added PyTorch RNG seed setting before inference runs.
- Handled 'retry_badcase' seed adjustment by incrementing the seed value on retries.
- Exposed 'self.last_successful_seed' as a model attribute for UI integrations.
- Propagated 'seed' parameter to high-level pipeline class and CLI tools (cli.py).
- Added '--seed' flag to full-finetune and LoRA inference scripts.
- Configured validation audio generation in training script to use a fixed seed for objective comparison on TensorBoard.
- Added comprehensive unit tests in CLI test files to validate seed parsing and propagation.
- Updated English and Chinese READMEs with seed usage examples.
This commit is contained in:
Eliseu Silva
2026-06-01 15:34:09 -03:00
parent f3b65758c6
commit 9f1548b631
10 changed files with 159 additions and 5 deletions
+5
View File
@@ -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)
+6 -1
View File
@@ -10,7 +10,7 @@
<a href="https://voxcpm.readthedocs.io/zh-cn/latest/"><img src="https://img.shields.io/badge/Docs-ReadTheDocs-8CA1AF" alt="Documentation"></a>
<a href="https://huggingface.co/openbmb/VoxCPM2"><img src="https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-VoxCPM2-yellow" alt="Hugging Face"></a>
<a href="https://modelscope.cn/models/OpenBMB/VoxCPM2"><img src="https://img.shields.io/badge/ModelScope-VoxCPM2-purple" alt="ModelScope"></a>
<a href="https://openbmb.github.io/voxcpm2-demopage/"><img src="https://img.shields.io/badge/DemoPage-Audio Samples-red"></a>
<a href="https://openbmb.github.io/voxcpm2-demopage/"><img src="https://img.shields.io/badge/DemoPage-Audio Samples-red" alt="DemoPage"></a>
</p>
@@ -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
# 声音克隆(参考音频)
+12 -1
View File
@@ -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()
+15 -1
View File
@@ -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()
+1 -1
View File
@@ -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)
+15 -1
View File
@@ -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()
+3
View File
@@ -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:
+22
View File
@@ -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
+22
View File
@@ -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
+58
View File
@@ -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