Fix axis-based WAN DTC gating

This commit is contained in:
xmarre
2026-05-12 20:29:36 +02:00
parent 3d15074771
commit a54e64c72c
4 changed files with 112 additions and 31 deletions
+17 -1
View File
@@ -67,5 +67,21 @@ class TIDEConfig:
def is_extrapolating(self) -> bool:
return self.target_image_tokens > self.base_image_tokens
def should_apply(self) -> bool:
@property
def is_axis_extrapolating(self) -> bool:
return self.scale_x > 1.0 or self.scale_y > 1.0
def should_apply_text_anchor(self) -> bool:
# Text anchoring is derived from total text/image token dilution, so keep
# the area-based gate unless the user explicitly forces native/smaller use.
return self.apply_to_native_or_smaller or self.is_extrapolating
def should_apply_temperature(self) -> bool:
# Dynamic Temperature Control is RoPE-axis based. Wide or tall WAN/Flux
# generations can extrapolate on one spatial axis while total pixel count
# is <= the square base area, e.g. 832x480 vs 640x640. In that case the
# width axis still needs the DTC path.
return self.apply_to_native_or_smaller or self.is_axis_extrapolating
def should_apply(self) -> bool:
return self.should_apply_text_anchor() or self.should_apply_temperature()
+1 -1
View File
@@ -34,7 +34,7 @@ def adaptive_text_bias(config: TIDEConfig) -> float:
paper value linearly; values <= base resolution return zero by default.
"""
if not config.should_apply():
if not config.should_apply_text_anchor():
return 0.0
beta = math.log(config.target_pixel_ratio)
if beta < 0.0 and not config.apply_to_native_or_smaller:
+24 -23
View File
@@ -86,36 +86,37 @@ class TIDEAttentionPatch:
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}
img_slice = extra_options.get("img_slice")
if not img_slice or len(img_slice) != 2:
return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask}
try:
text_tokens = int(img_slice[0])
total_tokens = int(k.shape[2])
except Exception:
return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask}
if text_tokens <= 0 or total_tokens <= text_tokens:
if not _block_enabled(self.config, extra_options):
return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask}
out_mask = attn_mask
beta = adaptive_text_bias(self.config)
if beta != 0.0:
out_mask = _add_text_bias_mask(
attn_mask,
text_tokens=text_tokens,
key_tokens=total_tokens,
beta=beta,
device=k.device,
dtype=q.dtype if torch.is_floating_point(q) else torch.float32,
)
img_slice = extra_options.get("img_slice")
if img_slice and len(img_slice) == 2:
try:
text_tokens = int(img_slice[0])
total_tokens = int(k.shape[2])
except Exception:
text_tokens = 0
total_tokens = 0
if text_tokens > 0 and total_tokens > text_tokens:
out_mask = _add_text_bias_mask(
attn_mask,
text_tokens=text_tokens,
key_tokens=total_tokens,
beta=beta,
device=k.device,
dtype=q.dtype if torch.is_floating_point(q) else torch.float32,
)
elif self.config.debug:
_LOG.warning("TIDE skipped text anchoring because Flux text/image token slices were unavailable.")
elif self.config.debug:
_LOG.warning("TIDE skipped text anchoring because Flux img_slice metadata was unavailable.")
out_pe = pe
if pe is not None and self.config.temperature_strength != 0.0:
if pe is not None and self.config.temperature_strength != 0.0 and self.config.should_apply_temperature():
tide_opts = extra_options.get("tide", {})
timestep = _safe_timestep01(tide_opts.get("timestep", extra_options.get("timestep")))
try:
+70 -6
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import logging
import sys
from dataclasses import asdict, replace
from typing import Any, Optional
@@ -17,6 +18,12 @@ _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"
_DEBUG_SCALE_LOG_LIMIT = 12
def _debug_print(config: Optional[TIDEConfig], message: str) -> None:
if config is not None and config.debug:
print(message, file=sys.stderr, flush=True)
try: # pragma: no cover - ComfyUI is not importable in standalone tests.
import comfy.patcher_extension as _comfy_patcher_extension
@@ -93,7 +100,20 @@ def _scale_wan_freqs(
if not torch.is_tensor(freqs):
return freqs
if config.temperature_strength == 0.0 or not config.should_apply():
if config.temperature_strength == 0.0:
if config.debug:
_debug_print(config, "[ComfyUI-TIDE] WAN DTC disabled: temperature_strength=0")
return freqs
if not config.should_apply_temperature():
if config.debug:
_debug_print(
config,
"[ComfyUI-TIDE] WAN DTC inactive: "
f"width={config.width} height={config.height} "
f"base={config.base_width}x{config.base_height} "
f"scale_x={config.scale_x:.4f} scale_y={config.scale_y:.4f} "
"and apply_to_native_or_smaller=false",
)
return freqs
axes_dim = _read_wan_axes_dim(inner, tuple(config.axes_dim))
@@ -118,10 +138,31 @@ def _scale_wan_freqs(
freqs.shape[-3] if freqs.ndim >= 3 else None,
scale.numel(),
)
_debug_print(
config,
"[ComfyUI-TIDE] WAN DTC skipped: "
f"freqs_shape={tuple(freqs.shape)} scale_len={int(scale.numel())}",
)
return freqs
view_shape = (1,) * (freqs.ndim - 3) + (scale.numel(), 1, 1)
return freqs * scale.reshape(view_shape)
out = freqs * scale.reshape(view_shape)
if config.debug:
count = int(getattr(inner, "_tide_wan_scale_log_count", 0))
if count < _DEBUG_SCALE_LOG_LIMIT:
setattr(inner, "_tide_wan_scale_log_count", count + 1)
_debug_print(
config,
"[ComfyUI-TIDE] WAN DTC applied: "
f"step_log={count + 1}/{_DEBUG_SCALE_LOG_LIMIT} "
f"timestep={float(timestep):.6f} "
f"shape={tuple(freqs.shape)} axes_dim={axes_dim} "
f"scale_x={config.scale_x:.4f} scale_y={config.scale_y:.4f} "
f"scale_min={float(scale.min().detach().cpu()):.6f} "
f"scale_max={float(scale.max().detach().cpu()):.6f}",
)
return out
def _mark_scaled_freqs(transformer_options: Optional[dict[str, Any]], freqs: Any) -> None:
@@ -175,6 +216,7 @@ def _ensure_wan_rope_encode_wrapped(inner: Any, fallback_config: Optional[TIDECo
inner._tide_wan_original_rope_encode = original_rope_encode
inner.rope_encode = tide_wan_rope_encode
inner._tide_wan_rope_encode_wrapped = True
_debug_print(fallback_config, f"[ComfyUI-TIDE] WAN wrapped rope_encode on {type(inner).__name__} id={id(inner)}")
return True
@@ -216,6 +258,7 @@ def _ensure_wan_forward_orig_wrapped(inner: Any, fallback_config: Optional[TIDEC
inner._tide_wan_original_forward_orig = original_forward_orig
inner.forward_orig = tide_wan_forward_orig
inner._tide_wan_forward_orig_wrapped = True
_debug_print(fallback_config, f"[ComfyUI-TIDE] WAN wrapped forward_orig on {type(inner).__name__} id={id(inner)}")
return True
@@ -242,7 +285,7 @@ class TIDEWanDiffusionWrapper:
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
transformer_options[_ENABLED_KEY] = self.config.should_apply_temperature() and self.config.temperature_strength != 0.0
tide_opts = transformer_options.get("tide", {})
if not isinstance(tide_opts, dict):
@@ -255,8 +298,19 @@ class TIDEWanDiffusionWrapper:
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)
if self.config.debug:
if not getattr(self, "_logged_runtime", False):
_debug_print(
self.config,
"[ComfyUI-TIDE] WAN diffusion wrapper active: "
f"inner={type(inner).__name__ if inner is not None else None} "
f"wrapped_rope={bool(wrapped_rope)} wrapped_forward_orig={bool(wrapped_forward_orig)} "
f"enabled={bool(transformer_options[_ENABLED_KEY])}",
)
self._logged_runtime = True
if 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)
_debug_print(self.config, f"[ComfyUI-TIDE] WAN wrapper did not find WAN target: inner={type(inner).__name__ if inner is not None else None}")
return executor(
x,
@@ -274,7 +328,7 @@ def install_tide_wan_patch(model: Any, config: TIDEConfig) -> Any:
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[_ENABLED_KEY] = config.should_apply_temperature() and config.temperature_strength != 0.0
transformer_options["tide_wan"] = _config_dict(config)
wrapper = TIDEWanDiffusionWrapper(config)
@@ -291,6 +345,16 @@ def install_tide_wan_patch(model: Any, config: TIDEConfig) -> Any:
wrappers_for_type = wrappers.setdefault(wrapper_type, {})
wrappers_for_type[_WRAPPER_KEY] = [wrapper]
_debug_print(
config,
"[ComfyUI-TIDE] WAN patch installed: "
f"wrapper_type={wrapper_type} model_patcher_wrapper={installed_on_patcher} "
f"enabled={bool(transformer_options[_ENABLED_KEY])} "
f"width={config.width} height={config.height} base={config.base_width}x{config.base_height} "
f"scale_x={config.scale_x:.4f} scale_y={config.scale_y:.4f} "
f"temperature_strength={config.temperature_strength:.4f}",
)
# 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.