fix: track successful generation seed

This commit is contained in:
Eliseu Silva
2026-06-24 22:26:33 -03:00
parent b567707deb
commit 5ef0b3db4c
5 changed files with 130 additions and 45 deletions
+57 -5
View File
@@ -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",
)
+13
View File
@@ -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.
+19 -20
View File
@@ -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:
+19 -20
View File
@@ -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:
+22
View File
@@ -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)]