Import ComfyUI-TIDE into skeleton repo

This commit is contained in:
xmarre
2026-05-01 15:53:44 +02:00
parent 331f7c67d4
commit 2239bb7e73
13 changed files with 943 additions and 0 deletions
+8
View File
@@ -0,0 +1,8 @@
__pycache__/
*.py[cod]
.pytest_cache/
.venv/
env/
outputs/
.codex/
.vs/
+220
View File
@@ -0,0 +1,220 @@
# ComfyUI-TIDE
ComfyUI-TIDE is a ComfyUI custom node implementation of **TIDE: Text-Informed Dynamic Extrapolation with Step-Aware Temperature Control for Diffusion Transformers**.
The node is intended for FLUX-family DiT models in ComfyUI, including FLUX.2-style model paths that use the same joint text/image attention structure as ComfyUI's FLUX implementation. It patches the `MODEL` object and does not add sampling steps.
## What is implemented
TIDE has two core inference-time mechanisms:
1. **Text Anchoring**
- Adds a positive additive bias to attention logits whose keys are text tokens.
- The default adaptive bias is:
```text
beta = log((target_width * target_height) / (base_width * base_height))
```
- With the default base resolution of 1024x1024, this matches the official FLUX script's `log(width / 1024) + log(height / 1024)`.
2. **Dynamic Temperature Control**
- Applies the paper/official-code YaRN temperature concept as a per-frequency multiplier on ComfyUI's RoPE matrix.
- The default curve follows the released implementation's `dyheating` form:
```text
tau(t, f) = tau_max - (tau_max - tau_min) * t ** alpha(f)
alpha(f) = alpha_low + (alpha_high - alpha_low) * f
```
- Defaults: `tau_max=1.0`, `alpha_low=0.6`, `alpha_high=0.2`.
- `frequency_mode=official_raw` uses raw RoPE frequencies, matching the released code. `paper_normalized` is available for comparison because the paper text describes a normalized frequency variable.
## Implementation plan
1. Patch ComfyUI attention rather than cloning a full model implementation.
2. Use ComfyUI's existing `attn1_patch` hook to modify joint DiT attention inputs.
3. Use a model function wrapper to expose the current denoising timestep to the attention patch.
4. Keep TIDE state local to the cloned `MODEL` object.
5. Avoid global monkey-patching and avoid changing the scheduler or sampler.
6. Force a small PyTorch SDPA attention path only when an additive TIDE mask is present, preventing backends from rejecting or materializing a dense high-resolution mask.
## Paper / official-code / local-module mapping
| Paper component | Official implementation component | This repository |
|---|---|---|
| Text-token influence decay analysis | `run.py::get_attn_mask()` creates an additive text-token mask | `tide_core.math.adaptive_text_bias`, `tide_core.patches.TIDEAttentionPatch` |
| Text Anchoring, beta = log(lambda) | `log(width / 1024) + log(height / 1024)` for the first 512 text tokens | Adaptive beta computed from node `width`, `height`, `base_width`, `base_height`; text-token count is inferred from ComfyUI `img_slice` instead of hard-coded to 512 |
| YaRN temperature baseline | `get_mscale`, `get_default_temperature` | `tide_core.math.get_mscale`, `get_default_temperature` |
| Dynamic Temperature Control | `dyheating()` and `tuning_temperature()` multiply RoPE cos/sin by `1/sqrt(tau)` | `tide_core.math.rope_temperature_scale`, applied to ComfyUI's RoPE matrix in `TIDEAttentionPatch` |
| Denoising-step-aware RoPE | Official `FluxPosEmbed.update_timestep(timestep.item())` | `TIDEModelWrapper` injects normalized current timestep into `transformer_options["tide"]` |
| FLUX attention integration | Official fork modifies Diffusers `FluxTransformer2DModel` / processor | ComfyUI `attn1_patch` and `optimized_attention_override`; no fork of ComfyUI core |
| Logarithmic FLUX scheduler shift | Official `FluxPipeline(..., shift_mode="log")` | Not implemented as a node-level sampler change; see limitations |
| DyPE / NTK-by-parts / YaRN positional interpolation | Official fork implements custom RoPE interpolation | Not fully implemented; this node implements TIDE's attention-side mechanisms and optional frequency interpretation only |
## Installation
Clone or copy this repository into ComfyUI's custom node directory:
```bash
cd ComfyUI/custom_nodes
git clone <this-repo-url> ComfyUI-TIDE
```
Restart ComfyUI.
No extra Python dependency is required beyond the PyTorch/ComfyUI environment. `requirements.txt` lists `torch` only for standalone unit tests.
## Usage
1. Add **TIDE High-Resolution Extrapolation** after your model loader.
2. Connect its `model` output to the sampler model input.
3. Set `width` and `height` to the final generation size used by your latent node.
4. Use a FLUX-family DiT model.
Recommended starting values:
| Setting | Value |
|---|---:|
| `width` / `height` | final latent/image size |
| `base_width` / `base_height` | `1024` / `1024` for FLUX-family models |
| `text_anchor_strength` | `1.0` |
| `temperature_strength` | `1.0` |
| `alpha_low` | `0.6` |
| `alpha_high` | `0.2` |
| `frequency_mode` | `official_raw` |
| `force_pytorch_attention_with_mask` | `True` |
For ablations:
- Text Anchoring only: `temperature_strength=0.0`, `text_anchor_strength=1.0`
- Dynamic Temperature only: `text_anchor_strength=0.0`, `temperature_strength=1.0`
- Disable the node behavior without removing it: set both strengths to `0.0`
## Node inputs
### Required
- `model`: ComfyUI MODEL to patch.
- `width`, `height`: final target generation dimensions in pixels.
- `text_anchor_strength`: multiplier on beta. `1.0` follows paper/official behavior.
- `temperature_strength`: multiplier on the RoPE temperature scale. `1.0` follows official behavior; `0.0` disables DTC.
### Optional
- `base_width`, `base_height`: training/native resolution used for adaptive scaling. Default is 1024x1024.
- `alpha_low`, `alpha_high`: Dynamic Temperature Control exponents.
- `tau_max`: maximum temperature. Default 1.0.
- `frequency_mode`:
- `official_raw`: matches the released code's raw RoPE-frequency use.
- `paper_normalized`: normalizes frequency to [0, 1] as a literal reading of the paper notation.
- `apply_to_double_blocks`, `apply_to_single_blocks`: choose which FLUX block types receive the patch.
- `apply_to_native_or_smaller`: default false. Prevents non-positive or native-resolution anchoring from changing normal-resolution behavior.
- `force_pytorch_attention_with_mask`: default true. Uses an internal PyTorch SDPA path only when the TIDE additive mask is active.
- `preserve_existing_wrapper`: default true. Delegates to an existing model wrapper after injecting TIDE timestep metadata.
- `debug`: logs skipped dynamic-temperature shape mismatches and exceptions.
## Repository structure
```text
ComfyUI-TIDE/
├── __init__.py
├── nodes.py
├── requirements.txt
├── README.md
├── examples/
│ └── README.md
├── tide_core/
│ ├── __init__.py
│ ├── config.py
│ ├── math.py
│ └── patches.py
└── tests/
├── test_attention_patch.py
└── test_math.py
```
## Tests
Standalone math/patch tests:
```bash
cd ComfyUI-TIDE
python -m pip install pytest torch
python -m pytest -q
```
These tests validate:
- adaptive beta computation,
- YaRN temperature formula,
- RoPE scale shape and timestep progression,
- attention-mask creation,
- masked SDPA override behavior.
They do not validate visual quality or live ComfyUI model execution.
## Paper vs official code mismatches and resolutions
### 1. Frequency variable in Dynamic Temperature Control
- Paper: describes `f` as frequency normalized into a range used by `alpha(f)`.
- Official code: multiplies `(alpha_high - alpha_low)` by raw RoPE frequencies.
- Resolution: default `frequency_mode=official_raw` to match official behavior; `paper_normalized` is exposed for controlled comparison.
### 2. Dynamic Temperature implementation site
- Paper: formulates the final attention with a temperature term in the softmax denominator.
- Official FLUX code: implements temperature by multiplying RoPE cos/sin by `1/sqrt(tau)` inside YaRN RoPE generation.
- Resolution: ComfyUI exposes RoPE matrices through `pe`; this repo applies the official-code equivalent multiplier to `pe`.
### 3. Text token count
- Paper: text-token length is abstract `L_T`.
- Official FLUX script: hard-codes 512 text tokens for FLUX.1.
- Resolution: this repo infers text-token count from ComfyUI's `img_slice`, avoiding hard-coding 512 and making FLUX-family variants more likely to work.
### 4. Scheduler time shifting
- Paper appendix and official pipeline use a logarithmic FLUX time-shift schedule for high resolutions.
- This custom node receives an already-built sampler schedule and should not silently alter the user's sampler.
- Resolution: no sampler schedule rewrite is performed. This is a deliberate deviation. Use a ComfyUI sampler/scheduler setup that does not over-shift high-resolution FLUX timesteps.
### 5. Positional interpolation / DyPE / YaRN
- Paper experiments combine TIDE with YaRN/DyPE-style positional handling.
- Official code includes a Diffusers FLUX fork with NTK, NTK-by-parts, YaRN, and DyPE RoPE logic.
- Resolution: this repo implements TIDE's attention-side contribution in ComfyUI and does not clone the full official Diffusers transformer. This avoids replacing ComfyUI internals but means it is not a complete official YaRN/DyPE port.
### 6. Qwen/general DiT support
- Official code includes Qwen-Image support.
- ComfyUI Qwen's current patch path does not propagate a patch-returned additive attention mask to the final attention call in the same way as FLUX.
- Resolution: this repo is implemented for Flux-style ComfyUI joint attention first. Other DiTs may work only if their ComfyUI implementation exposes compatible `attn1_patch`, `img_slice`, and additive-mask propagation.
## Assumptions
- The model uses ComfyUI's Flux-style joint attention with text tokens before image tokens.
- ComfyUI provides `extra_options["img_slice"]` in attention patches.
- `width` and `height` passed to this node match the actual generated latent/image dimensions.
- The timestep seen by the wrapper is already normalized or sigma-like in [0, 1]. Values outside the interval are clamped.
- FLUX-family token granularity is 16 image pixels per transformer token.
## Limitations
- Not a full official repository clone.
- Does not implement official Diffusers pipeline scripts, benchmark code, metric evaluation, datasets, or training code.
- Does not modify the sampler's high-resolution time-shift schedule.
- Does not implement full NTK-by-parts, YaRN positional interpolation, or DyPE positional interpolation.
- Visual quality is unverified in this static repository export.
- Very large resolutions still require enough VRAM for the chosen model, sampler, attention backend, and VAE path.
## Unresolved uncertainties
- Exact FLUX.2 internal token layout may differ from FLUX.1 depending on the ComfyUI model implementation. If it still uses Flux-style `img_slice` with text tokens first, this patch should apply.
- Whether ComfyUI's current default scheduler for FLUX.2 already avoids the extreme high-resolution time-shift problem described in the paper needs live workflow verification.
- The paper notation and official code differ on frequency normalization; defaulting to official code is the safest reproducibility choice, but it is still a documented mismatch.
## License
This repository is an implementation scaffold for ComfyUI. ComfyUI itself is GPL-licensed. Review license compatibility before redistributing as part of a larger package.
+9
View File
@@ -0,0 +1,9 @@
try:
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
except ImportError:
# Allows standalone pytest execution from a directory whose name is not a
# valid Python package identifier. ComfyUI imports this file as a package,
# so the relative import path above remains the normal runtime path.
from nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+10
View File
@@ -0,0 +1,10 @@
# Example workflow wiring
Minimal ComfyUI graph placement:
1. Load a FLUX-family model as usual.
2. Insert **TIDE High-Resolution Extrapolation** immediately after the model loader and before the sampler.
3. Set `width` and `height` to the same final latent/image dimensions used by your empty latent node.
4. Use a high-resolution latent such as 2048x2048, 4096x2048, or another multiple of 16.
The node modifies only the MODEL object. It does not create latents, change the sampler, or decode images.
+113
View File
@@ -0,0 +1,113 @@
from __future__ import annotations
from typing import Any
try:
from .tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper
except ImportError:
from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper
class TIDEHighResolutionExtrapolation:
"""Patch a Flux-style DiT model with TIDE text anchoring and dynamic temperature."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"width": ("INT", {"default": 2048, "min": 16, "max": 16384, "step": 16}),
"height": ("INT", {"default": 2048, "min": 16, "max": 16384, "step": 16}),
"text_anchor_strength": (
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 4.0, "step": 0.05, "tooltip": "Multiplier on paper beta=log(target_pixels/base_pixels). 1.0 matches the paper/official code."},
),
"temperature_strength": (
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 4.0, "step": 0.05, "tooltip": "0 disables Dynamic Temperature Control; 1.0 matches the official dyheating curve."},
),
},
"optional": {
"base_width": ("INT", {"default": 1024, "min": 16, "max": 16384, "step": 16}),
"base_height": ("INT", {"default": 1024, "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_double_blocks": ("BOOLEAN", {"default": True}),
"apply_to_single_blocks": ("BOOLEAN", {"default": True}),
"apply_to_native_or_smaller": ("BOOLEAN", {"default": False}),
"force_pytorch_attention_with_mask": (
"BOOLEAN",
{"default": True, "tooltip": "Use PyTorch SDPA for masked TIDE attention to avoid backends that reject or densify additive masks."},
),
"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,
text_anchor_strength: float,
temperature_strength: float,
base_width: int = 1024,
base_height: int = 1024,
alpha_low: float = 0.6,
alpha_high: float = 0.2,
tau_max: float = 1.0,
frequency_mode: str = "official_raw",
apply_to_double_blocks: bool = True,
apply_to_single_blocks: bool = True,
apply_to_native_or_smaller: bool = False,
force_pytorch_attention_with_mask: bool = True,
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=float(text_anchor_strength),
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=bool(apply_to_double_blocks),
apply_to_single_blocks=bool(apply_to_single_blocks),
apply_to_native_or_smaller=bool(apply_to_native_or_smaller),
force_pytorch_attention_with_mask=bool(force_pytorch_attention_with_mask),
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))
patched.set_model_attn1_patch(TIDEAttentionPatch(config))
transformer_options = patched.model_options.get("transformer_options", {}).copy()
old_override = transformer_options.get("optimized_attention_override")
transformer_options["optimized_attention_override"] = TIDEAttentionOverride(config, old_override=old_override)
patched.model_options["transformer_options"] = transformer_options
return (patched,)
NODE_CLASS_MAPPINGS = {
"TIDEHighResolutionExtrapolation": TIDEHighResolutionExtrapolation,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TIDEHighResolutionExtrapolation": "TIDE High-Resolution Extrapolation",
}
+3
View File
@@ -0,0 +1,3 @@
[tool.pytest.ini_options]
testpaths = ["tests"]
pythonpath = ["."]
+1
View File
@@ -0,0 +1 @@
torch
+57
View File
@@ -0,0 +1,57 @@
import math
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, TIDEAttentionOverride
def test_attention_patch_adds_text_bias_and_scales_pe():
cfg = TIDEConfig(width=2048, height=2048, base_width=1024, base_height=1024)
patch = TIDEAttentionPatch(cfg)
text_tokens = 4
image_tokens = 16
total = text_tokens + image_tokens
q = torch.randn(1, 2, total, 8)
k = torch.randn(1, 2, total, 8)
v = torch.randn(1, 2, total, 8)
pe = torch.ones(1, 1, total, sum(cfg.axes_dim) // 2, 2, 2)
out = patch(q, k, v, pe=pe, attn_mask=None, extra_options={"img_slice": [text_tokens, total], "block_type": "double", "tide": {"timestep": 1.0}})
mask = out["attn_mask"]
assert mask.shape == (1, 1, 1, total)
assert torch.allclose(mask[..., :text_tokens], torch.full((1, 1, 1, text_tokens), math.log(4.0)))
assert torch.allclose(mask[..., text_tokens:], torch.zeros(1, 1, 1, image_tokens))
assert out["pe"].shape == pe.shape
assert torch.max(out["pe"]) > 1.0
def test_attention_patch_respects_block_toggles():
cfg = TIDEConfig(width=2048, height=2048, apply_to_double_blocks=False, apply_to_single_blocks=True)
patch = TIDEAttentionPatch(cfg)
q = torch.randn(1, 1, 8, 4)
k = torch.randn(1, 1, 8, 4)
v = torch.randn(1, 1, 8, 4)
out = patch(q, k, v, pe=None, attn_mask=None, extra_options={"img_slice": [2, 8], "block_type": "double"})
assert out["attn_mask"] is None
def test_attention_override_runs_masked_sdpa_without_comfy_imports():
cfg = TIDEConfig(width=2048, height=2048, force_pytorch_attention_with_mask=True)
override = TIDEAttentionOverride(cfg)
q = torch.randn(1, 2, 5, 8)
k = torch.randn(1, 2, 5, 8)
v = torch.randn(1, 2, 5, 8)
mask = torch.zeros(1, 1, 1, 5)
def should_not_run(*args, **kwargs):
raise AssertionError("delegate should not run when mask is present")
out = override(should_not_run, q, k, v, 2, mask=mask, skip_reshape=True)
assert out.shape == (1, 5, 16)
+45
View File
@@ -0,0 +1,45 @@
import math
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.math import adaptive_text_bias, get_default_temperature, rope_temperature_scale
def test_adaptive_text_bias_matches_paper_and_official_flux_script():
cfg = TIDEConfig(width=2048, height=2048, base_width=1024, base_height=1024)
assert math.isclose(adaptive_text_bias(cfg), math.log(4.0), rel_tol=1e-6)
def test_native_resolution_bias_is_zero_by_default():
cfg = TIDEConfig(width=1024, height=1024, base_width=1024, base_height=1024)
assert adaptive_text_bias(cfg) == 0.0
def test_yarn_temperature_formula():
expected_mscale = 0.1 * math.log(4.0) + 1.0
assert math.isclose(get_default_temperature(4.0), 1.0 / (expected_mscale * expected_mscale), rel_tol=1e-6)
def test_rope_temperature_scale_shape_and_progression():
cfg = TIDEConfig(width=4096, height=2048, base_width=1024, base_height=1024)
early = rope_temperature_scale(cfg, timestep=1.0, device=torch.device("cpu"))
late = rope_temperature_scale(cfg, timestep=0.0, device=torch.device("cpu"))
assert early.shape == (sum(cfg.axes_dim) // 2,)
assert late.shape == (sum(cfg.axes_dim) // 2,)
assert torch.all(torch.isfinite(early))
assert torch.all(torch.isfinite(late))
# At t=0, Eq. 21 reaches tau_max=1, so the multiplier is 1.
assert torch.allclose(late, torch.ones_like(late), atol=1e-6)
assert torch.max(early) > 1.0
def test_temperature_strength_zero_disables_scaling():
cfg = TIDEConfig(width=4096, height=4096, temperature_strength=0.0)
scale = rope_temperature_scale(cfg, timestep=1.0, device=torch.device("cpu"))
assert torch.allclose(scale, torch.ones_like(scale), atol=1e-6)
+19
View File
@@ -0,0 +1,19 @@
from .config import TIDEConfig
from .math import (
adaptive_text_bias,
get_default_temperature,
get_mscale,
rope_temperature_scale,
)
from .patches import TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper
__all__ = [
"TIDEConfig",
"adaptive_text_bias",
"get_default_temperature",
"get_mscale",
"rope_temperature_scale",
"TIDEAttentionOverride",
"TIDEAttentionPatch",
"TIDEModelWrapper",
]
+61
View File
@@ -0,0 +1,61 @@
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True)
class TIDEConfig:
"""Runtime configuration for the ComfyUI TIDE patch.
The defaults intentionally match the paper/official-code path where feasible:
base 1024x1024 FLUX training resolution, alpha_low=0.6, alpha_high=0.2,
tau_max=1.0, and Flux-style latent patch granularity of 16 image pixels per
transformer token.
"""
width: int
height: int
base_width: int = 1024
base_height: int = 1024
token_px: int = 16
text_anchor_strength: float = 1.0
temperature_strength: float = 1.0
alpha_low: float = 0.6
alpha_high: float = 0.2
tau_max: float = 1.0
theta: float = 10000.0
axes_dim: tuple[int, int, int] = (16, 56, 56)
frequency_mode: str = "official_raw"
apply_to_double_blocks: bool = True
apply_to_single_blocks: bool = True
apply_to_native_or_smaller: bool = False
force_pytorch_attention_with_mask: bool = True
preserve_existing_wrapper: bool = True
debug: bool = False
@property
def target_image_tokens(self) -> int:
return max(1, (int(self.width) // self.token_px) * (int(self.height) // self.token_px))
@property
def base_image_tokens(self) -> int:
return max(1, (int(self.base_width) // self.token_px) * (int(self.base_height) // self.token_px))
@property
def target_pixel_ratio(self) -> float:
return max(1.0e-12, (float(self.width) * float(self.height)) / (float(self.base_width) * float(self.base_height)))
@property
def scale_x(self) -> float:
return max(1.0e-12, float(self.width) / float(self.base_width))
@property
def scale_y(self) -> float:
return max(1.0e-12, float(self.height) / float(self.base_height))
@property
def is_extrapolating(self) -> bool:
return self.target_image_tokens > self.base_image_tokens
def should_apply(self) -> bool:
return self.apply_to_native_or_smaller or self.is_extrapolating
+172
View File
@@ -0,0 +1,172 @@
from __future__ import annotations
import math
from typing import Iterable
import torch
from .config import TIDEConfig
def get_mscale(scale: float) -> float:
"""YaRN mscale used by the official TIDE implementation.
Official code computes 0.1 * log(scale) + 1 for scale > 1 and 1 otherwise.
The corresponding attention temperature is 1 / mscale**2.
"""
scale = float(scale)
if scale <= 1.0:
return 1.0
return 0.1 * math.log(scale) + 1.0
def get_default_temperature(scale: float) -> float:
mscale = get_mscale(scale)
return 1.0 / (mscale * mscale)
def adaptive_text_bias(config: TIDEConfig) -> float:
"""Paper Eq. 17/18: beta = log(lambda), with lambda = pixel ratio.
The official FLUX script computes log(width / 1024) + log(height / 1024),
which is identical to log((width * height) / 1024**2). Strength scales the
paper value linearly; values <= base resolution return zero by default.
"""
if not config.should_apply():
return 0.0
beta = math.log(config.target_pixel_ratio)
if beta < 0.0 and not config.apply_to_native_or_smaller:
beta = 0.0
return float(config.text_anchor_strength) * beta
def _axis_frequencies(axis_dim: int, theta: float, device: torch.device) -> torch.Tensor:
if axis_dim <= 0:
return torch.empty(0, dtype=torch.float32, device=device)
if axis_dim % 2 != 0:
raise ValueError(f"axis_dim must be even, got {axis_dim}")
return 1.0 / (float(theta) ** (torch.arange(0, axis_dim, 2, dtype=torch.float32, device=device) / axis_dim))
def _frequency_parameter(freqs: torch.Tensor, mode: str) -> torch.Tensor:
if freqs.numel() == 0:
return freqs
if mode == "official_raw":
# Matches the released code: alpha = alpha_low + (alpha_high-alpha_low) * freqs.
return freqs
if mode == "paper_normalized":
# The paper describes a normalized frequency f, but the released FLUX code uses raw RoPE frequencies.
# This option provides the closest literal normalized interpretation for comparison.
lo = freqs.min()
hi = freqs.max()
if torch.isclose(hi, lo):
return torch.zeros_like(freqs)
return (freqs - lo) / (hi - lo)
raise ValueError(f"Unsupported frequency_mode: {mode!r}")
def axis_temperature_scale(
*,
axis_dim: int,
axis_scale: float,
timestep: float,
alpha_low: float,
alpha_high: float,
tau_max: float,
theta: float,
frequency_mode: str,
strength: float,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
"""Return per-RoPE-pair multiplier equivalent to 1 / sqrt(tau(t, f))."""
freqs = _axis_frequencies(axis_dim, theta, device)
if freqs.numel() == 0:
return freqs.to(dtype=dtype)
low_temperature = get_default_temperature(axis_scale)
t = float(max(0.0, min(1.0, timestep)))
tau_max = float(tau_max)
if tau_max <= 0.0:
raise ValueError(f"tau_max must be > 0, got {tau_max}")
if axis_scale <= 1.0 or strength <= 0.0:
target = torch.ones_like(freqs)
else:
f = _frequency_parameter(freqs, frequency_mode)
alphas = float(alpha_low) + (float(alpha_high) - float(alpha_low)) * f
# Paper Eq. 21. Official implementation names this dyheating.
temps = tau_max - ((tau_max - float(low_temperature)) * torch.pow(torch.tensor(t, device=device), alphas))
temps = temps.clamp_min(1.0e-6)
target = torch.rsqrt(temps)
if strength < 1.0:
target = 1.0 + float(strength) * (target - 1.0)
elif strength > 1.0:
target = 1.0 + float(strength) * (target - 1.0)
return target.to(dtype=dtype)
def rope_temperature_scale(
config: TIDEConfig,
*,
timestep: float,
device: torch.device,
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""Build a concatenated scale vector for ComfyUI's complex RoPE matrix.
ComfyUI's FLUX EmbedND concatenates RoPE axes in the order (t, y, x), and
each RoPE pair is represented by one 2x2 complex-rotation matrix. The output
length is therefore sum(axes_dim) // 2, matching pe.shape[-3].
"""
axes_dim = tuple(int(v) for v in config.axes_dim)
if len(axes_dim) != 3:
raise ValueError(f"Expected three axes dims (t, y, x), got {axes_dim!r}")
parts = [
axis_temperature_scale(
axis_dim=axes_dim[0],
axis_scale=1.0,
timestep=timestep,
alpha_low=config.alpha_low,
alpha_high=config.alpha_high,
tau_max=config.tau_max,
theta=config.theta,
frequency_mode=config.frequency_mode,
strength=config.temperature_strength,
device=device,
dtype=dtype,
),
axis_temperature_scale(
axis_dim=axes_dim[1],
axis_scale=config.scale_y,
timestep=timestep,
alpha_low=config.alpha_low,
alpha_high=config.alpha_high,
tau_max=config.tau_max,
theta=config.theta,
frequency_mode=config.frequency_mode,
strength=config.temperature_strength,
device=device,
dtype=dtype,
),
axis_temperature_scale(
axis_dim=axes_dim[2],
axis_scale=config.scale_x,
timestep=timestep,
alpha_low=config.alpha_low,
alpha_high=config.alpha_high,
tau_max=config.tau_max,
theta=config.theta,
frequency_mode=config.frequency_mode,
strength=config.temperature_strength,
device=device,
dtype=dtype,
),
]
return torch.cat(parts, dim=0)
+225
View File
@@ -0,0 +1,225 @@
from __future__ import annotations
import logging
from typing import Any, Callable, Optional
import torch
import torch.nn.functional as F
from .config import TIDEConfig
from .math import adaptive_text_bias, rope_temperature_scale
_LOG = logging.getLogger("ComfyUI-TIDE")
def _safe_timestep01(value: Any) -> float:
if value is None:
return 1.0
try:
if isinstance(value, torch.Tensor):
if value.numel() == 0:
return 1.0
v = value.detach().float().max().cpu().item()
else:
v = float(value)
except Exception:
return 1.0
# FLUX-family Comfy models normally receive flow timesteps/sigmas in [0, 1].
# Clamp rather than normalize by an unknown scheduler-specific maximum.
return max(0.0, min(1.0, float(v)))
def _block_enabled(config: TIDEConfig, extra_options: dict[str, Any]) -> bool:
block_type = extra_options.get("block_type")
if block_type == "double":
return config.apply_to_double_blocks
if block_type == "single":
return config.apply_to_single_blocks
# Unknown Flux-like DiT patch site. Apply conservatively if either stream type is enabled.
return config.apply_to_double_blocks or config.apply_to_single_blocks
def _add_text_bias_mask(
attn_mask: Optional[torch.Tensor],
*,
text_tokens: int,
key_tokens: int,
beta: float,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
bias = torch.zeros((1, 1, 1, key_tokens), device=device, dtype=dtype)
if text_tokens > 0 and beta != 0.0:
bias[..., :text_tokens] = float(beta)
if attn_mask is None:
return bias
if not torch.is_floating_point(attn_mask):
# Boolean masks represent validity, not additive logits. Preserve them rather than
# accidentally converting padding semantics into a dense additive bias. Current Flux
# T2I path normally has no boolean mask; Qwen's text mask is handled upstream.
_LOG.warning("TIDE skipped text anchoring because the existing attention mask is boolean/non-floating.")
return attn_mask
# Comfy attention accepts additive masks of shape [B, H|1, Q|1, K].
# Rely on PyTorch broadcasting and do not materialize full QxK masks.
return attn_mask.to(device=device, dtype=dtype) + bias
class TIDEAttentionPatch:
"""ComfyUI attn1_patch implementing TIDE text anchoring and RoPE temperature scaling."""
def __init__(self, config: TIDEConfig):
self.config = config
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):
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:
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,
)
out_pe = pe
if pe is not None and self.config.temperature_strength != 0.0:
tide_opts = extra_options.get("tide", {})
timestep = _safe_timestep01(tide_opts.get("timestep", extra_options.get("timestep")))
try:
scale = rope_temperature_scale(
self.config,
timestep=timestep,
device=pe.device,
dtype=pe.dtype if torch.is_floating_point(pe) else torch.float32,
)
# Comfy FLUX pe shape is [B, 1, N, Dpair, 2, 2] or broadcast-compatible.
if pe.shape[-3] == scale.numel():
view_shape = (1,) * (pe.ndim - 3) + (scale.numel(), 1, 1)
out_pe = pe * scale.reshape(view_shape)
elif self.config.debug:
_LOG.warning(
"TIDE skipped dynamic temperature: pe axis dimension %s != scale length %s",
pe.shape[-3], scale.numel(),
)
except Exception as exc:
if self.config.debug:
_LOG.exception("TIDE dynamic temperature failed and was skipped: %s", exc)
return {"q": q, "k": k, "v": v, "pe": out_pe, "attn_mask": out_mask}
def _sdpa_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
heads: int,
mask: Optional[torch.Tensor] = None,
attn_precision=None,
skip_reshape: bool = False,
skip_output_reshape: bool = False,
**kwargs,
) -> torch.Tensor:
"""Small unwrapped PyTorch SDPA path for additive TIDE masks.
This avoids xFormers/Sage/Flash backends that may reject additive masks or
materialize dense high-resolution masks. It intentionally mirrors ComfyUI's
attention_pytorch tensor layout without invoking the wrapped function again.
"""
if skip_reshape:
b, _, _, dim_head = q.shape
q_s, k_s, v_s = q, k, v
else:
b, _, dim_head = q.shape
dim_head //= heads
q_s, k_s, v_s = map(lambda t: t.view(b, -1, heads, dim_head).transpose(1, 2), (q, k, v))
if mask is not None:
if mask.ndim == 2:
mask = mask.unsqueeze(0)
if mask.ndim == 3:
mask = mask.unsqueeze(1)
mask = mask.to(device=q_s.device, dtype=q_s.dtype if torch.is_floating_point(mask) else mask.dtype)
out = F.scaled_dot_product_attention(q_s, k_s, v_s, attn_mask=mask, dropout_p=0.0, is_causal=False)
if skip_reshape:
if not skip_output_reshape:
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
else:
if skip_output_reshape:
out = out.transpose(1, 2)
else:
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
return out
class TIDEAttentionOverride:
"""Optional optimized_attention_override used only when an additive mask is present."""
def __init__(self, config: TIDEConfig, old_override: Optional[Callable] = None):
self.config = config
self.old_override = old_override
def to(self, device: torch.device | str):
return self
def __call__(self, original_func: Callable, *args, **kwargs):
mask = kwargs.get("mask", None)
if mask is None:
mask = kwargs.get("attn_mask", None)
if self.config.force_pytorch_attention_with_mask and mask is not None:
return _sdpa_attention(*args, **kwargs)
if self.old_override is not None:
return self.old_override(original_func, *args, **kwargs)
return original_func(*args, **kwargs)
class TIDEModelWrapper:
"""Inject current denoising timestep into transformer_options for DTC."""
def __init__(self, config: TIDEConfig, old_wrapper: Optional[Callable] = None):
self.config = config
self.old_wrapper = old_wrapper
def to(self, device: torch.device | str):
return self
def __call__(self, apply_model: Callable, args: dict[str, Any]) -> torch.Tensor:
c = args.get("c", {}).copy()
transformer_options = c.get("transformer_options", {}).copy()
tide_opts = transformer_options.get("tide", {}).copy()
tide_opts["timestep"] = _safe_timestep01(args.get("timestep"))
tide_opts["width"] = self.config.width
tide_opts["height"] = self.config.height
transformer_options["tide"] = tide_opts
c["transformer_options"] = transformer_options
if self.config.preserve_existing_wrapper and self.old_wrapper is not None:
return self.old_wrapper(apply_model, args | {"c": c})
return apply_model(args["input"], args["timestep"], **c)