From 01a324bb112dc8737b9bae6c41af364464aa5118 Mon Sep 17 00:00:00 2001 From: avtc Date: Thu, 23 Apr 2026 17:03:45 +0300 Subject: [PATCH] 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. --- __init__.py | 8 ++ advanced_tiling.py | 27 ++++- dit_tiling.py | 107 +++++++++++++++++ modes/__init__.py | 3 +- modes/hex.py | 39 +++++++ modes/rect.py | 39 +++++++ ray_tiling.py | 81 +++++++++++++ toroidal_attention.py | 264 ++++++++++++++++++++++++++++++++++++++++++ 8 files changed, 561 insertions(+), 7 deletions(-) create mode 100644 dit_tiling.py create mode 100644 modes/rect.py create mode 100644 ray_tiling.py create mode 100644 toroidal_attention.py diff --git a/__init__.py b/__init__.py index 2ab251e..74eba72 100644 --- a/__init__.py +++ b/__init__.py @@ -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"] diff --git a/advanced_tiling.py b/advanced_tiling.py index f33414f..8a11749 100644 --- a/advanced_tiling.py +++ b/advanced_tiling.py @@ -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,) diff --git a/dit_tiling.py b/dit_tiling.py new file mode 100644 index 0000000..dc323b1 --- /dev/null +++ b/dit_tiling.py @@ -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) diff --git a/modes/__init__.py b/modes/__init__.py index 12b3a6f..444e5e0 100644 --- a/modes/__init__.py +++ b/modes/__init__.py @@ -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, } diff --git a/modes/hex.py b/modes/hex.py index 379ade1..8689e96 100644 --- a/modes/hex.py +++ b/modes/hex.py @@ -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) diff --git a/modes/rect.py b/modes/rect.py new file mode 100644 index 0000000..b88c116 --- /dev/null +++ b/modes/rect.py @@ -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) diff --git a/ray_tiling.py b/ray_tiling.py new file mode 100644 index 0000000..fde6288 --- /dev/null +++ b/ray_tiling.py @@ -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,) diff --git a/toroidal_attention.py b/toroidal_attention.py new file mode 100644 index 0000000..bbb0fea --- /dev/null +++ b/toroidal_attention.py @@ -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)