fix: track successful generation seed
This commit is contained in:
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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)]
|
||||
|
||||
Reference in New Issue
Block a user