Fix axis-based WAN DTC gating
This commit is contained in:
+17
-1
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user