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
This commit is contained in:
Adrien Toupet
2025-12-03 12:43:35 -05:00
parent 71ac9ffe54
commit ed53581359
5 changed files with 34 additions and 51 deletions
+1
View File
@@ -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"]
@@ -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,
-48
View File
@@ -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],
)
+1 -1
View File
@@ -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,
+31
View File
@@ -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