Merge PR #421 from naxci1 - Optimize VAE and GGUF

This commit is contained in:
Adrien Toupet
2025-12-24 01:13:19 +01:00
6 changed files with 228 additions and 45 deletions
+19 -1
View File
@@ -5,5 +5,23 @@ Official SeedVR2 integration for ComfyUI
from .src.optimization.compatibility import ensure_triton_compat # noqa: F401
from .src.interfaces import comfy_entrypoint, SeedVR2Extension
from .src.interfaces.video_upscaler import SeedVR2VideoUpscaler
from .src.interfaces.dit_model_loader import SeedVR2LoadDiTModel
from .src.interfaces.vae_model_loader import SeedVR2LoadVAEModel
from .src.interfaces.torch_compile_settings import SeedVR2TorchCompileSettings
__all__ = ["comfy_entrypoint", "SeedVR2Extension"]
NODE_CLASS_MAPPINGS = {
"SeedVR2VideoUpscaler": SeedVR2VideoUpscaler,
"SeedVR2LoadDiTModel": SeedVR2LoadDiTModel,
"SeedVR2LoadVAEModel": SeedVR2LoadVAEModel,
"SeedVR2TorchCompileSettings": SeedVR2TorchCompileSettings
}
NODE_DISPLAY_NAME_MAPPINGS = {
"SeedVR2VideoUpscaler": "SeedVR2 Video Upscaler",
"SeedVR2LoadDiTModel": "SeedVR2 (Down)Load DiT Model",
"SeedVR2LoadVAEModel": "SeedVR2 (Down)Load VAE Model",
"SeedVR2TorchCompileSettings": "SeedVR2 Torch Compile Settings"
}
__all__ = ["comfy_entrypoint", "SeedVR2Extension", "NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+18 -15
View File
@@ -1317,23 +1317,26 @@ def _configure_torch_compile(compile_args: Dict[str, Any], model_type: str,
# Check Triton availability for inductor backend BEFORE attempting compilation
if settings['backend'] == 'inductor':
if not TRITON_AVAILABLE:
error_msg = (
warn_msg = (
f"Cannot use torch.compile with 'inductor' backend: Triton is not installed.\n"
f"\n"
f"Triton is required for the inductor backend which performs kernel fusion and optimization.\n"
f"\n"
f"To fix this issue:\n"
f" 1. Install Triton: pip install triton\n"
f" 2. OR change backend to 'cudagraphs' (lightweight, no Triton needed)\n"
f" 3. OR disable torch.compile\n"
f"\n"
f"For more info: https://github.com/triton-lang/triton"
)
debug.log(error_msg, level="ERROR", category="setup", force=True)
raise RuntimeError(
"torch.compile with inductor backend requires Triton. "
"Install with: pip install triton"
f"Automatically falling back to 'cudagraphs' backend if available, or disabling compilation.\n"
f"For best performance, install Triton: pip install triton"
)
debug.log(warn_msg, level="WARNING", category="setup", force=True)
# Check if cudagraphs is available as fallback
try:
import torch
# Simple check - if we can't use inductor, try to see if we can use cudagraphs
# But typically cudagraphs is built-in.
# Let's switch to cudagraphs and see.
settings['backend'] = 'cudagraphs'
debug.log("Switched backend to 'cudagraphs' due to missing Triton", category="setup", force=True)
except Exception:
debug.log("Could not switch to cudagraphs, disabling compilation", level="WARNING", category="setup")
raise RuntimeError("Triton missing and fallback failed") # Should be caught by caller if we wanted to
# If we switched, we don't raise RuntimeError anymore.
# Log compilation configuration
debug.log(f"Configuring torch.compile for {model_type}...", category="setup", force=True)
+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}:
@@ -313,14 +313,7 @@ class ResnetBlock3D(ResnetBlock2D):
):
hidden_states = input_tensor
hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
hidden_states = retry_on_oom(
self.nonlinearity,
hidden_states,
debug=getattr(self, 'debug', None),
operation_name="ResnetBlock3D.nonlinearity"
)
# Handle upsample/downsample first (usually involves 5D operations if temporal)
if self.upsample is not None:
# upsample_nearest_nhwc fails with large batch sizes.
# see https://github.com/huggingface/diffusers/issues/984
@@ -333,6 +326,78 @@ class ResnetBlock3D(ResnetBlock2D):
input_tensor = self.downsample(input_tensor, memory_state=memory_state)
hidden_states = self.downsample(hidden_states, memory_state=memory_state)
# Check if we can run the main ResNet path in 4D (flattened temporal)
# Conditions:
# 1. Input is 5D (B, C, T, H, W)
# 2. conv1 and conv2 are effectively 2D (spatial only)
# 3. conv_shortcut (if present) is effectively 2D
# 4. time_embedding_norm is "default" (simple add) - others might need more complex broadcasting logic in 4D
can_run_4d = (
hidden_states.ndim == 5
and self.conv1.check_effective_2d()
and self.conv2.check_effective_2d()
and (self.conv_shortcut is None or self.conv_shortcut.check_effective_2d())
and (self.time_embedding_norm == "default" or temb is None)
)
if can_run_4d:
B, C, T, H, W = hidden_states.shape
# Reshape to 4D: (B*T, C, H, W)
# Use contiguous to ensure efficient memory layout for conv2d
hidden_states = hidden_states.transpose(1, 2).reshape(B * T, C, H, W).contiguous()
if input_tensor.ndim == 5:
input_tensor = input_tensor.transpose(1, 2).reshape(B * T, C, H, W).contiguous()
# Prepare temb for 4D
temb_4d = None
if self.time_emb_proj is not None and temb is not None:
if not self.skip_time_act:
temb = self.nonlinearity(temb)
# temb: (B, C_emb) -> proj -> (B, C_out) -> (B, C_out, 1, 1)
# We need (B*T, C_out, 1, 1)
temb_proj = self.time_emb_proj(temb)[:, :, None, None]
# Expand B to B*T
temb_4d = temb_proj.repeat_interleave(T, dim=0)
# Main path in 4D
# causal_norm_wrapper handles 4D input by just calling GroupNorm (no rearrange overhead!)
hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
hidden_states = self.nonlinearity(hidden_states)
# InflatedCausalConv3d handles 4D input by dispatching to conv2d (cached check)
hidden_states = self.conv1(hidden_states, memory_state=memory_state)
if temb_4d is not None and self.time_embedding_norm == "default":
hidden_states = hidden_states + temb_4d
hidden_states = causal_norm_wrapper(self.norm2, hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states, memory_state=memory_state)
if self.conv_shortcut is not None:
input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
# Reshape back to 5D: (B, C, T, H, W)
# Output is (B*T, C, H, W)
_, C_out, H_out, W_out = output_tensor.shape
output_tensor = output_tensor.view(B, T, C_out, H_out, W_out).transpose(1, 2)
return output_tensor
# Fallback to standard 5D path
hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
hidden_states = retry_on_oom(
self.nonlinearity,
hidden_states,
debug=getattr(self, 'debug', None),
operation_name="ResnetBlock3D.nonlinearity"
)
hidden_states = self.conv1(hidden_states, memory_state=memory_state)
if self.time_emb_proj is not None:
@@ -1078,7 +1143,7 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
norm_num_groups: int = 32,
sample_size: int = 32,
scaling_factor: float = 0.18215,
force_upcast: float = True,
force_upcast: float = False,
attention: bool = True,
temporal_scale_num: int = 2,
slicing_up_num: int = 0,
@@ -1093,7 +1158,7 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
):
extra_cond_dim = kwargs.pop("extra_cond_dim") if "extra_cond_dim" in kwargs else None
self.slicing_sample_min_size = slicing_sample_min_size
self.slicing_latent_min_size = slicing_sample_min_size // (2**temporal_scale_num)
self.slicing_latent_min_size = max(1, slicing_sample_min_size // (2**temporal_scale_num))
super().__init__(
in_channels=in_channels,
@@ -1718,7 +1783,7 @@ class VideoAutoencoderKLWrapper(VideoAutoencoderKL):
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():
@@ -81,6 +81,63 @@ class InflatedCausalConv3d(Conv3d):
def set_memory_device(self, memory_device: _memory_device_t):
self.memory_device = memory_device
def check_effective_2d(self):
"""
Check if the convolution is effectively 2D (spatial only) AND compatible with 2D optimization path.
Returns True if:
1. Temporal kernel size is 1 OR weights are tail-inflated (only last temporal slice is non-zero).
2. AND stride[0] == 1.
3. AND dilation[0] == 1.
4. AND padding[0] == 0 (self.padding is modified in __init__, this checks the spatial-only padding).
"""
if hasattr(self, '_effective_2d'):
return self._effective_2d
# Check stride, dilation, padding first (fast checks)
if not (self.stride[0] == 1 and self.dilation[0] == 1 and self.padding[0] == 0):
self._effective_2d = False
return False
if self.kernel_size[0] == 1:
self._effective_2d = True
return True
# Check for tail inflation (all zeros except last slice in temporal dim)
# Weight shape: (Out, In, T, H, W)
with torch.no_grad():
# Check if all slices except the last one are zero
# Use a small threshold for float comparison or exact check for initialized weights
if torch.sum(torch.abs(self.weight[:, :, :-1, :, :])) < 1e-6:
self._effective_2d = True
return True
self._effective_2d = False
return False
def forward(
self,
input: Union[Tensor, List[Tensor]],
memory_state: MemoryState = MemoryState.UNSET
) -> Tensor:
assert memory_state != MemoryState.UNSET
if memory_state != MemoryState.ACTIVE:
self.memory = None
# Handle 4D input directly (optimization for spatial-only blocks)
if torch.is_tensor(input) and input.ndim == 4:
# If 4D input is passed, we MUST be in the "effective 2D" path.
# Skip extend_head and memory logic, go straight to conv.
# _conv_forward handles the actual 2D dispatch.
return super().forward(input)
if (
math.isinf(self.memory_limit)
and torch.is_tensor(input)
and get_sequence_parallel_group() is None
):
return self.basic_forward(input, memory_state)
return self.slicing_forward(input, memory_state)
def _conv_forward(self, input, weight, bias, *args, **kwargs):
"""
Override _conv_forward to work around NVIDIA Conv3d memory bug.
@@ -91,6 +148,47 @@ 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 OR tail-inflated weights AND stride/dilation/padding compatibility
can_use_2d_path = self.check_effective_2d()
if can_use_2d_path:
# Handle 4D input directly (B*T folded into batch)
if input.ndim == 4:
# Input: (BT, C, H, W)
input_2d = input
is_4d_input = True
else:
# 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)
is_4d_input = False
# Prepare 2D weights
if self.kernel_size[0] == 1:
weight_2d = weight.squeeze(2)
else:
# Tail inflation: use the last temporal slice
weight_2d = weight[:, :, -1, :, :]
# 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
)
if is_4d_input:
return out_2d
# 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
@@ -210,22 +308,6 @@ class InflatedCausalConv3d(Conv3d):
)
return output
def forward(
self,
input: Union[Tensor, List[Tensor]],
memory_state: MemoryState = MemoryState.UNSET
) -> Tensor:
assert memory_state != MemoryState.UNSET
if memory_state != MemoryState.ACTIVE:
self.memory = None
if (
math.isinf(self.memory_limit)
and torch.is_tensor(input)
and get_sequence_parallel_group() is None
):
return self.basic_forward(input, memory_state)
return self.slicing_forward(input, memory_state)
def basic_forward(self, input: Tensor, memory_state: MemoryState = MemoryState.UNSET):
mem_size = self.stride[0] - self.kernel_size[0]
if (self.memory is not None) and (memory_state == MemoryState.ACTIVE):
+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)