From 0619df8dff7877a1cdcdad5a4b9f1c2ad4de5249 Mon Sep 17 00:00:00 2001 From: debpalash <4178343+debpalash@users.noreply.github.com> Date: Thu, 20 Aug 2026 10:04:00 +0530 Subject: [PATCH] fix(cloning): retain selected passage volume --- omnivoice/models/omnivoice.py | 9 +++++++-- tests/test_omnivoice_reference_bound.py | 19 +++++++++++++++++++ 2 files changed, 26 insertions(+), 2 deletions(-) diff --git a/omnivoice/models/omnivoice.py b/omnivoice/models/omnivoice.py index 0da89c4c..99c942df 100644 --- a/omnivoice/models/omnivoice.py +++ b/omnivoice/models/omnivoice.py @@ -785,8 +785,10 @@ class OmniVoice(PreTrainedModel): # than the trim thresholds but real is recovered below, so this is # the only remaining hard failure for a reference clip. validate_clone_reference(ref_wav, ref_rms) + input_gain = 1.0 if 0 < ref_rms < 0.1: - ref_wav = ref_wav * 0.1 / ref_rms + input_gain = 0.1 / ref_rms + ref_wav = ref_wav * input_gain ref_duration = ref_wav.size(-1) / self.sampling_rate if ref_text is not None and ref_duration > 20.0: @@ -881,6 +883,9 @@ class OmniVoice(PreTrainedModel): chunk_size = self.audio_tokenizer.config.hop_length clip_size = int(ref_wav.size(-1) % chunk_size) ref_wav = ref_wav[:, :-clip_size] if clip_size > 0 else ref_wav + aligned_rms = torch.sqrt(torch.mean(torch.square(ref_wav))).item() + validate_clone_reference(ref_wav, aligned_rms) + selected_ref_rms = aligned_rms / input_gain ref_audio_tokens = self.audio_tokenizer.encode( ref_wav.unsqueeze(0).to(self.audio_tokenizer.device), ).audio_codes.squeeze( @@ -893,7 +898,7 @@ class OmniVoice(PreTrainedModel): return VoiceClonePrompt( ref_audio_tokens=ref_audio_tokens, ref_text=ref_text, - ref_rms=ref_rms, + ref_rms=selected_ref_rms, ) def _decode_and_post_process( diff --git a/tests/test_omnivoice_reference_bound.py b/tests/test_omnivoice_reference_bound.py index 0d99bb67..f0b688e3 100644 --- a/tests/test_omnivoice_reference_bound.py +++ b/tests/test_omnivoice_reference_bound.py @@ -171,6 +171,25 @@ def test_speech_aware_bound_prefers_quiet_voice_over_loud_non_speech(monkeypatch assert model.audio_tokenizer.seen_peak < 0.1 +def test_selected_passage_rms_undoes_input_gain(monkeypatch): + model = _model(reject_tokenization=False) + model._asr_pipe = object() + model.transcribe = lambda candidate: ( + "Selected speech." if float(candidate[0].abs().mean()) > 0.1 else "" + ) + monkeypatch.setattr( + "omnivoice.models.omnivoice.remove_silence_safe", + lambda audio, *_args, **_kwargs: audio, + ) + audio = torch.full((1, 30 * 24_000), 0.01) + audio[:, 15 * 24_000 :] = 0.02 + + prompt = model.create_voice_clone_prompt((audio, 24_000), ref_text=None) + + assert prompt.ref_text == "Selected speech." + assert prompt.ref_rms == pytest.approx(0.02) + + def test_far_late_speech_at_auto_select_limit_survives_with_five_asr_calls(monkeypatch): model = _model(reject_tokenization=False) model.sampling_rate = 100