Apply apply_group_norm_silu to LTX-2 latent upsampler (#26045)

Signed-off-by: pengdurice <pengduhit@gmail.com>
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
pengdurice
2026-06-02 17:28:57 +08:00
committed by GitHub
co-authored by github-actions[bot]
parent 3ea1ba5b15
commit f651b48764
3 changed files with 348 additions and 8 deletions
@@ -38,8 +38,28 @@ CASES = [
Case("video_3d_small", (1, 64, 4, 16, 16), 32),
Case("threshold_3d", (1, 128, 1, 256, 256), 32),
Case("hunyuan_video_large", (1, 128, 20, 256, 256), 32),
# LTX-2 latent upsampler (`LatentUpsampler` + `ResBlock`) operates on
# `[B, mid_channels=512, F, H, W]` tensors with num_groups=32. The
# `small` and `pre_720p` cases stay in the default set; the larger
# `post_720p` case is opt-in via LARGE_CASES below.
Case("ltx2_upsampler_small", (1, 512, 8, 45, 80), 32),
Case("ltx2_upsampler_pre_720p", (1, 512, 16, 90, 160), 32),
]
CASE_BY_NAME = {case.name: case for case in CASES}
# Cases too large to fit comfortably alongside the native-path intermediates
# on consumer GPUs (e.g. 24 GB L4). Opt in with `--cases large` (large only),
# `--cases all-large` (default + large), or by name.
#
# `ltx2_upsampler_post_720p` is ~471M bf16 elements (~940 MB tensor) and the
# eager `silu(group_norm(x))` reference materializes mean / variance /
# normalized / silu intermediates -- working set lands around 5 GB. On
# H100 / H200 this is fine and surfaces the asymptotic ~14x kernel speedup;
# on a 24 GB GPU it can OOM, so it's gated out of `--cases all`.
LARGE_CASES = [
Case("ltx2_upsampler_post_720p", (1, 512, 16, 180, 320), 32),
]
CASE_BY_NAME = {case.name: case for case in CASES + LARGE_CASES}
def dtype_from_name(name: str) -> torch.dtype:
@@ -70,6 +90,10 @@ def parse_dtypes(text: str) -> list[torch.dtype]:
def parse_cases(text: str) -> list[Case]:
if text == "all":
return CASES
if text == "large":
return LARGE_CASES
if text == "all-large":
return CASES + LARGE_CASES
names = [item.strip() for item in text.split(",") if item.strip()]
missing = sorted(set(names) - CASE_BY_NAME.keys())
if missing:
@@ -239,7 +263,15 @@ def main() -> None:
parser = argparse.ArgumentParser(
description="Benchmark fused GroupNorm+SiLU against PyTorch GroupNorm+SiLU."
)
parser.add_argument("--cases", default="all")
parser.add_argument(
"--cases",
default="all",
help=(
"Comma-separated case names, or one of: 'all' (default-sized "
"cases only), 'large' (high-memory cases only -- requires "
"H100/H200-class GPU), 'all-large' (both). See CASES + LARGE_CASES."
),
)
parser.add_argument("--dtypes", default="bf16,fp16")
parser.add_argument("--rounds", type=int, default=3)
parser.add_argument("--warmup", type=int, default=25)
@@ -8,6 +8,7 @@ import torch
import torch.nn.functional as F
from einops import rearrange
from sglang.jit_kernel.diffusion.group_norm_silu import apply_group_norm_silu
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
)
@@ -110,8 +111,11 @@ class ResBlock(torch.nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
residual = x
x = self.conv1(x)
x = self.norm1(x)
x = self.activation(x)
# Fused GroupNorm + SiLU on the first norm of the block. The second
# norm (line below) is followed by `silu(norm + residual)`, which the
# current `apply_group_norm_silu` helper does not cover -- left on
# the eager path until a `group_norm_add_silu` helper exists.
x = apply_group_norm_silu(x, self.norm1, self.activation)
x = self.conv2(x)
x = self.norm2(x)
x = self.activation(x + residual)
@@ -242,8 +246,7 @@ class LatentUpsampler(torch.nn.Module, LayerwiseOffloadableModuleMixin):
if self.dims == 2:
x = rearrange(latent, "b c f h w -> (b f) c h w")
x = self.initial_conv(x)
x = self.initial_norm(x)
x = self.initial_activation(x)
x = apply_group_norm_silu(x, self.initial_norm, self.initial_activation)
for block in self.res_blocks:
x = block(x)
x = self.upsampler(x)
@@ -253,8 +256,7 @@ class LatentUpsampler(torch.nn.Module, LayerwiseOffloadableModuleMixin):
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
else:
x = self.initial_conv(latent)
x = self.initial_norm(x)
x = self.initial_activation(x)
x = apply_group_norm_silu(x, self.initial_norm, self.initial_activation)
for block in self.res_blocks:
x = block(x)
@@ -0,0 +1,306 @@
from __future__ import annotations
from unittest.mock import patch
import pytest
import torch
import sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler as lu_mod
from sglang.jit_kernel.diffusion.group_norm_silu import apply_group_norm_silu
from sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler import (
LatentUpsampler,
ResBlock,
SpatialRationalResampler,
)
def _resblock_eager_reference(block: ResBlock, x: torch.Tensor) -> torch.Tensor:
residual = x
x = block.activation(block.norm1(block.conv1(x))) # fused site
x = block.norm2(block.conv2(x))
return block.activation(x + residual)
def _latent_upsampler_eager_reference(
upsampler: LatentUpsampler, latent: torch.Tensor
) -> torch.Tensor:
from einops import rearrange
b, _, f, _, _ = latent.shape
if upsampler.dims == 2:
x = rearrange(latent, "b c f h w -> (b f) c h w")
x = upsampler.initial_activation(
upsampler.initial_norm(upsampler.initial_conv(x))
)
for block in upsampler.res_blocks:
x = _resblock_eager_reference(block, x)
x = upsampler.upsampler(x)
for block in upsampler.post_upsample_res_blocks:
x = _resblock_eager_reference(block, x)
x = upsampler.final_conv(x)
return rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
x = upsampler.initial_activation(
upsampler.initial_norm(upsampler.initial_conv(latent))
)
for block in upsampler.res_blocks:
x = _resblock_eager_reference(block, x)
if upsampler.temporal_upsample:
x = upsampler.upsampler(x)[:, :, 1:, :, :]
elif isinstance(upsampler.upsampler, SpatialRationalResampler):
x = upsampler.upsampler(x)
else:
x = rearrange(x, "b c f h w -> (b f) c h w")
x = upsampler.upsampler(x)
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
for block in upsampler.post_upsample_res_blocks:
x = _resblock_eager_reference(block, x)
return upsampler.final_conv(x)
@pytest.mark.parametrize(
"batch,channels,dims,spatial",
[
(1, 64, 2, (16, 16)),
(2, 64, 2, (8, 24)),
(1, 128, 3, (2, 8, 8)),
],
)
def test_resblock_forward_parity(batch, channels, dims, spatial):
torch.manual_seed(0)
block = ResBlock(channels=channels, dims=dims).eval()
torch.manual_seed(1)
x = torch.randn(batch, channels, *spatial, dtype=torch.float32)
with torch.no_grad():
out = block(x)
ref = _resblock_eager_reference(block, x)
torch.testing.assert_close(out, ref, atol=0.0, rtol=0.0)
@pytest.mark.parametrize(
"dims,latent_shape,mid_channels,num_blocks_per_stage,rational_resampler",
[
(2, (1, 32, 2, 16, 16), 64, 2, False),
(3, (1, 32, 2, 16, 16), 64, 2, False),
(3, (1, 32, 2, 16, 16), 64, 2, True),
],
)
def test_latent_upsampler_forward_parity(
dims, latent_shape, mid_channels, num_blocks_per_stage, rational_resampler
):
torch.manual_seed(2)
upsampler = LatentUpsampler(
in_channels=latent_shape[1],
mid_channels=mid_channels,
num_blocks_per_stage=num_blocks_per_stage,
dims=dims,
spatial_upsample=True,
temporal_upsample=False,
spatial_scale=2.0,
rational_resampler=rational_resampler,
).eval()
torch.manual_seed(3)
latent = torch.randn(*latent_shape, dtype=torch.float32)
with torch.no_grad():
out = upsampler(latent)
ref = _latent_upsampler_eager_reference(upsampler, latent)
torch.testing.assert_close(out, ref, atol=0.0, rtol=0.0)
def test_resblock_fuses_exactly_one_site():
torch.manual_seed(4)
block = ResBlock(channels=64, dims=2).eval()
x = torch.randn(1, 64, 16, 16, dtype=torch.float32)
with patch.object(
lu_mod, "apply_group_norm_silu", wraps=lu_mod.apply_group_norm_silu
) as spy:
with torch.no_grad():
block(x)
assert spy.call_count == 1
@pytest.mark.parametrize(
"dims,num_blocks_per_stage,rational_resampler,expected_calls",
[
(2, 2, False, 1 + 2 * 2),
(2, 4, False, 1 + 4 * 2),
(3, 2, False, 1 + 2 * 2),
(3, 2, True, 1 + 2 * 2),
],
)
def test_latent_upsampler_fuses_expected_sites(
dims, num_blocks_per_stage, rational_resampler, expected_calls
):
torch.manual_seed(5)
upsampler = LatentUpsampler(
in_channels=32,
mid_channels=64,
num_blocks_per_stage=num_blocks_per_stage,
dims=dims,
spatial_upsample=True,
temporal_upsample=False,
spatial_scale=2.0,
rational_resampler=rational_resampler,
).eval()
latent = torch.randn(1, 32, 2, 16, 16, dtype=torch.float32)
with patch.object(
lu_mod, "apply_group_norm_silu", wraps=lu_mod.apply_group_norm_silu
) as spy:
with torch.no_grad():
upsampler(latent)
assert spy.call_count == expected_calls
# CUDA Triton fast path -------------------------------------------------------
requires_cuda = pytest.mark.skipif(
not torch.cuda.is_available(),
reason="Triton fused group_norm_silu requires CUDA",
)
# bf16 keeps kernel-level tolerance because its fp32-equivalent exponent range
# absorbs multi-layer conv drift; fp16 needs a looser tolerance on the e2e
# upsampler test where 8+ downstream convs amplify fused-vs-eager rounding.
_RESBLOCK_TOL = {torch.bfloat16: (7e-2, 2e-2), torch.float16: (3e-3, 3e-3)}
_UPSAMPLER_TOL = {torch.bfloat16: (7e-2, 2e-2), torch.float16: (2e-2, 1e-1)}
@requires_cuda
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize(
"batch,channels,dims,spatial",
[
(1, 64, 2, (16, 16)),
(1, 128, 3, (2, 8, 8)),
],
)
def test_resblock_forward_parity_cuda(dtype, batch, channels, dims, spatial):
torch.manual_seed(0)
device = torch.device("cuda")
block = ResBlock(channels=channels, dims=dims).to(device=device, dtype=dtype).eval()
torch.manual_seed(1)
x = torch.randn(batch, channels, *spatial, device=device, dtype=dtype)
with torch.no_grad():
out = block(x)
ref = _resblock_eager_reference(block, x)
atol, rtol = _RESBLOCK_TOL[dtype]
torch.testing.assert_close(out, ref, atol=atol, rtol=rtol)
@requires_cuda
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize(
"dims,latent_shape,mid_channels,num_blocks_per_stage,rational_resampler",
[
(2, (1, 32, 2, 16, 16), 64, 2, False),
(3, (1, 32, 2, 16, 16), 64, 2, False),
(3, (1, 32, 2, 16, 16), 64, 2, True),
],
)
def test_latent_upsampler_forward_parity_cuda(
dtype, dims, latent_shape, mid_channels, num_blocks_per_stage, rational_resampler
):
torch.manual_seed(2)
device = torch.device("cuda")
upsampler = (
LatentUpsampler(
in_channels=latent_shape[1],
mid_channels=mid_channels,
num_blocks_per_stage=num_blocks_per_stage,
dims=dims,
spatial_upsample=True,
temporal_upsample=False,
spatial_scale=2.0,
rational_resampler=rational_resampler,
)
.to(device=device, dtype=dtype)
.eval()
)
torch.manual_seed(3)
latent = torch.randn(*latent_shape, device=device, dtype=dtype)
with torch.no_grad():
out = upsampler(latent)
ref = _latent_upsampler_eager_reference(upsampler, latent)
atol, rtol = _UPSAMPLER_TOL[dtype]
torch.testing.assert_close(out, ref, atol=atol, rtol=rtol)
@requires_cuda
def test_resblock_actually_uses_triton_kernel_cuda():
from sglang.jit_kernel.diffusion.triton import group_norm_silu as triton_mod
torch.manual_seed(0)
device = torch.device("cuda")
dtype = torch.bfloat16
block = ResBlock(channels=64, dims=2).to(device=device, dtype=dtype).eval()
x = torch.randn(1, 64, 16, 16, device=device, dtype=dtype)
with patch.object(
triton_mod,
"triton_group_norm_silu",
wraps=triton_mod.triton_group_norm_silu,
) as spy:
with torch.no_grad():
block(x)
assert spy.call_count >= 1
@requires_cuda
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize(
"dims,latent_shape,mid_channels",
[
(2, (1, 32, 2, 16, 16), 64),
(3, (1, 32, 2, 16, 16), 64),
],
)
def test_initial_groupnorm_silu_parity_cuda_local(
dtype, dims, latent_shape, mid_channels
):
# Sharp parity at the fused initial-norm boundary, before downstream convs
# can amplify drift; uses kernel-level tolerance instead of the looser e2e.
torch.manual_seed(2)
device = torch.device("cuda")
upsampler = (
LatentUpsampler(
in_channels=latent_shape[1],
mid_channels=mid_channels,
num_blocks_per_stage=2,
dims=dims,
spatial_upsample=True,
temporal_upsample=False,
spatial_scale=2.0,
rational_resampler=False,
)
.to(device=device, dtype=dtype)
.eval()
)
torch.manual_seed(3)
latent = torch.randn(*latent_shape, device=device, dtype=dtype)
with torch.no_grad():
if dims == 2:
from einops import rearrange
b, _, f, _, _ = latent.shape
x_in = rearrange(latent, "b c f h w -> (b f) c h w")
else:
x_in = latent
x_after_conv = upsampler.initial_conv(x_in)
out_fused = apply_group_norm_silu(
x_after_conv, upsampler.initial_norm, upsampler.initial_activation
)
out_eager = upsampler.initial_activation(upsampler.initial_norm(x_after_conv))
atol, rtol = _RESBLOCK_TOL[dtype]
torch.testing.assert_close(out_fused, out_eager, atol=atol, rtol=rtol)