[diffusion] fix: fix Wan channels_last_3d VAE decode corruption (#25985)

This commit is contained in:
HuangJi
2026-05-22 23:34:41 +08:00
committed by GitHub
parent 6baa859a86
commit 9bbd519f34
3 changed files with 31 additions and 3 deletions
@@ -7,6 +7,20 @@ import torch.nn.functional as F
from sglang.multimodal_gen.runtime.platforms import current_platform
def _conv3d_weight_is_channels_last_3d(weight: torch.Tensor) -> bool:
return (
weight.dim() == 5
and hasattr(torch, "channels_last_3d")
and weight.is_contiguous(memory_format=torch.channels_last_3d)
)
def match_conv3d_input_format(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
if x.dim() == 5 and _conv3d_weight_is_channels_last_3d(weight):
return x.contiguous(memory_format=torch.channels_last_3d)
return x
class AvgDown3D(nn.Module):
def __init__(
self,
@@ -151,6 +165,7 @@ class WanCausalConv3d(nn.Conv3d):
x = (
x if current_platform.is_amp_supported() else x.to(self.weight.dtype)
) # casting needed if amp isn't supported
x = match_conv3d_input_format(x, self.weight)
return super().forward(x)
@@ -18,6 +18,7 @@ from sglang.multimodal_gen.runtime.models.vaes.parallel.wan_common_utils import
WanRMS_norm,
WanUpsample,
attention_block_forward,
match_conv3d_input_format,
mid_block_forward,
resample_forward,
residual_block_forward,
@@ -91,10 +92,21 @@ def split_for_parallel_decode(
return x, expected_height
def _maybe_contiguous_for_sp_gather(x: torch.Tensor) -> torch.Tensor:
if (
x.dim() == 5
and hasattr(torch, "channels_last_3d")
and x.is_contiguous(memory_format=torch.channels_last_3d)
and not x.is_contiguous()
):
return x.contiguous()
return x
def gather_and_trim_height(x: torch.Tensor, expected_height: int | None):
if expected_height is None:
return x
x = get_sp_group().all_gather(x, dim=-2)
x = get_sp_group().all_gather(_maybe_contiguous_for_sp_gather(x), dim=-2)
if x.shape[-2] != expected_height:
x = x[..., :expected_height, :].contiguous()
return x
@@ -323,6 +335,7 @@ class WanDistCausalConv3d(nn.Conv3d):
x_padded = x_padded[..., shift:, :]
global_start += shift
x_padded = match_conv3d_input_format(x_padded, self.weight)
out = super().forward(x_padded)
if self.height_halo_size == 0:
@@ -484,7 +497,7 @@ class WanDistAttentionBlock(nn.Module):
def forward(self, x):
if self.world_size > 1:
x = self.sp_group.all_gather(x, dim=-2)
x = self.sp_group.all_gather(_maybe_contiguous_for_sp_gather(x), dim=-2)
x = x.contiguous()
x = attention_block_forward(self, x)
if self.world_size > 1:
@@ -33,7 +33,7 @@ if TYPE_CHECKING:
logger = init_logger(__name__)
SGL_TEST_FILES_CI_DATA_REVISION = "94eab4fcca6d4ddc77cdb3622f13033b61e81002"
SGL_TEST_FILES_CI_DATA_REVISION = "5b728ad6dee869c0c720ef3b668ac7a0d98b0f9a"
SGL_TEST_FILES_CONSISTENCY_GT_ROOT = (
"https://raw.githubusercontent.com/"
f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/"