From ed5358135990e157a2ff8344b94eaeea852d439a Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 3 Dec 2025 12:43:35 -0500 Subject: [PATCH] Fix triton.ops compatibility for bitsandbytes 0.45+ / triton 3.0+ Fixes #340 - Installation error with PyTorch 2.7+cu126 and triton_windows Add compatibility shim for missing triton.ops.matmul_perf_model module. Reverts local VAE types approach --- __init__.py | 1 + .../video_vae_v3/modules/attn_video_vae.py | 3 +- src/models/video_vae_v3/modules/types.py | 48 ------------------- src/models/video_vae_v3/modules/video_vae.py | 2 +- src/optimization/compatibility.py | 31 ++++++++++++ 5 files changed, 34 insertions(+), 51 deletions(-) diff --git a/__init__.py b/__init__.py index 6f8c296..c38b50c 100644 --- a/__init__.py +++ b/__init__.py @@ -3,6 +3,7 @@ ComfyUI-SeedVR2_VideoUpscaler Official SeedVR2 integration for ComfyUI """ +from .src.optimization.compatibility import ensure_triton_compat # noqa: F401 from .src.interfaces import comfy_entrypoint, SeedVR2Extension __all__ = ["comfy_entrypoint", "SeedVR2Extension"] \ No newline at end of file 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 7713dd7..342543a 100644 --- a/src/models/video_vae_v3/modules/attn_video_vae.py +++ b/src/models/video_vae_v3/modules/attn_video_vae.py @@ -17,6 +17,7 @@ import torch import torch.nn as nn import torch.nn.functional as F from diffusers.models.attention_processor import Attention, SpatialNorm +from diffusers.models.autoencoders.vae import DecoderOutput, DiagonalGaussianDistribution from diffusers.models.downsampling import Downsample2D from diffusers.models.lora import LoRACompatibleConv from diffusers.models.modeling_outputs import AutoencoderKLOutput @@ -45,8 +46,6 @@ from .types import ( CausalAutoencoderOutput, CausalDecoderOutput, CausalEncoderOutput, - DecoderOutput, - DiagonalGaussianDistribution, MemoryState, _inflation_mode_t, _memory_device_t, diff --git a/src/models/video_vae_v3/modules/types.py b/src/models/video_vae_v3/modules/types.py index 9b2160a..5a030d2 100644 --- a/src/models/video_vae_v3/modules/types.py +++ b/src/models/video_vae_v3/modules/types.py @@ -74,51 +74,3 @@ class CausalEncoderOutput(NamedTuple): class CausalDecoderOutput(NamedTuple): sample: torch.Tensor - - -class DecoderOutput: - """Output of decoding method - matches diffusers.models.autoencoders.vae.DecoderOutput""" - def __init__(self, sample: torch.Tensor, commit_loss: Optional[torch.Tensor] = None): - self.sample = sample - self.commit_loss = commit_loss - - -class DiagonalGaussianDistribution: - """Matches diffusers.models.autoencoders.vae.DiagonalGaussianDistribution exactly.""" - def __init__(self, parameters: torch.Tensor, deterministic: bool = False): - self.parameters = parameters - self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) - self.logvar = torch.clamp(self.logvar, -30.0, 20.0) - self.deterministic = deterministic - self.std = torch.exp(0.5 * self.logvar) - self.var = torch.exp(self.logvar) - if self.deterministic: - self.var = self.std = torch.zeros_like( - self.mean, device=self.parameters.device, dtype=self.parameters.dtype - ) - - def sample(self, generator: Optional[torch.Generator] = None) -> torch.Tensor: - if self.deterministic: - return self.mode() - sample = torch.randn( - self.mean.shape, - generator=generator, - device=self.parameters.device, - dtype=self.parameters.dtype, - ) - return self.mean + self.std * sample - - def mode(self) -> torch.Tensor: - return self.mean - - def kl(self, other: Optional["DiagonalGaussianDistribution"] = None) -> torch.Tensor: - if other is None: - return 0.5 * torch.sum( - self.mean.pow(2) + self.var - 1.0 - self.logvar, - dim=[1, 2, 3], - ) - return 0.5 * torch.sum( - (self.mean - other.mean).pow(2) / other.var - + self.var / other.var - 1.0 - self.logvar + other.logvar, - dim=[1, 2, 3], - ) diff --git a/src/models/video_vae_v3/modules/video_vae.py b/src/models/video_vae_v3/modules/video_vae.py index 8077d1e..daa1fc0 100644 --- a/src/models/video_vae_v3/modules/video_vae.py +++ b/src/models/video_vae_v3/modules/video_vae.py @@ -15,6 +15,7 @@ from typing import Optional, Tuple, Literal, Callable, Union import torch import torch.nn as nn import torch.nn.functional as F +from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution from einops import rearrange from ....common.half_precision_fixes import safe_pad_operation @@ -35,7 +36,6 @@ from .types import ( CausalAutoencoderOutput, CausalDecoderOutput, CausalEncoderOutput, - DiagonalGaussianDistribution, MemoryState, _inflation_mode_t, _memory_device_t, diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index 94a18ee..486a516 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -5,6 +5,37 @@ Contains FP8/FP16 compatibility layers and wrappers for different model architec Extracted from: seedvr2.py (lines 1045-1630) """ +# Triton compatibility shim for bitsandbytes 0.45+ with triton 3.0+ +# Must be called before any diffusers import +import sys + +def ensure_triton_compat(): + """Create minimal triton.ops stubs only if missing, to allow bitsandbytes import.""" + if 'triton.ops.matmul_perf_model' in sys.modules: + return + + try: + from triton.ops.matmul_perf_model import early_config_prune # noqa: F401 + return + except (ImportError, ModuleNotFoundError, AttributeError): + pass + + import types + + if 'triton.ops' not in sys.modules: + sys.modules['triton.ops'] = types.ModuleType('triton.ops') + + matmul_perf = types.ModuleType('triton.ops.matmul_perf_model') + matmul_perf.early_config_prune = lambda configs, *a, **kw: configs + matmul_perf.estimate_matmul_time = lambda *a, **kw: 0.0 + + sys.modules['triton.ops'].matmul_perf_model = matmul_perf + sys.modules['triton.ops.matmul_perf_model'] = matmul_perf + +# Run immediately on import +ensure_triton_compat() + + import torch import types import os