From b93c7515f7a102de8bfd87df2e579e4a79874e6f Mon Sep 17 00:00:00 2001 From: xmarre Date: Sat, 2 May 2026 01:41:43 +0200 Subject: [PATCH] Add SDXL high-resolution TIDE attention override node --- __init__.py | 13 +++ nodes_sdxl.py | 228 ++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 241 insertions(+) create mode 100644 nodes_sdxl.py diff --git a/__init__.py b/__init__.py index e295e25..c081c16 100644 --- a/__init__.py +++ b/__init__.py @@ -15,4 +15,17 @@ except ImportError as exc: # so the relative import path above remains the normal runtime path. from nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +# SDXL/UNet support is intentionally kept in a separate module because the +# faithful TIDE text-anchor path is MM-DiT-specific. This registers the SDXL +# adaptation without changing existing FLUX nodes. +from .nodes_sdxl import NODE_CLASS_MAPPINGS as _SDXL_NODE_CLASS_MAPPINGS +from .nodes_sdxl import NODE_DISPLAY_NAME_MAPPINGS as _SDXL_NODE_DISPLAY_NAME_MAPPINGS + +NODE_CLASS_MAPPINGS.update(_SDXL_NODE_CLASS_MAPPINGS) +try: + NODE_DISPLAY_NAME_MAPPINGS +except NameError: + NODE_DISPLAY_NAME_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS.update(_SDXL_NODE_DISPLAY_NAME_MAPPINGS) + __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/nodes_sdxl.py b/nodes_sdxl.py new file mode 100644 index 0000000..92db716 --- /dev/null +++ b/nodes_sdxl.py @@ -0,0 +1,228 @@ +import math +from dataclasses import dataclass +from typing import Any, Callable, Optional + +import torch + + +@dataclass(frozen=True) +class TIDESDXLConfig: + width: int + height: int + base_width: int + base_height: int + temperature_strength: float + alpha: float + tau_max: float + apply_to: str + + @property + def pixel_ratio(self) -> float: + base_pixels = max(1.0, float(self.base_width) * float(self.base_height)) + target_pixels = max(1.0, float(self.width) * float(self.height)) + return max(1.0, target_pixels / base_pixels) + + @property + def extrapolation_scale(self) -> float: + return math.sqrt(self.pixel_ratio) + + @property + def inv_tau_min(self) -> float: + # TIDE inherits YaRN's attention-temperature convention: + # sqrt(1/tau) = 0.1 * ln(scale) + 1. + scale = self.extrapolation_scale + if scale <= 1.0: + return 1.0 + yarn_mscale = 0.1 * math.log(scale) + 1.0 + return yarn_mscale * yarn_mscale + + +def _token_count(x: torch.Tensor, skip_reshape: bool) -> int: + if skip_reshape and x.ndim >= 4: + return int(x.shape[-2]) + return int(x.shape[1]) + + +def _normalised_sigma_t(transformer_options: dict[str, Any]) -> float: + """Return t in [0, 1], where 1 is the noisy/start side and 0 is the clean/end side.""" + cur = transformer_options.get("sigmas", None) + all_sigmas = transformer_options.get("sample_sigmas", None) + if cur is None or all_sigmas is None: + return 1.0 + + try: + cur_f = float(cur.detach().float().mean().item()) if torch.is_tensor(cur) else float(cur) + if torch.is_tensor(all_sigmas): + sigmas = all_sigmas.detach().float().flatten() + sigmas = sigmas[torch.isfinite(sigmas)] + if sigmas.numel() == 0: + return 1.0 + sigma_max = float(sigmas.max().item()) + positive = sigmas[sigmas > 0] + sigma_min = float(positive.min().item()) if positive.numel() else float(sigmas.min().item()) + else: + vals = [float(v) for v in all_sigmas if math.isfinite(float(v))] + if not vals: + return 1.0 + sigma_max = max(vals) + positives = [v for v in vals if v > 0] + sigma_min = min(positives) if positives else min(vals) + + denom = sigma_max - sigma_min + if denom <= 1e-12: + return 1.0 + return max(0.0, min(1.0, (cur_f - sigma_min) / denom)) + except Exception: + return 1.0 + + +def _temperature_q_scale(config: TIDESDXLConfig, transformer_options: dict[str, Any]) -> float: + if config.temperature_strength <= 0.0: + return 1.0 + + inv_tau_min = config.inv_tau_min + if inv_tau_min <= 1.0: + return 1.0 + + tau_min = 1.0 / inv_tau_min + tau_max = max(tau_min, float(config.tau_max)) + t = _normalised_sigma_t(transformer_options) + alpha = max(1e-6, float(config.alpha)) + tau = tau_max - (tau_max - tau_min) * (t ** alpha) + inv_tau = 1.0 / max(tau, 1e-6) + + # Strength blends from no-op to the full TIDE/YaRN temperature. + return 1.0 + (inv_tau - 1.0) * max(0.0, min(1.0, float(config.temperature_strength))) + + +def _looks_like_unet_spatial_transformer(transformer_options: dict[str, Any]) -> bool: + # SDXL/SD1 UNet SpatialTransformer sets activations_shape before calling + # BasicTransformerBlock. FLUX-style DiT paths instead use block_type/img_slice. + if "activations_shape" not in transformer_options: + return False + if "block_type" in transformer_options or "img_slice" in transformer_options: + return False + return True + + +def _attention_kind(q_tokens: int, k_tokens: int) -> str: + # In SDXL UNet blocks, self-attention has image-token K/V (q_len == k_len), + # while cross-attention has text-token K/V (usually 77 tokens). + if q_tokens == k_tokens: + return "self" + return "cross" + + +def _enabled_for_kind(apply_to: str, kind: str) -> bool: + return apply_to == "both" or apply_to == kind + + +def build_sdxl_attention_override( + config: TIDESDXLConfig, + previous_override: Optional[Callable[..., torch.Tensor]] = None, +) -> Callable[..., torch.Tensor]: + def tide_sdxl_attention_override(func: Callable[..., torch.Tensor], *args: Any, **kwargs: Any) -> torch.Tensor: + if len(args) < 4: + if previous_override is not None: + return previous_override(func, *args, **kwargs) + return func(*args, **kwargs) + + q, k, v, heads = args[:4] + rest = args[4:] + transformer_options = kwargs.get("transformer_options", {}) or {} + + if ( + not torch.is_tensor(q) + or not torch.is_tensor(k) + or not _looks_like_unet_spatial_transformer(transformer_options) + ): + if previous_override is not None: + return previous_override(func, *args, **kwargs) + return func(*args, **kwargs) + + skip_reshape = bool(kwargs.get("skip_reshape", False)) + q_tokens = _token_count(q, skip_reshape) + k_tokens = _token_count(k, skip_reshape) + kind = _attention_kind(q_tokens, k_tokens) + + if _enabled_for_kind(config.apply_to, kind): + q_scale = _temperature_q_scale(config, transformer_options) + if q_scale != 1.0: + q = q * q_scale + args = (q, k, v, heads, *rest) + + if previous_override is not None: + return previous_override(func, *args, **kwargs) + return func(*args, **kwargs) + + return tide_sdxl_attention_override + + +class TIDESDXLHighRes: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "width": ("INT", {"default": 1536, "min": 64, "max": 16384, "step": 8}), + "height": ("INT", {"default": 1536, "min": 64, "max": 16384, "step": 8}), + "temperature_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.05}), + "base_width": ("INT", {"default": 1024, "min": 64, "max": 16384, "step": 8}), + "base_height": ("INT", {"default": 1024, "min": 64, "max": 16384, "step": 8}), + "alpha": ("FLOAT", {"default": 0.6, "min": 0.01, "max": 4.0, "step": 0.05}), + "tau_max": ("FLOAT", {"default": 1.0, "min": 0.05, "max": 4.0, "step": 0.05}), + "apply_to": (["cross", "self", "both"], {"default": "both"}), + } + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "apply" + CATEGORY = "model_patches/TIDE" + + def apply( + self, + model, + width: int, + height: int, + temperature_strength: float, + base_width: int, + base_height: int, + alpha: float, + tau_max: float, + apply_to: str, + ): + patched = model.clone() + config = TIDESDXLConfig( + width=int(width), + height=int(height), + base_width=int(base_width), + base_height=int(base_height), + temperature_strength=float(temperature_strength), + alpha=float(alpha), + tau_max=float(tau_max), + apply_to=str(apply_to), + ) + + transformer_options = patched.model_options.setdefault("transformer_options", {}) + previous_override = transformer_options.get("optimized_attention_override", None) + transformer_options["optimized_attention_override"] = build_sdxl_attention_override(config, previous_override) + transformer_options["tide_sdxl"] = { + "width": config.width, + "height": config.height, + "base_width": config.base_width, + "base_height": config.base_height, + "temperature_strength": config.temperature_strength, + "alpha": config.alpha, + "tau_max": config.tau_max, + "apply_to": config.apply_to, + } + return (patched,) + + +NODE_CLASS_MAPPINGS = { + "TIDESDXLHighRes": TIDESDXLHighRes, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "TIDESDXLHighRes": "TIDE SDXL High-Resolution Extrapolation", +}