From fc584aaa226560ccdfe70c2bcfe9424af1adeb04 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Tue, 22 Sep 2026 04:56:39 +0300 Subject: [PATCH] Blend H3 VAE tiles against composited neighbours (#16436) --- comfy/ldm/minimax/vae.py | 23 ++++++++++++----------- 1 file changed, 12 insertions(+), 11 deletions(-) diff --git a/comfy/ldm/minimax/vae.py b/comfy/ldm/minimax/vae.py index 36d9abb57..cc4c16de7 100644 --- a/comfy/ldm/minimax/vae.py +++ b/comfy/ldm/minimax/vae.py @@ -599,38 +599,39 @@ class MiniMaxH3VideoVAE(nn.Module): y_idx, y_len, y_overlap = self.split_tiles(height) x_idx, x_len, x_overlap = self.split_tiles(width) - # Blended tiles are written straight into a pre-allocated canvas. + # Blended tiles are written straight into a pre-allocated canvas. Tiles blend against + # their neighbours as already blended, so seams stay continuous where overlaps cross. canvas = None - row_tails = [] + strip = None out_y = 0 for i, (i_pos, i_len) in enumerate(zip(y_idx, y_len)): zi, zl = i_pos // self.vae_ratio, i_len // self.vae_ratio tiles = self._decode_tile_row(z[..., zi:zi + zl, :], x_idx, x_len) - new_tails = [] + new_strip = None left_tail = None out_x = 0 # enumerate would retain the previous tile while the next batch decodes. for j in range(len(x_idx)): tile = next(tiles) - if i < len(y_idx) - 1: - new_tails.append(tile[..., -y_overlap[i]:, :].clone()) - next_left_tail = tile[..., :, -x_overlap[j]:].clone() if j < len(x_idx) - 1 else None if i > 0: - tile = self.blend(row_tails[j], tile, y_overlap[i - 1], dim=-2) + tile = self.blend(strip[..., :, x_idx[j]:x_idx[j] + x_len[j]], tile, y_overlap[i - 1], dim=-2) if j > 0: tile = self.blend(left_tail, tile, x_overlap[j - 1], dim=-1) - left_tail = next_left_tail - if i < len(y_idx) - 1: - tile = tile[..., :-y_overlap[i], :] + left_tail = tile[..., :, -x_overlap[j]:].clone() if j < len(x_idx) - 1 else None if j < len(x_idx) - 1: tile = tile[..., :, :-x_overlap[j]] if canvas is None: canvas = torch.empty(*tile.shape[:-2], height, width, dtype=tile.dtype, device=tile.device) + if i < len(y_idx) - 1: + if new_strip is None: + new_strip = torch.empty(*tile.shape[:-2], y_overlap[i], width, dtype=tile.dtype, device=tile.device) + new_strip[..., :, out_x:out_x + tile.shape[-1]] = tile[..., -y_overlap[i]:, :] + tile = tile[..., :-y_overlap[i], :] canvas[..., out_y:out_y + tile.shape[-2], out_x:out_x + tile.shape[-1]].copy_(tile) tile_height = tile.shape[-2] out_x += tile.shape[-1] del tile - row_tails = new_tails + strip = new_strip out_y += tile_height return canvas