feature: Add DiT model support with toroidal attention, latent wrapping, and Raylight integration

Extends tiling beyond UNet (Conv2d) models to support DiT architectures
using toroidal attention patches and latent content wrapping for seamless
infinite tiling. Adds rectangular tiling mode, WanVAE 5D tensor handling,
and an AdvancedTilingRay node for distributed Raylight workers.
This commit is contained in:
avtc
2026-04-23 17:03:45 +03:00
parent 4c673c71a3
commit 01a324bb11
8 changed files with 561 additions and 7 deletions
+8
View File
@@ -22,4 +22,12 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"AdvancedTilingVAEDecode": "Advanced Tiling VAE Decode",
}
try:
from .ray_tiling import AdvancedTilingRay, HAS_RAYLIGHT
if HAS_RAYLIGHT:
NODE_CLASS_MAPPINGS["AdvancedTilingRay"] = AdvancedTilingRay
NODE_DISPLAY_NAME_MAPPINGS["AdvancedTilingRay"] = "Advanced Tiling (Raylight)"
except ImportError:
pass
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+21 -6
View File
@@ -12,6 +12,7 @@ from torch.nn import Conv2d
from torch.nn import functional as F
from torch.nn.modules.utils import _pair
from .modes import modes, Settings
from .dit_tiling import patch_dit_model, _has_conv2d
@functools.cache
@@ -123,10 +124,15 @@ class AdvancedTilingSettings:
return {
"required": {
"mode": (list(modes.keys()),),
"mode": (list(modes.keys()), {
"tooltip": "Tiling mode. 'None' disables tiling, 'Hexagon' wraps edges in a hexagonal pattern, 'Rectangular' wraps right→left and bottom→top.",
}),
"rotation": (
"FLOAT",
{"default": 0.0, "min": 0.0, "max": 360.0, "step": 0.01},
{
"default": 0.0, "min": 0.0, "max": 360.0, "step": 0.01,
"tooltip": "Rotation angle in degrees for the tiling pattern.",
},
),
},
}
@@ -147,7 +153,7 @@ class AdvancedTilingSettings:
class AdvancedTiling:
"""
Patches Conv2D layers in a model to perform tiling
Patches model to perform tiling - supports both UNet (Conv2d) and DiT models
"""
# pylint: disable=invalid-name
@@ -174,8 +180,12 @@ class AdvancedTiling:
Does the actual patching of the model
"""
model_copy = copy.deepcopy(model)
patch_model(model_copy.model, settings)
model_copy = model.clone()
if _has_conv2d(model_copy.model.diffusion_model):
patch_model(model_copy.model, settings)
else:
patch_dit_model(model_copy, settings)
return (model_copy,)
@@ -223,9 +233,14 @@ class AdvancedTilingVAEDecode:
patch_model(vae_copy.first_stage_model, settings)
# Decode latents to image
image = vae_copy.decode(samples["samples"])
# WanVAE returns 5D (B, T, H, W, C), standard VAE returns 4D (B, H, W, C)
if image.ndim == 5:
image = image.squeeze(1)
if crop:
# Crop image based on tiling settings
mask = create_crop_mask(image.shape[2], image.shape[1], settings)
image = torch.cat((image, mask), dim=3)
image = torch.cat((image, mask.to(device=image.device)), dim=3)
return (image,)
+107
View File
@@ -0,0 +1,107 @@
"""
DiT model tiling via toroidal attention and latent wrapping.
For Hexagon mode:
- Toroidal attention: injects wrapped-neighbor K/V entries into every attention
layer, making boundary patches structurally "see" opposite-edge content as
spatially adjacent.
- Latent wrapping: copies source content to waste positions in the input latent
on each denoising step (for KSampler preview).
For Rectangular mode:
- Toroidal attention: injects wrapped K/V entries from opposite edges (right↔left,
top↔bottom).
"""
import torch
import torch.nn as nn
from torch.nn import Conv2d
from .modes import Settings
def _has_conv2d(model: nn.Module) -> bool:
"""Check if the model has any Conv2d layers (UNet vs DiT)."""
return any(isinstance(m, Conv2d) for m in model.modules())
def _create_content_wrapper(settings: Settings):
"""
Create a model function wrapper that copies source content to waste
positions in the latent on each denoising step. This makes the KSampler
preview show wrapped content instead of noise.
Only used for Hexagon mode.
:param settings: Tiling settings
"""
from .advanced_tiling import calculate_mapping
_mapping_cache = {}
def wrapper(apply_model, args):
x = args["input"]
is_5d = x.ndim == 5
if is_5d:
_, _, _, H, W = x.shape
else:
_, _, H, W = x.shape
cache_key = (W, H, hash(settings))
if cache_key not in _mapping_cache:
_mapping_cache[cache_key] = calculate_mapping(
(W, H), (W, H), settings
)
mapping = _mapping_cache[cache_key]
# Content replacement in latent
if is_5d:
x[:, :, :, mapping[1], mapping[0]] = x[:, :, :, mapping[3], mapping[2]]
else:
x[:, :, mapping[1], mapping[0]] = x[:, :, mapping[3], mapping[2]]
return apply_model(args["input"], args["timestep"], **args["c"])
return wrapper
def patch_dit_model(model_patcher, settings: Settings):
"""
Apply tiling to a DiT model.
For Hexagon mode: applies toroidal attention + latent content wrapping.
For Rectangular mode: applies toroidal attention only.
:param model_patcher: ComfyUI ModelPatcher instance
:param settings: Tiling settings
"""
diff_model = model_patcher.model.diffusion_model
if settings.mode == "Hexagon":
from .toroidal_attention import HexToroidalAttentionPatch
if not hasattr(diff_model, 'pe_embedder'):
raise ValueError(
"Model does not have pe_embedder. "
"Toroidal attention requires a model with RoPE position embeddings."
)
patch = HexToroidalAttentionPatch(settings, diff_model.pe_embedder)
model_patcher.set_model_attn1_patch(patch)
# Latent content wrapping for KSampler preview
wrapper = _create_content_wrapper(settings)
model_patcher.set_model_unet_function_wrapper(wrapper)
elif settings.mode == "Rectangular":
from .toroidal_attention import RectToroidalAttentionPatch
if not hasattr(diff_model, 'pe_embedder'):
raise ValueError(
"Model does not have pe_embedder. "
"Toroidal attention requires a model with RoPE position embeddings."
)
patch = RectToroidalAttentionPatch(diff_model.pe_embedder)
model_patcher.set_model_attn1_patch(patch)
+2 -1
View File
@@ -18,16 +18,17 @@ class Settings:
self.rotation = rotation
def __hash__(self):
# We don't care about the tiling function, because it's determined by the mode
return hash((self.mode, self.rotation))
from .hex import hex_tiling
from .none import none_tiling
from .rect import rect_tiling
modes = {
"None": none_tiling,
"Hexagon": hex_tiling,
"Rectangular": rect_tiling,
}
+39
View File
@@ -187,3 +187,42 @@ def hex_tiling(
new_y = (new_y + padded_size[1] // 2) % padded_size[1]
return (new_x, new_y)
def hex_patch_tiling(
patch_h: int,
patch_w: int,
original_h_patches: int,
original_w_patches: int,
padded_h_patches: int,
padded_w_patches: int,
settings: Settings,
) -> tuple[int, int]:
"""
Hexagonal tiling at patch granularity for DiT models.
Maps a position in the padded patch grid to its source position
in the original patch grid using hexagonal coordinate remapping.
:param patch_h: Row index in padded patch grid
:param patch_w: Column index in padded patch grid
:param original_h_patches: Number of patch rows in original grid
:param original_w_patches: Number of patch columns in original grid
:param padded_h_patches: Number of patch rows in padded grid
:param padded_w_patches: Number of patch columns in padded grid
:param settings: Tiling settings
:return: (source_h, source_w) in the original patch grid
"""
size = min(original_h_patches, original_w_patches) // 2
q, r = pixel_to_hex(
(patch_w - padded_w_patches // 2, patch_h - padded_h_patches // 2),
size,
settings,
)
rounded = axial_round((q, r))
q -= rounded[0]
r -= rounded[1]
new_w, new_h = hex_to_pixel((q, r), size, settings)
new_h = (new_h + padded_h_patches // 2) % padded_h_patches
new_w = (new_w + padded_w_patches // 2) % padded_w_patches
return (new_h, new_w)
+39
View File
@@ -0,0 +1,39 @@
"""
Rectangular (toroidal) tiling implementation
Wraps coordinates modularly: right edge wraps to left, bottom wraps to top.
"""
from . import Settings
def rect_tiling(
x: int,
y: int,
original_size: tuple[int, int],
padded_size: tuple[int, int],
_settings: Settings,
) -> tuple[int, int]:
"""
Rectangular tiling: wraps coordinates modularly around the original area.
Positions inside original_size are identity. Positions in the padding area
wrap to the opposite side of the original content.
:param x: X coordinate in padded space
:param y: Y coordinate in padded space
:param original_size: (width, height) of original content
:param padded_size: (width, height) of padded tensor
:param _settings: Tiling settings (unused for rectangular)
:return: (new_x, new_y) source coordinates in padded space
"""
ow, oh = original_size
pw, ph = padded_size
pad_x = (pw - ow) // 2
pad_y = (ph - oh) // 2
rel_x = (x - pad_x) % ow
rel_y = (y - pad_y) % oh
return (rel_x + pad_x, rel_y + pad_y)
+81
View File
@@ -0,0 +1,81 @@
"""
Raylight integration for DiT tiling.
Provides AdvancedTilingRay node that works with raylight's RAY_ACTORS type,
applying toroidal attention tiling on distributed workers.
Uses manual Ray remote calls instead of @ray_patch to avoid serialization
issues with custom node packages that have non-standard import paths
(e.g. hyphens in directory names that prevent normal Python imports).
"""
try:
from raylight.comfy_extra_dist.ray_patch_decorator import ray_patch
HAS_RAYLIGHT = True
except ImportError:
HAS_RAYLIGHT = False
if HAS_RAYLIGHT:
import ray
class AdvancedTilingRay:
"""
Applies tiling to a raylight RAY_ACTORS model.
Dispatches per-worker patching via Ray remote calls.
"""
# pylint: disable=invalid-name
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"settings": ("ADVANCED_TILING_SETTINGS",),
"ray_actors": ("RAY_ACTORS",),
},
}
RETURN_TYPES = ("RAY_ACTORS",)
RETURN_NAMES = ("ray_actors",)
FUNCTION = "run"
CATEGORY = "conditioning"
def run(self, settings, ray_actors):
# Defined inside run() so cloudpickle serializes it as a nested
# function (by value) instead of by module reference. Module-level
# functions get serialized by reference, which requires importing
# the module by name on the worker -- but the module name is the
# filesystem path (with a hyphen), causing ModuleNotFoundError.
def _patch(model, mode, rotation):
import importlib.util
import os
import sys
_node_dir = os.path.dirname(os.path.abspath(__file__))
_pkg = "ComfyUI_AdvancedTiling"
if _pkg not in sys.modules:
_init = os.path.join(_node_dir, "__init__.py")
spec = importlib.util.spec_from_file_location(
_pkg, _init,
submodule_search_locations=[_node_dir],
)
mod = importlib.util.module_from_spec(spec)
sys.modules[_pkg] = mod
spec.loader.exec_module(mod)
from ComfyUI_AdvancedTiling.dit_tiling import patch_dit_model
from ComfyUI_AdvancedTiling.modes import Settings
patch_dit_model(model, Settings(mode, rotation))
return model
gpu_workers = ray_actors["workers"]
futures = [
actor.model_function_runner.remote(
_patch, settings.mode, settings.rotation
)
for actor in gpu_workers
]
ray.get(futures)
return (ray_actors,)
+264
View File
@@ -0,0 +1,264 @@
"""
Toroidal attention patches for hex and rectangular tiling in DiT models.
Injects wrapped-neighbor K/V entries into every attention layer via attn1_patch,
so boundary patches structurally "see" opposite-edge content as spatially adjacent.
Analogous to how Conv2d circular padding works for UNet models.
"""
import torch
from .modes import Settings
from .modes.hex import hex_tiling
def _factorize(n: int) -> tuple[int, int]:
"""Factorize n into h * w, preferring square."""
s = int(n ** 0.5)
while s > 0:
if n % s == 0:
return s, n // s
s -= 1
return 1, n
# ---------------------------------------------------------------------------
# Hex toroidal attention
# ---------------------------------------------------------------------------
_boundary_cache: dict = {}
def _compute_hex_boundary_pairs(
h_patches: int,
w_patches: int,
settings: Settings,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Compute hex boundary neighbor relationships at patch granularity.
For each patch inside the hex that has a 4-connected neighbor outside
the hex, record the boundary patch index, the wrapped source index,
and the direction offset.
Cached by (h_patches, w_patches, hash(settings)).
:param h_patches: Number of patch rows
:param w_patches: Number of patch columns
:param settings: Tiling settings
:return: (boundary_idx, source_idx, off_h, off_w) as LongTensors
"""
cache_key = (h_patches, w_patches, hash(settings))
if cache_key in _boundary_cache:
return _boundary_cache[cache_key]
boundary_idx = []
source_idx = []
offsets_h = []
offsets_w = []
for h in range(h_patches):
for w in range(w_patches):
src_w, src_h = hex_tiling(
w, h,
(w_patches, h_patches),
(w_patches, h_patches),
settings,
)
if src_w != w or src_h != h:
continue
for dh, dw in [(-1, 0), (1, 0), (0, -1), (0, 1)]:
nh, nw = h + dh, w + dw
n_src_w, n_src_h = hex_tiling(
nw, nh,
(w_patches, h_patches),
(w_patches, h_patches),
settings,
)
if n_src_w != nw or n_src_h != nh:
boundary_idx.append(h * w_patches + w)
source_idx.append(n_src_h * w_patches + n_src_w)
offsets_h.append(dh)
offsets_w.append(dw)
if not boundary_idx:
empty = torch.tensor([], dtype=torch.long)
result = (empty, empty.clone(), empty.clone(), empty.clone())
else:
result = (
torch.tensor(boundary_idx, dtype=torch.long),
torch.tensor(source_idx, dtype=torch.long),
torch.tensor(offsets_h, dtype=torch.long),
torch.tensor(offsets_w, dtype=torch.long),
)
_boundary_cache[cache_key] = result
return result
# ---------------------------------------------------------------------------
# Rectangular toroidal attention
# ---------------------------------------------------------------------------
_rect_boundary_cache: dict = {}
def _compute_rect_boundary_pairs(
h_patches: int,
w_patches: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Compute rectangular boundary pairs for toroidal wrapping.
For each patch on the edge of the grid, records the boundary patch index,
the wrapped source index from the opposite edge, and the direction offset.
Cached by (h_patches, w_patches).
:param h_patches: Number of patch rows
:param w_patches: Number of patch columns
:return: (boundary_idx, source_idx, off_h, off_w) as LongTensors
"""
cache_key = (h_patches, w_patches)
if cache_key in _rect_boundary_cache:
return _rect_boundary_cache[cache_key]
boundary_idx = []
source_idx = []
offsets_h = []
offsets_w = []
for h in range(h_patches):
for w in range(w_patches):
for dh, dw in [(-1, 0), (1, 0), (0, -1), (0, 1)]:
nh, nw = h + dh, w + dw
if nh < 0 or nh >= h_patches or nw < 0 or nw >= w_patches:
wrapped_h = nh % h_patches
wrapped_w = nw % w_patches
boundary_idx.append(h * w_patches + w)
source_idx.append(wrapped_h * w_patches + wrapped_w)
offsets_h.append(dh)
offsets_w.append(dw)
if not boundary_idx:
empty = torch.tensor([], dtype=torch.long)
result = (empty, empty.clone(), empty.clone(), empty.clone())
else:
result = (
torch.tensor(boundary_idx, dtype=torch.long),
torch.tensor(source_idx, dtype=torch.long),
torch.tensor(offsets_h, dtype=torch.long),
torch.tensor(offsets_w, dtype=torch.long),
)
_rect_boundary_cache[cache_key] = result
return result
# ---------------------------------------------------------------------------
# Shared base for attention patches
# ---------------------------------------------------------------------------
class _BaseToroidalAttentionPatch:
"""Shared logic for hex and rectangular toroidal attention patches."""
def __init__(self, pe_embedder):
self.pe_embedder = pe_embedder
self._initialized = False
self._boundary_idx = None
self._source_idx = None
self._n_extra = 0
self._synthetic_pe = None
def _compute_boundary_pairs(self, h_patches, w_patches):
raise NotImplementedError
def _initialize(self, n_img: int):
h_patches, w_patches = _factorize(n_img)
boundary_idx, source_idx, off_h, off_w = self._compute_boundary_pairs(
h_patches, w_patches,
)
self._boundary_idx = boundary_idx
self._source_idx = source_idx
self._n_extra = len(boundary_idx)
if self._n_extra == 0:
self._initialized = True
return
boundary_h = boundary_idx // w_patches
boundary_w = boundary_idx % w_patches
h_center = h_patches // 2
w_center = w_patches // 2
syn_ids = torch.zeros(1, self._n_extra, 3, dtype=torch.float32)
syn_ids[0, :, 0] = 0
syn_ids[0, :, 1] = (boundary_h + off_h).float() - h_center
syn_ids[0, :, 2] = (boundary_w + off_w).float() - w_center
with torch.no_grad():
self._synthetic_pe = self.pe_embedder(syn_ids)
self._initialized = True
def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None):
if extra_options is None or "img_slice" not in extra_options:
return {}
if attn_mask is not None:
return {}
if not self._initialized:
img_slice = extra_options["img_slice"]
n_img = img_slice[1] - img_slice[0]
self._initialize(n_img)
if self._n_extra == 0:
return {}
n_txt = extra_options["img_slice"][0]
src_joint_idx = self._source_idx.to(k.device) + n_txt
extra_k = k[:, :, src_joint_idx, :]
extra_v = v[:, :, src_joint_idx, :]
new_k = torch.cat([k, extra_k], dim=2)
new_v = torch.cat([v, extra_v], dim=2)
synthetic_pe = self._synthetic_pe.to(device=pe.device, dtype=pe.dtype)
new_pe = torch.cat([
pe,
synthetic_pe.expand(pe.shape[0], -1, -1, -1, -1, -1),
], dim=2)
return {
"k": new_k,
"v": new_v,
"pe": new_pe,
}
class HexToroidalAttentionPatch(_BaseToroidalAttentionPatch):
"""attn1_patch for hex tiling: injects wrapped K/V for hex boundary patches."""
def __init__(self, settings: Settings, pe_embedder):
super().__init__(pe_embedder)
self.settings = settings
def _compute_boundary_pairs(self, h_patches, w_patches):
return _compute_hex_boundary_pairs(h_patches, w_patches, self.settings)
class RectToroidalAttentionPatch(_BaseToroidalAttentionPatch):
"""attn1_patch for rectangular tiling: injects wrapped K/V from opposite edges."""
def _compute_boundary_pairs(self, h_patches, w_patches):
return _compute_rect_boundary_pairs(h_patches, w_patches)