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:
co-authored by
github-actions[bot]
parent
3ea1ba5b15
commit
f651b48764
@@ -38,8 +38,28 @@ CASES = [
|
|||||||
Case("video_3d_small", (1, 64, 4, 16, 16), 32),
|
Case("video_3d_small", (1, 64, 4, 16, 16), 32),
|
||||||
Case("threshold_3d", (1, 128, 1, 256, 256), 32),
|
Case("threshold_3d", (1, 128, 1, 256, 256), 32),
|
||||||
Case("hunyuan_video_large", (1, 128, 20, 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:
|
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]:
|
def parse_cases(text: str) -> list[Case]:
|
||||||
if text == "all":
|
if text == "all":
|
||||||
return CASES
|
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()]
|
names = [item.strip() for item in text.split(",") if item.strip()]
|
||||||
missing = sorted(set(names) - CASE_BY_NAME.keys())
|
missing = sorted(set(names) - CASE_BY_NAME.keys())
|
||||||
if missing:
|
if missing:
|
||||||
@@ -239,7 +263,15 @@ def main() -> None:
|
|||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="Benchmark fused GroupNorm+SiLU against PyTorch GroupNorm+SiLU."
|
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("--dtypes", default="bf16,fp16")
|
||||||
parser.add_argument("--rounds", type=int, default=3)
|
parser.add_argument("--rounds", type=int, default=3)
|
||||||
parser.add_argument("--warmup", type=int, default=25)
|
parser.add_argument("--warmup", type=int, default=25)
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import torch
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from einops import rearrange
|
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 (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
LayerwiseOffloadableModuleMixin,
|
LayerwiseOffloadableModuleMixin,
|
||||||
)
|
)
|
||||||
@@ -110,8 +111,11 @@ class ResBlock(torch.nn.Module):
|
|||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
residual = x
|
residual = x
|
||||||
x = self.conv1(x)
|
x = self.conv1(x)
|
||||||
x = self.norm1(x)
|
# Fused GroupNorm + SiLU on the first norm of the block. The second
|
||||||
x = self.activation(x)
|
# 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.conv2(x)
|
||||||
x = self.norm2(x)
|
x = self.norm2(x)
|
||||||
x = self.activation(x + residual)
|
x = self.activation(x + residual)
|
||||||
@@ -242,8 +246,7 @@ class LatentUpsampler(torch.nn.Module, LayerwiseOffloadableModuleMixin):
|
|||||||
if self.dims == 2:
|
if self.dims == 2:
|
||||||
x = rearrange(latent, "b c f h w -> (b f) c h w")
|
x = rearrange(latent, "b c f h w -> (b f) c h w")
|
||||||
x = self.initial_conv(x)
|
x = self.initial_conv(x)
|
||||||
x = self.initial_norm(x)
|
x = apply_group_norm_silu(x, self.initial_norm, self.initial_activation)
|
||||||
x = self.initial_activation(x)
|
|
||||||
for block in self.res_blocks:
|
for block in self.res_blocks:
|
||||||
x = block(x)
|
x = block(x)
|
||||||
x = self.upsampler(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)
|
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
|
||||||
else:
|
else:
|
||||||
x = self.initial_conv(latent)
|
x = self.initial_conv(latent)
|
||||||
x = self.initial_norm(x)
|
x = apply_group_norm_silu(x, self.initial_norm, self.initial_activation)
|
||||||
x = self.initial_activation(x)
|
|
||||||
for block in self.res_blocks:
|
for block in self.res_blocks:
|
||||||
x = block(x)
|
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)
|
||||||
Reference in New Issue
Block a user