diff --git a/__init__.py b/__init__.py index c38b50c..705ead0 100644 --- a/__init__.py +++ b/__init__.py @@ -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"] \ No newline at end of file +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"] diff --git a/src/core/model_configuration.py b/src/core/model_configuration.py index 6129762..f3a92c3 100644 --- a/src/core/model_configuration.py +++ b/src/core/model_configuration.py @@ -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) diff --git a/src/core/model_loader.py b/src/core/model_loader.py index 8a0e1c6..705ee83 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/attn_video_vae.py b/src/models/video_vae_v3/modules/attn_video_vae.py index 90df8d0..dbe452e 100644 --- a/src/models/video_vae_v3/modules/attn_video_vae.py +++ b/src/models/video_vae_v3/modules/attn_video_vae.py @@ -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(): 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 0a1b189..5071a5a 100644 --- a/src/models/video_vae_v3/modules/causal_inflation_lib.py +++ b/src/models/video_vae_v3/modules/causal_inflation_lib.py @@ -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): 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