[diffusion] optimize: optimize causal conv3d vae padding (#28204)

This commit is contained in:
Mick
2026-06-15 20:18:08 +08:00
committed by GitHub
parent f768344b1a
commit 818808d152
4 changed files with 177 additions and 16 deletions
@@ -0,0 +1,122 @@
from __future__ import annotations
import torch
import triton # type: ignore
import triton.language as tl # type: ignore
@triton.jit
def _fused_cat_pad_5d_kernel(
x_ptr,
cache_ptr,
out_ptr,
total,
channels,
t_size,
h_size,
w_size,
cache_t,
out_t,
out_h,
out_w,
pad_d_left,
pad_h_top,
pad_w_left,
block_size: tl.constexpr,
):
offsets = tl.program_id(0) * block_size + tl.arange(0, block_size)
mask = offsets < total
ow = offsets % out_w
tmp = offsets // out_w
oh = tmp % out_h
tmp = tmp // out_h
out = tmp % out_t
tmp = tmp // out_t
oc = tmp % channels
ob = tmp // channels
iw = ow - pad_w_left
ih = oh - pad_h_top
src_t = out - pad_d_left
valid = (
mask
& (iw >= 0)
& (iw < w_size)
& (ih >= 0)
& (ih < h_size)
& (src_t >= 0)
& (src_t < cache_t + t_size)
)
from_cache = src_t < cache_t
x_t = src_t - cache_t
clamped_iw = tl.minimum(tl.maximum(iw, 0), w_size - 1)
clamped_ih = tl.minimum(tl.maximum(ih, 0), h_size - 1)
clamped_x_t = tl.minimum(tl.maximum(x_t, 0), t_size - 1)
clamped_src_t = tl.minimum(tl.maximum(src_t, 0), cache_t - 1)
x_offsets = (
((ob * channels + oc) * t_size + clamped_x_t) * h_size + clamped_ih
) * w_size + clamped_iw
cache_offsets = (
((ob * channels + oc) * cache_t + clamped_src_t) * h_size + clamped_ih
) * w_size + clamped_iw
x_vals = tl.load(x_ptr + x_offsets, mask=valid & ~from_cache, other=0.0)
cache_vals = tl.load(cache_ptr + cache_offsets, mask=valid & from_cache, other=0.0)
vals = tl.where(from_cache, cache_vals, x_vals)
tl.store(out_ptr + offsets, vals, mask=mask)
def fused_causal_conv3d_cat_pad(
x: torch.Tensor,
cache_x: torch.Tensor,
padding: list[int] | tuple[int, ...],
) -> torch.Tensor:
width_left, width_right, height_top, height_bottom, depth_left, depth_right = (
padding
)
depth_left -= cache_x.shape[2]
assert depth_left >= 0
assert depth_right == 0
assert width_left == width_right
assert height_top == height_bottom
bsz, channels, t_size, h_size, w_size = x.shape
cache_t = cache_x.shape[2]
out = torch.empty(
(
bsz,
channels,
t_size + cache_t + depth_left + depth_right,
h_size + height_top + height_bottom,
w_size + width_left + width_right,
),
device=x.device,
dtype=x.dtype,
)
block_size = 256
total = out.numel()
grid = (triton.cdiv(total, block_size),)
with torch.get_device_module().device(x.device):
_fused_cat_pad_5d_kernel[grid](
x,
cache_x,
out,
total,
channels,
t_size,
h_size,
w_size,
cache_t,
out.shape[2],
out.shape[3],
out.shape[4],
depth_left,
height_top,
width_left,
block_size,
)
return out
@@ -14,6 +14,13 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
)
from sglang.multimodal_gen.runtime.platforms import current_platform
if current_platform.is_cuda():
from sglang.jit_kernel.diffusion.triton.causal_conv3d_pad import (
fused_causal_conv3d_cat_pad,
)
else:
fused_causal_conv3d_cat_pad = None
_SPATIAL_PARALLEL_DECODE_DISABLED = contextvars.ContextVar(
"spatial_parallel_decode_disabled", default=False
)
@@ -58,6 +65,49 @@ def _tensor_chunk(x: torch.Tensor, dim: int = -2, world_size: int = 1, rank: int
)
def _can_fuse_causal_conv3d_cat_pad(
x: torch.Tensor,
cache_x: torch.Tensor | None,
padding: list[int],
) -> bool:
if cache_x is None or fused_causal_conv3d_cat_pad is None:
return False
if not x.is_cuda or not x.is_contiguous() or not cache_x.is_contiguous():
return False
if x.dim() != 5 or cache_x.dim() != 5 or x.dtype != cache_x.dtype:
return False
if x.shape[0] != cache_x.shape[0] or x.shape[1] != cache_x.shape[1]:
return False
if x.shape[3:] != cache_x.shape[3:]:
return False
width_left, width_right, height_top, height_bottom, depth_left, depth_right = (
padding
)
if width_left != width_right or height_top != height_bottom or depth_right != 0:
return False
if depth_left < cache_x.shape[2]:
return False
return bool(width_left or height_top)
def causal_conv3d_cat_pad(
x: torch.Tensor,
cache_x: torch.Tensor | None,
padding: list[int],
) -> torch.Tensor:
if cache_x is not None and padding[4] > 0:
if cache_x.device != x.device:
cache_x = cache_x.to(x.device)
if _can_fuse_causal_conv3d_cat_pad(x, cache_x, padding):
return fused_causal_conv3d_cat_pad(x, cache_x, padding)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
if any(padding):
x = F.pad(x, padding)
return x
def split_for_parallel_decode(
x: torch.Tensor, upsample_count: int, world_size: int, rank: int
):
@@ -551,12 +601,7 @@ class SpatialParallelCausalConv3d(nn.Conv3d):
if spatial_parallel_decode_disabled():
padding[2] = self.height_pad_top
padding[3] = self.height_pad_bottom
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
x = causal_conv3d_cat_pad(x, cache_x, padding)
x = x if current_platform.is_amp_supported() else x.to(self.weight.dtype)
if spatial_parallel_decode_disabled():
@@ -24,6 +24,7 @@ from sglang.multimodal_gen.runtime.layers.parallel_conv import (
SpatialParallelCausalConv3d,
SpatialParallelConv2d,
SpatialParallelZeroPad2d,
causal_conv3d_cat_pad,
chunk_height_for_parallel_decode,
disable_spatial_parallel_decode,
gather_and_trim_height,
@@ -87,11 +88,7 @@ class QwenImageCausalConv3d(nn.Conv3d):
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
x = causal_conv3d_cat_pad(x, cache_x, padding)
return super().forward(x)
@@ -40,6 +40,7 @@ from sglang.multimodal_gen.runtime.layers.parallel_conv import (
SpatialParallelCausalConv3d,
SpatialParallelConv2d,
SpatialParallelZeroPad2d,
causal_conv3d_cat_pad,
chunk_height_for_parallel_decode,
disable_spatial_parallel_decode,
gather_and_trim_height,
@@ -217,11 +218,7 @@ class WanCausalConv3d(nn.Conv3d):
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
x = causal_conv3d_cat_pad(x, cache_x, padding)
x = (
x if current_platform.is_amp_supported() else x.to(self.weight.dtype)
) # casting needed if amp isn't supported