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)