Import ComfyUI-TIDE into skeleton repo
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
.pytest_cache/
|
||||
.venv/
|
||||
env/
|
||||
outputs/
|
||||
.codex/
|
||||
.vs/
|
||||
@@ -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.
|
||||
@@ -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"]
|
||||
@@ -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.
|
||||
@@ -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",
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
pythonpath = ["."]
|
||||
@@ -0,0 +1 @@
|
||||
torch
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user