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 <[email protected]>
This commit is contained in:
fooSynaptic
2026-07-26 23:31:20 -07:00
committed by GitHub
co-authored by fooSynaptic
parent 669f85024f
commit 624f6a2140
2 changed files with 74 additions and 2 deletions
+6 -2
View File
@@ -248,9 +248,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 = {}
@@ -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"