diff --git a/tide_core/config.py b/tide_core/config.py index 2ce04f1..c9b7a67 100644 --- a/tide_core/config.py +++ b/tide_core/config.py @@ -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() diff --git a/tide_core/math.py b/tide_core/math.py index d98fd59..c25c60a 100644 --- a/tide_core/math.py +++ b/tide_core/math.py @@ -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: diff --git a/tide_core/patches.py b/tide_core/patches.py index 3bf077a..159f9ee 100644 --- a/tide_core/patches.py +++ b/tide_core/patches.py @@ -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: diff --git a/tide_core/wan.py b/tide_core/wan.py index 01f0a1b..14a7540 100644 --- a/tide_core/wan.py +++ b/tide_core/wan.py @@ -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.