Optimize VAE and improve GGUF support for 50-series GPUs

- Implemented 2D convolution optimization in `InflatedCausalConv3d` to speed up spatial-only operations by reshaping effectively 2D tensors and using `F.conv2d` instead of `Conv3d`.
- Added `torch.nn.functional.conv2d` and `torch.nn.functional.conv3d` to `GGUFTensor`'s `__torch_function__` dispatch to enable automatic dequantization of weights, allowing the VAE to utilize GGUF quantization.
- Fixed a bug in `VideoAutoencoderKL` where `slicing_latent_min_size` could become 0 with small split sizes, now clamping it to a minimum of 1.
This commit is contained in:
google-labs-jules[bot]
2025-12-15 15:38:14 +00:00
parent 4fc3296c81
commit 04475bdc24
3 changed files with 47 additions and 2 deletions
+15
View File
@@ -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}:
@@ -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
+2 -2
View File
@@ -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)