diff --git a/src/core/model_loader.py b/src/core/model_loader.py index 9b84a0b..895cab5 100644 --- a/src/core/model_loader.py +++ b/src/core/model_loader.py @@ -378,6 +378,21 @@ class GGUFTensor(torch.Tensor): debug.log(f"Args: {[arg.shape if hasattr(arg, 'shape') else type(arg) for arg in args]}", level="WARNING", category="dit", force=True, indent_level=1) raise + # Handle conv2d/conv3d operations specially + if func in {torch.nn.functional.conv2d, torch.nn.functional.conv3d}: + if len(args) >= 2 and isinstance(args[1], cls): # weight is the second argument + try: + weight_tensor = args[1] + dequantized_weight = weight_tensor.dequantize(device=args[0].device, dtype=args[0].dtype) + new_args = (args[0], dequantized_weight) + args[2:] + return func(*new_args, **kwargs) + except Exception as e: + if debug: + debug.log(f"Error in conv dequantization: {e}", level="WARNING", category="dit", force=True) + debug.log(f"Function: {func}", level="WARNING", category="dit", force=True, indent_level=1) + debug.log(f"Args: {[arg.shape if hasattr(arg, 'shape') else type(arg) for arg in args]}", level="WARNING", category="dit", force=True, indent_level=1) + raise + # Handle matrix multiplication operations that need dequantization if func in {torch.matmul, torch.mm, torch.bmm, torch.addmm, torch.addmv, torch.addr, torch.baddbmm, torch.chain_matmul}: diff --git a/src/models/video_vae_v3/modules/causal_inflation_lib.py b/src/models/video_vae_v3/modules/causal_inflation_lib.py index 664f8b9..64a347b 100644 --- a/src/models/video_vae_v3/modules/causal_inflation_lib.py +++ b/src/models/video_vae_v3/modules/causal_inflation_lib.py @@ -91,6 +91,36 @@ class InflatedCausalConv3d(Conv3d): Workaround: Call torch.cudnn_convolution directly to bypass buggy layer. Status is logged at startup in compatibility.py. """ + # Optimization: Use fast 2D conv for spatial-only operations (no temporal mixing) + # Check: kernel_time=1, stride_time=1, dilation_time=1, padding_time=0 + if ( + self.kernel_size[0] == 1 + and self.stride[0] == 1 + and self.dilation[0] == 1 + and self.padding[0] == 0 + ): + # Reshape input: (B, C, T, H, W) -> (B*T, C, H, W) + B, C, T, H, W = input.shape + input_2d = input.permute(0, 2, 1, 3, 4).reshape(B * T, C, H, W) + + # Reshape weight: (Out, In, 1, kH, kW) -> (Out, In, kH, kW) + weight_2d = weight.squeeze(2) + + # Conv2D params + stride_2d = self.stride[1:] + padding_2d = self.padding[1:] + dilation_2d = self.dilation[1:] + + # Execute standard Conv2d + out_2d = F.conv2d( + input_2d, weight_2d, bias, stride_2d, padding_2d, dilation_2d, self.groups + ) + + # Reshape output back: (B*T, Out, H', W') -> (B, Out, T, H', W') + _, C_out, H_out, W_out = out_2d.shape + out = out_2d.view(B, T, C_out, H_out, W_out).permute(0, 2, 1, 3, 4) + return out + if (NVIDIA_CONV3D_MEMORY_BUG_WORKAROUND and weight.dtype in (torch.float16, torch.bfloat16) and hasattr(torch.backends.cudnn, 'is_available') and diff --git a/src/models/video_vae_v3/modules/video_vae.py b/src/models/video_vae_v3/modules/video_vae.py index daa1fc0..66a1896 100644 --- a/src/models/video_vae_v3/modules/video_vae.py +++ b/src/models/video_vae_v3/modules/video_vae.py @@ -886,7 +886,7 @@ class VideoAutoencoderKL(nn.Module): if split_size is not None: self.enable_slicing() self.slicing_sample_min_size = split_size - self.slicing_latent_min_size = split_size // self.temporal_downsample_factor + self.slicing_latent_min_size = max(1, split_size // self.temporal_downsample_factor) else: self.disable_slicing() for module in self.modules(): @@ -950,7 +950,7 @@ class VideoAutoencoderKLWrapper(VideoAutoencoderKL): self.disable_slicing() self.slicing_sample_min_size = split_size if split_size is not None: - self.slicing_latent_min_size = split_size // self.temporal_downsample_factor + self.slicing_latent_min_size = max(1, split_size // self.temporal_downsample_factor) for module in self.modules(): if isinstance(module, InflatedCausalConv3d): module.set_memory_device(memory_device) \ No newline at end of file