[diffusion] fix: stabilize H3 reference audio across repeated requests (#38225)
This commit is contained in:
@@ -140,8 +140,8 @@ def WNConv1d(*args, **kwargs):
|
|||||||
return weight_norm(nn.Conv1d(*args, **kwargs))
|
return weight_norm(nn.Conv1d(*args, **kwargs))
|
||||||
|
|
||||||
|
|
||||||
@torch.jit.script
|
|
||||||
def snake(x, alpha):
|
def snake(x, alpha):
|
||||||
|
# profiling JIT changes rounding after the first call when it fuses this graph
|
||||||
shape = x.shape
|
shape = x.shape
|
||||||
x = x.reshape(shape[0], shape[1], -1)
|
x = x.reshape(shape[0], shape[1], -1)
|
||||||
x = x + (alpha + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2)
|
x = x + (alpha + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2)
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ logger = init_logger(__name__)
|
|||||||
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
|
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
|
||||||
# publish.
|
# publish.
|
||||||
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
|
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "11783b3fbd8ebb1e3509cc4590c1fff476e65511"
|
SGL_TEST_FILES_CI_DATA_REVISION = "285ffa7fa7a8afcfb90e2344e3f9d6ca12a5e2ea"
|
||||||
|
|
||||||
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
|
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
|
||||||
# when it's regenerated on its own cadence.
|
# when it's regenerated on its own cadence.
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
"""MiniMax-H3 released VAE decode contract."""
|
"""MiniMax-H3 released VAE decode contract."""
|
||||||
|
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
import textwrap
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -117,3 +120,35 @@ def test_audio_vae_attention_defaults_to_local_sdpa_and_allows_fa():
|
|||||||
}
|
}
|
||||||
assert recording_fa.input_dtype == torch.bfloat16
|
assert recording_fa.input_dtype == torch.bfloat16
|
||||||
assert output.dtype == torch.float32
|
assert output.dtype == torch.float32
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA")
|
||||||
|
def test_audio_snake_first_call_matches_repeated_calls():
|
||||||
|
# a fresh process prevents earlier tests from warming a profiling JIT graph
|
||||||
|
subprocess.run(
|
||||||
|
[
|
||||||
|
sys.executable,
|
||||||
|
"-c",
|
||||||
|
textwrap.dedent(
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
from sglang.multimodal_gen.runtime.models.vaes.minimax_h3_audio_vae.audio_vae import Snake1d
|
||||||
|
|
||||||
|
torch.manual_seed(42)
|
||||||
|
activation = Snake1d(64).cuda().eval()
|
||||||
|
with torch.inference_mode():
|
||||||
|
activation.alpha.uniform_(0.1, 2.0)
|
||||||
|
x = torch.randn(2, 64, 4096, device="cuda")
|
||||||
|
original_x = x.clone()
|
||||||
|
original_alpha = activation.alpha.clone()
|
||||||
|
first = activation(x)
|
||||||
|
for _ in range(4):
|
||||||
|
torch.testing.assert_close(activation(x), first, rtol=0, atol=0)
|
||||||
|
torch.testing.assert_close(x, original_x, rtol=0, atol=0)
|
||||||
|
torch.testing.assert_close(activation.alpha, original_alpha, rtol=0, atol=0)
|
||||||
|
"""
|
||||||
|
),
|
||||||
|
],
|
||||||
|
check=True,
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user