[diffusion] optimize: optimize causal conv3d vae padding (#28204)
This commit is contained in:
@@ -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
|
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 = contextvars.ContextVar(
|
||||||
"spatial_parallel_decode_disabled", default=False
|
"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(
|
def split_for_parallel_decode(
|
||||||
x: torch.Tensor, upsample_count: int, world_size: int, rank: int
|
x: torch.Tensor, upsample_count: int, world_size: int, rank: int
|
||||||
):
|
):
|
||||||
@@ -551,12 +601,7 @@ class SpatialParallelCausalConv3d(nn.Conv3d):
|
|||||||
if spatial_parallel_decode_disabled():
|
if spatial_parallel_decode_disabled():
|
||||||
padding[2] = self.height_pad_top
|
padding[2] = self.height_pad_top
|
||||||
padding[3] = self.height_pad_bottom
|
padding[3] = self.height_pad_bottom
|
||||||
if cache_x is not None and self._padding[4] > 0:
|
x = causal_conv3d_cat_pad(x, cache_x, padding)
|
||||||
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 = x if current_platform.is_amp_supported() else x.to(self.weight.dtype)
|
x = x if current_platform.is_amp_supported() else x.to(self.weight.dtype)
|
||||||
|
|
||||||
if spatial_parallel_decode_disabled():
|
if spatial_parallel_decode_disabled():
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ from sglang.multimodal_gen.runtime.layers.parallel_conv import (
|
|||||||
SpatialParallelCausalConv3d,
|
SpatialParallelCausalConv3d,
|
||||||
SpatialParallelConv2d,
|
SpatialParallelConv2d,
|
||||||
SpatialParallelZeroPad2d,
|
SpatialParallelZeroPad2d,
|
||||||
|
causal_conv3d_cat_pad,
|
||||||
chunk_height_for_parallel_decode,
|
chunk_height_for_parallel_decode,
|
||||||
disable_spatial_parallel_decode,
|
disable_spatial_parallel_decode,
|
||||||
gather_and_trim_height,
|
gather_and_trim_height,
|
||||||
@@ -87,11 +88,7 @@ class QwenImageCausalConv3d(nn.Conv3d):
|
|||||||
|
|
||||||
def forward(self, x, cache_x=None):
|
def forward(self, x, cache_x=None):
|
||||||
padding = list(self._padding)
|
padding = list(self._padding)
|
||||||
if cache_x is not None and self._padding[4] > 0:
|
x = causal_conv3d_cat_pad(x, cache_x, padding)
|
||||||
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)
|
|
||||||
return super().forward(x)
|
return super().forward(x)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ from sglang.multimodal_gen.runtime.layers.parallel_conv import (
|
|||||||
SpatialParallelCausalConv3d,
|
SpatialParallelCausalConv3d,
|
||||||
SpatialParallelConv2d,
|
SpatialParallelConv2d,
|
||||||
SpatialParallelZeroPad2d,
|
SpatialParallelZeroPad2d,
|
||||||
|
causal_conv3d_cat_pad,
|
||||||
chunk_height_for_parallel_decode,
|
chunk_height_for_parallel_decode,
|
||||||
disable_spatial_parallel_decode,
|
disable_spatial_parallel_decode,
|
||||||
gather_and_trim_height,
|
gather_and_trim_height,
|
||||||
@@ -217,11 +218,7 @@ class WanCausalConv3d(nn.Conv3d):
|
|||||||
|
|
||||||
def forward(self, x, cache_x=None):
|
def forward(self, x, cache_x=None):
|
||||||
padding = list(self._padding)
|
padding = list(self._padding)
|
||||||
if cache_x is not None and self._padding[4] > 0:
|
x = causal_conv3d_cat_pad(x, cache_x, padding)
|
||||||
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 = (
|
x = (
|
||||||
x if current_platform.is_amp_supported() else x.to(self.weight.dtype)
|
x if current_platform.is_amp_supported() else x.to(self.weight.dtype)
|
||||||
) # casting needed if amp isn't supported
|
) # casting needed if amp isn't supported
|
||||||
|
|||||||
Reference in New Issue
Block a user