[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
|
||||
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user