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:
@@ -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.
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
Reference in New Issue
Block a user