[diffusion] improve: use conv2d width padding in wanvae (#27081)

This commit is contained in:
Mick
2026-06-03 10:14:50 +08:00
committed by GitHub
parent 13852d3f31
commit f0e18be0ba
@@ -202,22 +202,23 @@ class WanDistConv2d(nn.Conv2d):
self.padding: tuple[int, int] self.padding: tuple[int, int]
if self.height_halo_size > 0: if self.height_halo_size > 0:
self._padding = (self.padding[1], self.padding[1], 0, 0) self._padding = (0, 0, 0, 0)
else: else:
self._padding = ( self._padding = (
self.padding[1], 0,
self.padding[1], 0,
self.padding[0], self.padding[0],
self.padding[0], self.padding[0],
) )
self.padding = (0, 0) self.padding = (0, self.padding[1])
self._halo_recv_top_buf: torch.Tensor | None = None self._halo_recv_top_buf: torch.Tensor | None = None
self._halo_recv_bottom_buf: torch.Tensor | None = None self._halo_recv_bottom_buf: torch.Tensor | None = None
self.rank = get_sp_parallel_rank() self.rank = get_sp_parallel_rank()
self.world_size = get_sp_world_size() self.world_size = get_sp_world_size()
def forward(self, x): def forward(self, x):
if any(self._padding):
x = F.pad(x, self._padding) x = F.pad(x, self._padding)
x_padded, self._halo_recv_top_buf, self._halo_recv_bottom_buf = halo_exchange( x_padded, self._halo_recv_top_buf, self._halo_recv_bottom_buf = halo_exchange(