diff --git a/README.md b/README.md index 6dfc1e8..f4da50a 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ ComfyUI-TIDE is a ComfyUI custom node implementation of inference-time mechanisms from **TIDE: Text-Informed Dynamic Extrapolation with Step-Aware Temperature Control for Diffusion Transformers**. -The primary implementation targets Flux-style DiT attention in ComfyUI. The repository also includes an experimental SDXL/UNet adaptation that applies the usable attention-temperature part of the method to SDXL-style `SpatialTransformer` attention. +The primary implementation targets Flux-style DiT attention in ComfyUI. The repository also includes a WAN 2.1/2.2 path for ComfyUI WAN video DiTs and an experimental SDXL/UNet adaptation that applies the usable attention-temperature part of the method to SDXL-style `SpatialTransformer` attention. The nodes patch a cloned ComfyUI `MODEL` object. They do **not** add extra sampling steps, replace the sampler, replace the scheduler, or fork ComfyUI core. @@ -42,6 +42,12 @@ For Flux-style joint text/image attention, it implements: * a lightweight ComfyUI model wrapper to pass the current denoising timestep into the attention patch; * a small PyTorch SDPA fallback used only when an additive TIDE attention mask is active. +For WAN 2.1/2.2-style video DiT attention, it implements: + +* step-aware RoPE temperature scaling for WAN self-attention; +* lazy wrapping of ComfyUI WAN `rope_encode` / `forward_orig` paths through a cloned-model diffusion wrapper; +* chaining through ComfyUI `WrappersMP.DIFFUSION_MODEL` without modifying ComfyUI core. + For SDXL-style UNet attention, it implements: * step-aware attention-temperature scaling through ComfyUI's `optimized_attention_override`; @@ -64,6 +70,17 @@ The main target path is FLUX-family DiT models in ComfyUI, including FLUX.2-styl Other DiT models may require model-specific patch paths. They should not be assumed to work unless their ComfyUI implementation exposes the same attention-patch contract. +### WAN 2.1 / 2.2 video DiT path + +The WAN node targets ComfyUI WAN-family implementations based on `comfy.ldm.wan.model.WanModel` and close subclasses, including WAN 2.1 and WAN 2.2 paths that expose: + +* `rope_encode`; +* `forward_orig`; +* `rope_embedder.axes_dim`; +* ComfyUI diffusion-model wrapper execution. + +The WAN path applies Dynamic Temperature Control to the WAN self-attention RoPE matrices. It does **not** apply TIDE Text Anchoring because ComfyUI WAN uses separate self-attention over video tokens and cross-attention over text/context tokens. In that architecture, text keys and image/video keys are not competing inside one joint softmax, and adding the same positive bias to every text key in pure cross-attention would be cancelled by softmax shift invariance. + ### SDXL / UNet path The SDXL node targets ComfyUI's UNet `SpatialTransformer` attention path: @@ -88,6 +105,18 @@ Implemented mechanisms: * Dynamic Temperature Control; * optional PyTorch SDPA fallback when the additive text-anchor mask is active. +### TIDE WAN High-Resolution Extrapolation + +Use this node for WAN 2.1 / WAN 2.2-style video DiT models in ComfyUI. + +Implemented mechanism: + +* Dynamic Temperature Control on WAN self-attention RoPE. + +Not implemented for WAN: + +* Text Anchoring, because WAN text conditioning is separate cross-attention rather than Flux-style joint text/image attention. + ### TIDE SDXL High-Resolution Extrapolation Use this node for SDXL-style UNet models. @@ -151,7 +180,13 @@ Defaults: `frequency_mode=official_raw` uses raw RoPE frequencies, matching the released implementation behavior this port was written against. `paper_normalized` is exposed for comparison because the paper notation describes a normalized frequency variable. -### 3. Dynamic attention temperature for SDXL/UNet attention +### 3. Dynamic Temperature Control for WAN 2.1/2.2 RoPE attention + +The WAN node applies the same step-aware RoPE temperature multiplier used by the Flux path, but at ComfyUI WAN's `rope_encode` / `forward_orig` boundary. WAN uses a three-axis RoPE layout `(time, height, width)`, discovered from `rope_embedder.axes_dim` at runtime. Spatial scaling uses the node's `width`, `height`, `base_width`, and `base_height`; the temporal axis is left unscaled. + +Because WAN does not expose a Flux-style joint text/image attention softmax, the WAN node sets `text_anchor_strength=0.0` internally and only applies Dynamic Temperature Control. + +### 4. Dynamic attention temperature for SDXL/UNet attention The SDXL node applies the attention-temperature part of the method by scaling the attention query tensor before ComfyUI's optimized attention function: @@ -206,6 +241,15 @@ SDXL defaults: 6. No global monkey-patching is used. 7. No sampler or scheduler rewrite is performed. +### WAN path + +1. The WAN node clones and patches the incoming ComfyUI `MODEL`. +2. It installs a ComfyUI `WrappersMP.DIFFUSION_MODEL` wrapper on the cloned model. +3. The wrapper injects current timestep metadata into `transformer_options["tide"]`. +4. The wrapper lazily wraps the live WAN model's `rope_encode` and `forward_orig` methods. +5. `rope_encode` output or externally supplied `freqs` are multiplied by the TIDE per-frequency temperature scale once per forward path. +6. The implementation preserves ComfyUI's native WAN block loop, sampler, scheduler, attention backend, and dynamic-VRAM lifetime. + ### SDXL path 1. The SDXL node clones and patches the incoming ComfyUI `MODEL`. @@ -229,6 +273,7 @@ SDXL defaults: | Dynamic Temperature Control | `dyheating()` / temperature-aware RoPE scaling | `tide_core.math.rope_temperature_scale`, applied to ComfyUI `pe` | | Denoising-step-aware behavior | Update position embedding state from current timestep | `TIDEModelWrapper` injects normalized timestep into `transformer_options` | | FLUX attention integration | Modified Diffusers FLUX transformer/processor | ComfyUI `attn1_patch` plus optional `optimized_attention_override` | +| WAN attention integration | Not part of the paper's main Flux/MM-DiT implementation | `tide_core.wan`, ComfyUI diffusion-model wrapper, RoPE scaling for WAN self-attention | | SDXL attention integration | Not part of the paper's main Flux/MM-DiT implementation | `nodes_sdxl.py`, ComfyUI `optimized_attention_override`, query scaling for UNet attention | | Logarithmic FLUX scheduler shift | Scheduler/pipeline-level change | Not implemented by this node | | DyPE / NTK-by-parts / YaRN positional interpolation | Custom positional interpolation stack | Not fully implemented; this node implements the TIDE attention-side mechanisms only | @@ -278,6 +323,34 @@ Ablation settings: | Dynamic Temperature only | `text_anchor_strength=0.0`, `temperature_strength=1.0` | | Disabled | `text_anchor_strength=0.0`, `temperature_strength=0.0` | +### WAN 2.1 / 2.2 usage + +1. Load a WAN 2.1 or WAN 2.2 model as usual. +2. Add **TIDE WAN High-Resolution Extrapolation** after the model loader. +3. Connect the patched `model` output to your sampler. +4. Set `width` and `height` to the final video frame dimensions in pixels. +5. Set `base_width` and `base_height` to the resolution you want to treat as the model's native/reference resolution for this workflow. The node defaults to `640x640` because ComfyUI's WAN 2.2 text-to-video blueprint currently uses that size, but WAN checkpoints and workflows vary. + +Recommended starting values: + +| Setting | Value | +| ---------------------------- | ---------------------------: | +| `width` / `height` | final video frame dimensions | +| `base_width` / `base_height` | workflow/model reference | +| `temperature_strength` | `1.0` | +| `alpha_low` | `0.6` | +| `alpha_high` | `0.2` | +| `tau_max` | `1.0` | +| `frequency_mode` | `official_raw` | + +WAN ablation settings: + +| Test | Settings | +| ----------------------- | -------------------------------- | +| Dynamic Temperature | `temperature_strength=1.0` | +| Disabled | `temperature_strength=0.0` | +| Force native-size patch | `apply_to_native_or_smaller=True` | + ### SDXL usage 1. Load an SDXL checkpoint as usual. @@ -335,6 +408,21 @@ SDXL ablation settings: | `preserve_existing_wrapper` | `True` | Delegate to an existing ComfyUI model wrapper after injecting TIDE metadata. | | `debug` | `False` | Log skipped DTC shape mismatches and exceptions. | +### TIDE WAN High-Resolution Extrapolation + +| Input | Default | Description | +| ---------------------------- | -------------: | ---------------------------------------------------------------------------------------------------- | +| `model` | required | ComfyUI `MODEL` object to patch. | +| `width`, `height` | `1280`, `720` | Final target video frame dimensions in pixels. Must match the latent/video size used by the workflow. | +| `temperature_strength` | `1.0` | Strength of WAN RoPE Dynamic Temperature Control. `0.0` disables the WAN patch. | +| `base_width`, `base_height` | `640`, `640` | Reference/native resolution used for adaptive scaling. Adjust for the checkpoint/workflow. | +| `alpha_low`, `alpha_high` | `0.6`, `0.2` | DTC exponents for low/high RoPE frequency behavior. | +| `tau_max` | `1.0` | Maximum temperature reached near the end of denoising. | +| `frequency_mode` | `official_raw` | `official_raw` or `paper_normalized`. | +| `apply_to_native_or_smaller` | `False` | Allow patching even when target token count is not above base token count. | +| `preserve_existing_wrapper` | `True` | Delegate to an existing ComfyUI model wrapper after injecting TIDE metadata. | +| `debug` | `False` | Log skipped WAN wrapping/shape cases. | + ### TIDE SDXL High-Resolution Extrapolation | Input | Default | Description | @@ -362,10 +450,12 @@ ComfyUI-TIDE/ │ ├── __init__.py │ ├── config.py │ ├── math.py -│ └── patches.py +│ ├── patches.py +│ └── wan.py └── tests/ ├── test_attention_patch.py - └── test_math.py + ├── test_math.py + └── test_wan.py ``` ## Tests @@ -383,7 +473,8 @@ The tests cover: * YaRN/default temperature formula; * RoPE temperature scale shape and timestep progression; * additive attention-mask creation; -* masked SDPA override behavior. +* masked SDPA override behavior; +* WAN RoPE temperature scaling helper behavior. The tests do not validate visual quality or live ComfyUI execution. @@ -395,6 +486,10 @@ python -m py_compile nodes_sdxl.py ## Paper vs implementation differences +### WAN support is an adaptation, not full Flux/MM-DiT TIDE + +The full Text Anchoring mechanism is defined for MM-DiT joint attention where text keys and image keys compete inside one softmax. ComfyUI WAN uses self-attention for video/image tokens and separate cross-attention for text/context tokens. Therefore this repository applies the TIDE Dynamic Temperature Control mechanism to WAN self-attention RoPE, but does not apply Text Anchoring to WAN cross-attention. + ### SDXL support is an adaptation, not full TIDE The full TIDE method is designed around DiT/MM-DiT attention where text tokens and image tokens are present in the same attention sequence. That makes Text Anchoring meaningful because a positive bias on text-key logits changes the balance between text keys and image keys. @@ -432,6 +527,15 @@ The method is architecture-relevant to DiTs, but this implementation is tied to * The timestep passed through the model wrapper is normalized or sigma-like in `[0, 1]`; values outside the interval are clamped. * FLUX-family image token granularity is 16 pixels per transformer token. +### WAN path + +* The model uses a ComfyUI WAN implementation with `rope_encode` and `forward_orig`. +* The node `width` and `height` match the actual generated video frame dimensions. +* `base_width` and `base_height` are chosen as the intended native/reference dimensions for the specific WAN checkpoint/workflow. +* WAN RoPE axes are ordered as `(time, height, width)`, matching ComfyUI's `WanModel.rope_encode`. +* The temporal RoPE axis is not scaled by this node. +* The timestep passed through the model wrapper is normalized or sigma-like in `[0, 1]`; values outside the interval are clamped. + ### SDXL path * The model uses ComfyUI's UNet `SpatialTransformer` attention path. @@ -449,6 +553,8 @@ The method is architecture-relevant to DiTs, but this implementation is tied to * Does not modify the sampler's high-resolution time-shift schedule. * Does not fully implement NTK-by-parts, YaRN positional interpolation, or DyPE positional interpolation. * Flux support is tied to ComfyUI's Flux-style attention patch contract. +* WAN support applies Dynamic Temperature Control only; WAN Text Anchoring is intentionally not implemented. +* WAN visual quality needs live workflow testing across WAN 2.1/2.2 variants, frame counts, resolutions, samplers, and attention backends. * SDXL support is an experimental attention-temperature adaptation, not full Text Anchoring. * SDXL visual quality needs live workflow testing across checkpoints, resolutions, samplers, and attention backends. * Very large resolutions still require sufficient VRAM for the selected model, sampler, attention path, latent size, and VAE path. diff --git a/nodes.py b/nodes.py index 7429a25..6b7e139 100644 --- a/nodes.py +++ b/nodes.py @@ -3,15 +3,15 @@ from __future__ import annotations from typing import Any try: - from .tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper + from .tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, install_tide_wan_patch except ModuleNotFoundError as exc: if exc.name not in {f"{__package__}.tide_core", "tide_core"}: raise - from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper + from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, install_tide_wan_patch except ImportError as exc: if "attempted relative import with no known parent package" not in str(exc): raise - from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper + from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, install_tide_wan_patch class TIDEHighResolutionExtrapolation: @@ -110,10 +110,89 @@ class TIDEHighResolutionExtrapolation: return (patched,) +class TIDEWANHighResolutionExtrapolation: + """Patch a WAN 2.1/2.2-style DiT model with TIDE Dynamic Temperature Control.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "width": ("INT", {"default": 1280, "min": 16, "max": 16384, "step": 16}), + "height": ("INT", {"default": 720, "min": 16, "max": 16384, "step": 16}), + "temperature_strength": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 4.0, "step": 0.05, "tooltip": "0 disables WAN Dynamic Temperature Control; 1.0 uses the TIDE/YaRN curve."}, + ), + }, + "optional": { + "base_width": ("INT", {"default": 640, "min": 16, "max": 16384, "step": 16}), + "base_height": ("INT", {"default": 640, "min": 16, "max": 16384, "step": 16}), + "alpha_low": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 8.0, "step": 0.05}), + "alpha_high": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 8.0, "step": 0.05}), + "tau_max": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 4.0, "step": 0.01}), + "frequency_mode": (["official_raw", "paper_normalized"], {"default": "official_raw"}), + "apply_to_native_or_smaller": ("BOOLEAN", {"default": False}), + "preserve_existing_wrapper": ("BOOLEAN", {"default": True}), + "debug": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("MODEL",) + RETURN_NAMES = ("model",) + FUNCTION = "patch" + CATEGORY = "model_patches/TIDE" + + def patch( + self, + model, + width: int, + height: int, + temperature_strength: float, + base_width: int = 640, + base_height: int = 640, + alpha_low: float = 0.6, + alpha_high: float = 0.2, + tau_max: float = 1.0, + frequency_mode: str = "official_raw", + apply_to_native_or_smaller: bool = False, + preserve_existing_wrapper: bool = True, + debug: bool = False, + ): + config = TIDEConfig( + width=int(width), + height=int(height), + base_width=int(base_width), + base_height=int(base_height), + text_anchor_strength=0.0, + temperature_strength=float(temperature_strength), + alpha_low=float(alpha_low), + alpha_high=float(alpha_high), + tau_max=float(tau_max), + frequency_mode=str(frequency_mode), + apply_to_double_blocks=True, + apply_to_single_blocks=False, + apply_to_native_or_smaller=bool(apply_to_native_or_smaller), + force_pytorch_attention_with_mask=False, + preserve_existing_wrapper=bool(preserve_existing_wrapper), + debug=bool(debug), + ) + + patched = model.clone() + + old_wrapper = patched.model_options.get("model_function_wrapper") + patched.set_model_unet_function_wrapper(TIDEModelWrapper(config, old_wrapper=old_wrapper)) + install_tide_wan_patch(patched, config) + + return (patched,) + + NODE_CLASS_MAPPINGS = { "TIDEHighResolutionExtrapolation": TIDEHighResolutionExtrapolation, + "TIDEWANHighResolutionExtrapolation": TIDEWANHighResolutionExtrapolation, } NODE_DISPLAY_NAME_MAPPINGS = { "TIDEHighResolutionExtrapolation": "TIDE High-Resolution Extrapolation", + "TIDEWANHighResolutionExtrapolation": "TIDE WAN High-Resolution Extrapolation", } diff --git a/tests/test_wan.py b/tests/test_wan.py new file mode 100644 index 0000000..9b0a919 --- /dev/null +++ b/tests/test_wan.py @@ -0,0 +1,49 @@ +import pathlib +import sys + +import torch + +ROOT = pathlib.Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) + +from tide_core.config import TIDEConfig +from tide_core.patches import TIDEAttentionPatch +from tide_core.wan import _scale_wan_freqs + + +class _FakeRopeEmbedder: + axes_dim = (8, 12, 12) + + +class _FakeWanInner: + rope_embedder = _FakeRopeEmbedder() + + +def test_wan_rope_temperature_scale_preserves_shape_and_relaxes_late(): + cfg = TIDEConfig(width=1280, height=640, base_width=640, base_height=640) + freqs = torch.ones(1, 32, 1, sum(_FakeRopeEmbedder.axes_dim) // 2, 2, 2) + + early = _scale_wan_freqs(cfg, _FakeWanInner(), freqs, timestep=1.0) + late = _scale_wan_freqs(cfg, _FakeWanInner(), freqs, timestep=0.0) + + assert early.shape == freqs.shape + assert late.shape == freqs.shape + assert torch.max(early) > 1.0 + assert torch.allclose(late, freqs, atol=1e-6) + + +def test_flux_attn_patch_tolerates_wan_post_attention_patch_payload(): + cfg = TIDEConfig(width=1280, height=720, base_width=640, base_height=640) + patch = TIDEAttentionPatch(cfg) + x = torch.randn(1, 8, 16) + payload = {"x": x, "q": torch.randn(1, 8, 2, 8), "k": torch.randn(1, 8, 2, 8), "transformer_options": {}} + + assert patch(payload) is x + + +def test_flux_attn_patch_only_treats_full_wan_payload_as_passthrough(): + cfg = TIDEConfig(width=1280, height=720, base_width=640, base_height=640) + patch = TIDEAttentionPatch(cfg) + payload = {"x": torch.randn(1, 8, 16)} + + assert patch(payload)["q"] is payload diff --git a/tide_core/__init__.py b/tide_core/__init__.py index 7048e09..958964c 100644 --- a/tide_core/__init__.py +++ b/tide_core/__init__.py @@ -6,6 +6,7 @@ from .math import ( rope_temperature_scale, ) from .patches import TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper +from .wan import TIDEWanDiffusionWrapper, install_tide_wan_patch __all__ = [ "TIDEConfig", @@ -16,4 +17,6 @@ __all__ = [ "TIDEAttentionOverride", "TIDEAttentionPatch", "TIDEModelWrapper", + "TIDEWanDiffusionWrapper", + "install_tide_wan_patch", ] diff --git a/tide_core/patches.py b/tide_core/patches.py index 80807f0..3bf077a 100644 --- a/tide_core/patches.py +++ b/tide_core/patches.py @@ -76,7 +76,15 @@ class TIDEAttentionPatch: def to(self, device: torch.device | str): # Comfy calls .to on patches during model moves. return self - def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None): + def __call__(self, q, k=None, v=None, pe=None, attn_mask=None, extra_options=None): + # ComfyUI WAN currently calls attn1_patch after self-attention with a + # single dict payload: {"x", "q", "k", "transformer_options"}. + # The Flux TIDE patch cannot modify that already-computed attention, so + # leave it as a no-op instead of failing when the generic node is used + # on a WAN model. WAN support is installed through tide_core.wan. + if k is None and v is None and isinstance(q, dict) and {"x", "q", "k", "transformer_options"}.issubset(q): + return q.get("x", q) + extra_options = extra_options or {} if not self.config.should_apply() or not _block_enabled(self.config, extra_options): return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask} diff --git a/tide_core/wan.py b/tide_core/wan.py new file mode 100644 index 0000000..01f0a1b --- /dev/null +++ b/tide_core/wan.py @@ -0,0 +1,308 @@ +from __future__ import annotations + +import logging +from dataclasses import asdict, replace +from typing import Any, Optional + +import torch + +from .config import TIDEConfig +from .math import rope_temperature_scale +from .patches import _safe_timestep01 + +_LOG = logging.getLogger("ComfyUI-TIDE") + +_CONFIG_KEY = "tide_wan_config" +_ENABLED_KEY = "tide_wan_enabled" +_WRAPPER_KEY = "tide_wan_rope_temperature" +_SCALED_FREQS_ID_KEY = "_tide_wan_scaled_freqs_id" +_DIFFUSION_MODEL_WRAPPER_TYPE = "diffusion_model" + +try: # pragma: no cover - ComfyUI is not importable in standalone tests. + import comfy.patcher_extension as _comfy_patcher_extension +except ModuleNotFoundError as exc: # pragma: no cover + if exc.name and exc.name.startswith("comfy"): + _comfy_patcher_extension = None + else: + raise + + +def _diffusion_wrapper_type() -> str: + if _comfy_patcher_extension is None: + return _DIFFUSION_MODEL_WRAPPER_TYPE + return _comfy_patcher_extension.WrappersMP.DIFFUSION_MODEL + + +def _ensure_transformer_options(model: Any) -> dict[str, Any]: + if not hasattr(model, "model_options") or model.model_options is None: + model.model_options = {} + transformer_options = model.model_options.get("transformer_options") + if not isinstance(transformer_options, dict): + transformer_options = {} + model.model_options["transformer_options"] = transformer_options + return transformer_options + + +def _config_dict(config: TIDEConfig) -> dict[str, Any]: + data = asdict(config) + data["axes_dim"] = tuple(config.axes_dim) + data["backend"] = "wan" + data["text_anchoring"] = "not_applicable_cross_attention_only" + return data + + +def _looks_like_wan_inner(inner: Any) -> bool: + # ComfyUI WAN 2.1/2.2 models expose this narrow runtime contract; avoid + # class-name checks so compatible WAN subclasses can use the same path. + return ( + inner is not None + and callable(getattr(inner, "rope_encode", None)) + and callable(getattr(inner, "forward_orig", None)) + and hasattr(inner, "rope_embedder") + ) + + +def _read_wan_axes_dim(inner: Any, fallback: tuple[int, int, int]) -> tuple[int, int, int]: + rope_embedder = getattr(inner, "rope_embedder", None) + axes_dim = getattr(rope_embedder, "axes_dim", None) + if axes_dim is None: + return tuple(int(v) for v in fallback) + try: + axes = tuple(int(v) for v in axes_dim) + except Exception: + return tuple(int(v) for v in fallback) + if len(axes) != 3 or any(v <= 0 for v in axes): + return tuple(int(v) for v in fallback) + return axes + + +def _scale_wan_freqs( + config: TIDEConfig, + inner: Any, + freqs: torch.Tensor, + *, + timestep: float, +) -> torch.Tensor: + """Apply TIDE Dynamic Temperature Control to WAN RoPE matrices. + + ComfyUI WAN computes RoPE as a broadcastable matrix with shape compatible + with [B, tokens, 1, rope_pairs, 2, 2]. The paper's Text Anchoring term is + not applied here: WAN uses separate self-attention and cross-attention, so + text and image keys do not compete in one joint softmax. + """ + + if not torch.is_tensor(freqs): + return freqs + if config.temperature_strength == 0.0 or not config.should_apply(): + return freqs + + axes_dim = _read_wan_axes_dim(inner, tuple(config.axes_dim)) + local_config = replace(config, axes_dim=axes_dim) + + try: + scale = rope_temperature_scale( + local_config, + timestep=timestep, + device=freqs.device, + dtype=freqs.dtype if torch.is_floating_point(freqs) else torch.float32, + ) + except Exception as exc: + if config.debug: + _LOG.exception("TIDE WAN dynamic temperature scale failed and was skipped: %s", exc) + return freqs + + if freqs.ndim < 3 or freqs.shape[-3] != scale.numel(): + if config.debug: + _LOG.warning( + "TIDE WAN skipped dynamic temperature: freqs axis dimension %s != scale length %s", + freqs.shape[-3] if freqs.ndim >= 3 else None, + scale.numel(), + ) + return freqs + + view_shape = (1,) * (freqs.ndim - 3) + (scale.numel(), 1, 1) + return freqs * scale.reshape(view_shape) + + +def _mark_scaled_freqs(transformer_options: Optional[dict[str, Any]], freqs: Any) -> None: + if isinstance(transformer_options, dict) and torch.is_tensor(freqs): + transformer_options[_SCALED_FREQS_ID_KEY] = id(freqs) + + +def _freqs_already_scaled(transformer_options: Optional[dict[str, Any]], freqs: Any) -> bool: + return isinstance(transformer_options, dict) and torch.is_tensor(freqs) and transformer_options.get(_SCALED_FREQS_ID_KEY) == id(freqs) + + +def _resolve_config(transformer_options: Optional[dict[str, Any]], fallback: Optional[TIDEConfig]) -> Optional[TIDEConfig]: + cfg = transformer_options.get(_CONFIG_KEY) if isinstance(transformer_options, dict) else None + if isinstance(cfg, TIDEConfig): + return cfg + return fallback + + +def _resolve_timestep(transformer_options: Optional[dict[str, Any]], timestep: Any = None) -> float: + if isinstance(transformer_options, dict): + tide_opts = transformer_options.get("tide", {}) + if isinstance(tide_opts, dict): + value = tide_opts.get("timestep", None) + if value is not None: + return _safe_timestep01(value) + value = transformer_options.get("timestep", None) + if value is not None: + return _safe_timestep01(value) + return _safe_timestep01(timestep) + + +def _ensure_wan_rope_encode_wrapped(inner: Any, fallback_config: Optional[TIDEConfig]) -> bool: + if not _looks_like_wan_inner(inner): + return False + if getattr(inner, "_tide_wan_rope_encode_wrapped", False): + return True + + original_rope_encode = inner.rope_encode + + def tide_wan_rope_encode(*args: Any, **kwargs: Any): + transformer_options = kwargs.get("transformer_options", None) + freqs = original_rope_encode(*args, **kwargs) + config = _resolve_config(transformer_options, fallback_config) + if config is None: + return freqs + timestep = _resolve_timestep(transformer_options) + out = _scale_wan_freqs(config, inner, freqs, timestep=timestep) + _mark_scaled_freqs(transformer_options, out) + return out + + inner._tide_wan_original_rope_encode = original_rope_encode + inner.rope_encode = tide_wan_rope_encode + inner._tide_wan_rope_encode_wrapped = True + return True + + +def _ensure_wan_forward_orig_wrapped(inner: Any, fallback_config: Optional[TIDEConfig]) -> bool: + if not _looks_like_wan_inner(inner): + return False + if getattr(inner, "_tide_wan_forward_orig_wrapped", False): + return True + + original_forward_orig = inner.forward_orig + + def tide_wan_forward_orig( + x, + t, + context, + clip_fea=None, + freqs=None, + transformer_options=None, + **kwargs, + ): + config = _resolve_config(transformer_options, fallback_config) + if config is not None and torch.is_tensor(freqs) and not _freqs_already_scaled(transformer_options, freqs): + timestep = _resolve_timestep(transformer_options, t) + freqs = _scale_wan_freqs(config, inner, freqs, timestep=timestep) + _mark_scaled_freqs(transformer_options, freqs) + + if transformer_options is None: + return original_forward_orig(x, t, context, clip_fea=clip_fea, freqs=freqs, **kwargs) + return original_forward_orig( + x, + t, + context, + clip_fea=clip_fea, + freqs=freqs, + transformer_options=transformer_options, + **kwargs, + ) + + inner._tide_wan_original_forward_orig = original_forward_orig + inner.forward_orig = tide_wan_forward_orig + inner._tide_wan_forward_orig_wrapped = True + return True + + +class TIDEWanDiffusionWrapper: + """ComfyUI diffusion_model wrapper that enables WAN RoPE DTC hooks lazily.""" + + def __init__(self, config: TIDEConfig): + self.config = config + + def to(self, device: torch.device | str): + return self + + def __call__( + self, + executor, + x, + timestep, + context, + clip_fea=None, + time_dim_concat=None, + transformer_options=None, + **kwargs, + ): + if transformer_options is None: + transformer_options = {} + transformer_options[_CONFIG_KEY] = self.config + transformer_options[_ENABLED_KEY] = self.config.should_apply() and self.config.temperature_strength != 0.0 + + tide_opts = transformer_options.get("tide", {}) + if not isinstance(tide_opts, dict): + tide_opts = {} + tide_opts["timestep"] = _safe_timestep01(timestep) + tide_opts["width"] = self.config.width + tide_opts["height"] = self.config.height + transformer_options["tide"] = tide_opts + + inner = getattr(executor, "class_obj", None) + wrapped_rope = _ensure_wan_rope_encode_wrapped(inner, self.config) + wrapped_forward_orig = _ensure_wan_forward_orig_wrapped(inner, self.config) + if self.config.debug and not (wrapped_rope or wrapped_forward_orig): + _LOG.warning("TIDE WAN wrapper did not find a WAN-like rope_encode/forward_orig target on %s", type(inner).__name__ if inner is not None else None) + + return executor( + x, + timestep, + context, + clip_fea, + time_dim_concat, + transformer_options, + **kwargs, + ) + + +def install_tide_wan_patch(model: Any, config: TIDEConfig) -> Any: + """Install the WAN-specific TIDE DTC path on a cloned ComfyUI MODEL.""" + + transformer_options = _ensure_transformer_options(model) + transformer_options[_CONFIG_KEY] = config + transformer_options[_ENABLED_KEY] = config.should_apply() and config.temperature_strength != 0.0 + transformer_options["tide_wan"] = _config_dict(config) + + wrapper = TIDEWanDiffusionWrapper(config) + wrapper_type = _diffusion_wrapper_type() + + installed_on_patcher = False + if callable(getattr(model, "remove_wrappers_with_key", None)) and callable(getattr(model, "add_wrapper_with_key", None)): + model.remove_wrappers_with_key(wrapper_type, _WRAPPER_KEY) + model.add_wrapper_with_key(wrapper_type, _WRAPPER_KEY, wrapper) + installed_on_patcher = True + + if not installed_on_patcher: + wrappers = transformer_options.setdefault("wrappers", {}) + wrappers_for_type = wrappers.setdefault(wrapper_type, {}) + wrappers_for_type[_WRAPPER_KEY] = [wrapper] + + # Best-effort eager wrapping for already-materialized WAN models. The runtime + # diffusion wrapper repeats this lazily because ComfyUI can replace/live-wrap + # the diffusion module during dynamic loading. + outer = getattr(model, "model", None) + inner = getattr(outer, "diffusion_model", None) if outer is not None else getattr(model, "diffusion_model", None) + _ensure_wan_rope_encode_wrapped(inner, config) + _ensure_wan_forward_orig_wrapped(inner, config) + return model + + +__all__ = [ + "TIDEWanDiffusionWrapper", + "install_tide_wan_patch", + "_scale_wan_freqs", +]