Merge pull request #3 from xmarre/codex/wan21-wan22-support

[codex] Add WAN 2.1 and 2.2 TIDE support
This commit is contained in:
xmarre
2026-05-12 19:23:26 +02:00
committed by GitHub
6 changed files with 562 additions and 9 deletions
+111 -5
View File
@@ -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.
+82 -3
View File
@@ -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",
}
+49
View File
@@ -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
+3
View File
@@ -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",
]
+9 -1
View File
@@ -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}
+308
View File
@@ -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",
]