Merge pull request #4 from avtc/feature/qwen-image-and-raylight-support
feature: Add DiT model support with toroidal attention, latent wrapping, and Raylight integration
This commit is contained in:
@@ -9,19 +9,27 @@
|
||||
## Implemented tiling modes
|
||||
|
||||
- [x] Hexagon
|
||||
- [x] Rectangular (toroidal — right edge wraps to left, bottom wraps to top)
|
||||
- [x] None (normal generation)
|
||||
|
||||
## Supported models
|
||||
|
||||
### UNet-based
|
||||
|
||||
- [x] Stable Diffusion 1.5 (also 1.4)
|
||||
- [x] Stable Diffusion 2.1 (also 2.0)
|
||||
- [x] Stable Diffusion XL (SDXL)
|
||||
|
||||
### DiT-based (via toroidal attention patching)
|
||||
|
||||
- [x] FLUX.2 — tested on flux.2-klein-4b
|
||||
- [x] Qwen Image — tested on qwen-image-2512 (Q6 GGUF and fp16)
|
||||
|
||||
DiT models are automatically detected and use toroidal attention instead of latent padding for seamless tiling.
|
||||
|
||||
## TODO
|
||||
|
||||
- More tiling modes
|
||||
- Optimize VAE decode (first pass is very slow)
|
||||
- Support DiT based models (SD3, PixArt-Σ, FLUX.1)
|
||||
|
||||
## Credits
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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,)
|
||||
@@ -0,0 +1,267 @@
|
||||
"""
|
||||
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
|
||||
|
||||
n_axes = len(self.pe_embedder.axes_dim)
|
||||
syn_ids = torch.zeros(1, self._n_extra, n_axes, 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)
|
||||
expand_shape = list(pe.shape)
|
||||
expand_shape[2] = synthetic_pe.shape[2]
|
||||
new_pe = torch.cat([
|
||||
pe,
|
||||
synthetic_pe.expand(expand_shape),
|
||||
], 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)
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
Reference in New Issue
Block a user