diff --git a/python/sglang/jit_kernel/diffusion/triton/causal_conv3d_pad.py b/python/sglang/jit_kernel/diffusion/triton/causal_conv3d_pad.py new file mode 100644 index 000000000..5846a1084 --- /dev/null +++ b/python/sglang/jit_kernel/diffusion/triton/causal_conv3d_pad.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py b/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py index 068a6b925..780d354c9 100644 --- a/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py +++ b/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py @@ -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(): diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py index cad3530f1..ea7fff1fa 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py index 0f3e2f424..1bbad1241 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py @@ -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