From 0d316876386079d0bbaf5bcd849899cbec698d8a Mon Sep 17 00:00:00 2001 From: fooSynaptic <2313990450@qq.com> Date: Mon, 27 Jul 2026 14:31:20 +0800 Subject: [PATCH] fix(tada): run voice-prompt encode under torch.inference_mode (#955) Encoder.eval() alone still builds an autograd graph because parameters require grad by default. On 8GB GPUs that ballooned TADA encode VRAM far past the model footprint (issue 890). Wrap the encode forward in inference_mode and add a unit test that asserts the flag is set. Co-authored-by: fooSynaptic <19420328+fooSynaptic@users.noreply.github.com> --- backend/backends/hume_backend.py | 8 ++- .../tests/test_hume_encode_inference_mode.py | 68 +++++++++++++++++++ 2 files changed, 74 insertions(+), 2 deletions(-) create mode 100644 backend/tests/test_hume_encode_inference_mode.py diff --git a/backend/backends/hume_backend.py b/backend/backends/hume_backend.py index 85ddb4cf..f6df707d 100644 --- a/backend/backends/hume_backend.py +++ b/backend/backends/hume_backend.py @@ -247,9 +247,13 @@ class HumeTadaBackend: audio = audio.T # (samples, channels) -> (channels, samples) audio = audio.to(device) - # Encode with forced alignment + # Encode with forced alignment. + # Must run under inference_mode: encoder params still require + # grad by default, and an autograd graph across the DAC/Snake + # stack can balloon VRAM far past the model footprint (#890). text_arg = [reference_text] if reference_text else None - prompt = self.encoder(audio, text=text_arg, sample_rate=sr) + with torch.inference_mode(): + prompt = self.encoder(audio, text=text_arg, sample_rate=sr) # Serialize EncoderOutput to a dict of CPU tensors for caching prompt_dict = {} diff --git a/backend/tests/test_hume_encode_inference_mode.py b/backend/tests/test_hume_encode_inference_mode.py new file mode 100644 index 00000000..9da1120d --- /dev/null +++ b/backend/tests/test_hume_encode_inference_mode.py @@ -0,0 +1,68 @@ +"""Ensure TADA voice-prompt encoding disables autograd (#890).""" + +from __future__ import annotations + +from dataclasses import dataclass +from unittest.mock import AsyncMock + +import numpy as np +import pytest +import soundfile as sf +import torch + +from backend.backends.hume_backend import HumeTadaBackend + + +@dataclass +class _FakeEncoderOutput: + emb: torch.Tensor + + +class _GradTrackingEncoder: + """Raises unless called under torch.inference_mode().""" + + def __init__(self) -> None: + self.called_under_inference_mode = False + + def __call__(self, audio, text=None, sample_rate=None): + self.called_under_inference_mode = torch.is_inference_mode_enabled() + if not self.called_under_inference_mode: + raise AssertionError("encoder forward must run under inference_mode") + # Touch a requires_grad tensor the way Snake1d alpha would. + alpha = torch.nn.Parameter(torch.ones(1, device=audio.device)) + _ = audio.mean() * alpha + return _FakeEncoderOutput(emb=torch.zeros(1, 4, device=audio.device)) + + +@pytest.mark.asyncio +async def test_create_voice_prompt_runs_encoder_under_inference_mode(tmp_path, monkeypatch): + wav = tmp_path / "ref.wav" + sf.write(str(wav), np.zeros(24000, dtype=np.float32), 24000) + + backend = HumeTadaBackend() + backend.model = object() # mark loaded + backend.model_size = "1B" + backend._device = "cpu" + encoder = _GradTrackingEncoder() + backend.encoder = encoder + + monkeypatch.setattr(backend, "load_model", AsyncMock(return_value=None)) + monkeypatch.setattr( + "backend.backends.hume_backend.get_cached_voice_prompt", + lambda key: None, + ) + monkeypatch.setattr( + "backend.backends.hume_backend.cache_voice_prompt", + lambda key, value: None, + ) + + prompt, from_cache = await backend.create_voice_prompt( + str(wav), + reference_text="hello world", + use_cache=False, + ) + + assert from_cache is False + assert encoder.called_under_inference_mode is True + assert isinstance(prompt["emb"], torch.Tensor) + assert prompt["emb"].device.type == "cpu"