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:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user