From 901e485204e1c519c16483ace713f293b4d2a6ff Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 26 Jul 2026 08:22:58 +0200 Subject: [PATCH] possible fix seedvr Signed-off-by: Vladimir Mandic --- .../video_vae_v3/modules/attn_video_vae.py | 164 +++++------------- scripts/postprocessing_seedvr.py | 2 +- 2 files changed, 48 insertions(+), 118 deletions(-) diff --git a/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py b/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py index 78a2d294d..39cc26dc4 100644 --- a/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py +++ b/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py @@ -114,26 +114,14 @@ class Upsample3D(Upsample2D): hidden_states = [hidden_states] # ADD BY NUMZ for i in range(len(hidden_states)): - if self.use_conv and hasattr(self, "upscale_conv") and self.upscale_conv.kernel_size == (1, 1, 1): - hidden_states[i] = hidden_states[i].repeat_interleave(self.temporal_ratio, dim=2) - if self.spatial_ratio != 1: - hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=3) - hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=4) - elif self.use_conv: - hidden_states[i] = self.upscale_conv(hidden_states[i]) - hidden_states[i] = rearrange( - hidden_states[i], - "b (x y z c) f h w -> b c (f z) (h x) (w y)", - x=self.spatial_ratio, - y=self.spatial_ratio, - z=self.temporal_ratio, - ).contiguous() - else: - if self.temporal_ratio != 1: - hidden_states[i] = hidden_states[i].repeat_interleave(self.temporal_ratio, dim=2) - if self.spatial_ratio != 1: - hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=3) - hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=4) + hidden_states[i] = self.upscale_conv(hidden_states[i]) + hidden_states[i] = rearrange( + hidden_states[i], + "b (x y z c) f h w -> b c (f z) (h x) (w y)", + x=self.spatial_ratio, + y=self.spatial_ratio, + z=self.temporal_ratio, + ) # [Overridden] For causal temporal conv if self.temporal_up and memory_state != MemoryState.ACTIVE: @@ -1165,10 +1153,8 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): else: encoded = self._encode(x) posterior = DiagonalGaussianDistribution(encoded) - if not return_dict: return (posterior,) - return AutoencoderKLOutput(latent_dist=posterior) @apply_forward_hook @@ -1181,16 +1167,15 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): decoded = self.tiled_decode(z) else: decoded = self._decode(z) - if not return_dict: return (decoded,) - return DecoderOutput(sample=decoded) def _encode( self, x: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED ) -> torch.Tensor: - _x = causal_conv_slice_inputs(x.to(self.device), self.slicing_sample_min_size, memory_state=memory_state) + _x = x.to(self.device) + _x = causal_conv_slice_inputs(_x, self.slicing_sample_min_size, memory_state=memory_state) h = self.encoder(_x, memory_state=memory_state) if self.quant_conv is not None: output = self.quant_conv(h, memory_state=memory_state) @@ -1262,109 +1247,53 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) row_limit = self.tile_latent_min_size - blend_extent - prev_row = None + rows = [] self.tiles = 0 - - row_positions = list(range(0, x.shape[3], overlap_size)) - col_positions = list(range(0, x.shape[4], overlap_size)) - enc = None - output_width = 0 - h_cursor = 0 - - for _row_idx, i in enumerate(row_positions): - row_tiles = [] - for tile_idx, j in enumerate(col_positions): + for i in range(0, x.shape[3], overlap_size): + row = [] + for j in range(0, x.shape[4], overlap_size): tile = x[:, :, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] tile = self._encode(tile) - if tile.ndim == 4: - tile = tile.unsqueeze(0) - if prev_row is not None: - tile = self.blend_v(prev_row[tile_idx], tile, blend_extent) - if tile_idx > 0: - tile = self.blend_h(row_tiles[-1], tile, blend_extent) - row_tiles.append(tile) + row.append(tile) self.tiles += 1 - - cropped_tiles = [tile[:, :, :, :row_limit, :row_limit] for tile in row_tiles] - row_width = 0 - for cropped in cropped_tiles: - row_width += cropped.shape[-1] - if output_width < cropped.shape[-1]: - output_width = cropped.shape[-1] - - if enc is None: - enc = torch.empty( - cropped_tiles[0].shape[0], - cropped_tiles[0].shape[1], - cropped_tiles[0].shape[2], - len(row_positions) * row_limit, - len(col_positions) * row_limit, - dtype=cropped_tiles[0].dtype, - device=cropped_tiles[0].device, - ) - - w_cursor = 0 - for cropped in cropped_tiles: - enc[:, :, :, h_cursor : h_cursor + cropped.shape[-2], w_cursor : w_cursor + cropped.shape[-1]] = cropped - w_cursor += cropped.shape[-1] - - h_cursor += cropped_tiles[0].shape[-2] - prev_row = row_tiles - - return enc[:, :, :, :h_cursor, :w_cursor] + rows.append(row) + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + if i > 0: + tile = self.blend_v(rows[i - 1][j], tile, blend_extent) + if j > 0: + tile = self.blend_h(row[j - 1], tile, blend_extent) + result_row.append(tile[:, :, :, :row_limit, :row_limit]) + result_rows.append(torch.cat(result_row, dim=4)) + enc = torch.cat(result_rows, dim=3) + return enc def tiled_decode(self, z: torch.Tensor) -> torch.Tensor: overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor)) blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor) row_limit = self.tile_sample_min_size - blend_extent - prev_row = None - - row_positions = list(range(0, z.shape[3], overlap_size)) - col_positions = list(range(0, z.shape[4], overlap_size)) - dec = None - output_width = 0 - h_cursor = 0 - - for _row_idx, i in enumerate(row_positions): - row_tiles = [] - for tile_idx, j in enumerate(col_positions): + rows = [] + for i in range(0, z.shape[3], overlap_size): + row = [] + for j in range(0, z.shape[4], overlap_size): tile = z[:, :, :, i : i + self.tile_latent_min_size, j : j + self.tile_latent_min_size] decoded = self.decoder(tile) - if decoded.ndim == 4: - decoded = decoded.unsqueeze(0) - if prev_row is not None: - decoded = self.blend_v(prev_row[tile_idx], decoded, blend_extent) - if tile_idx > 0: - decoded = self.blend_h(row_tiles[-1], decoded, blend_extent) - row_tiles.append(decoded) - - cropped_tiles = [tile[:, :, :, :row_limit, :row_limit] for tile in row_tiles] - row_width = 0 - for cropped in cropped_tiles: - row_width += cropped.shape[-1] - if output_width < cropped.shape[-1]: - output_width = cropped.shape[-1] - - if dec is None: - dec = torch.empty( - cropped_tiles[0].shape[0], - cropped_tiles[0].shape[1], - cropped_tiles[0].shape[2], - len(row_positions) * row_limit, - len(col_positions) * row_limit, - dtype=cropped_tiles[0].dtype, - device=cropped_tiles[0].device, - ) - - w_cursor = 0 - for cropped in cropped_tiles: - dec[:, :, :, h_cursor : h_cursor + cropped.shape[-2], w_cursor : w_cursor + cropped.shape[-1]] = cropped - w_cursor += cropped.shape[-1] - - h_cursor += cropped_tiles[0].shape[-2] - prev_row = row_tiles - - return dec[:, :, :, :h_cursor, :w_cursor] + row.append(decoded) + rows.append(row) + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + if i > 0: + tile = self.blend_v(rows[i - 1][j], tile, blend_extent) + if j > 0: + tile = self.blend_h(row[j - 1], tile, blend_extent) + result_row.append(tile[:, :, :, :row_limit, :row_limit]) + result_rows.append(torch.cat(result_row, dim=4)) + dec = torch.cat(result_rows, dim=3) + return dec def forward( self, x: torch.FloatTensor, mode: Literal["encode", "decode", "all"] = "all", **kwargs @@ -1405,6 +1334,7 @@ class VideoAutoencoderKLWrapper(VideoAutoencoderKL): ): self.spatial_downsample_factor = spatial_downsample_factor self.temporal_downsample_factor = temporal_downsample_factor + self.freeze_encoder = freeze_encoder self.freeze_encoder = True super().__init__(*args, **kwargs) diff --git a/scripts/postprocessing_seedvr.py b/scripts/postprocessing_seedvr.py index 2796a3d46..13dd0a8e7 100644 --- a/scripts/postprocessing_seedvr.py +++ b/scripts/postprocessing_seedvr.py @@ -33,7 +33,7 @@ class ScriptSeedVR(scripts_postprocessing.ScriptPostprocessing): seedvr_tile_overlap = gr.Slider(minimum=0, maximum=1.0, step=0.01, value=0.25, label="SeedVR tile overlap", elem_id="extras_seedvr_tile_overlap") with gr.Row(): seedvr_vae_memory = gr.Slider(minimum=0.1, maximum=1.0, step=0.01, value=1.0, label="SeedVR VAE memory", elem_id="extras_seedvr_vae_memory") - with gr.Accordion('SeedVR video', open = False, elem_id="postprocess_seedvr_video_accordion"): + with gr.Accordion('SeedVR video', open = False, elem_id="postprocess_seedvr_video_accordion", visible=False): with gr.Row(): seedvr_batch_size = gr.Slider(minimum=1, maximum=64, step=1, value=1, label="SeedVR batch size", elem_id="extras_seedvr_batch_size") seedvr_batch_overlap = gr.Slider(minimum=0, maximum=16, step=1, value=0, label="SeedVR batch overlap", elem_id="extras_seedvr_batch_overlap")