feat(attention): add sampler-derived concept regions
This commit is contained in:
@@ -12,6 +12,9 @@ from . import simple_syrup as _simple_syrup_package
|
||||
|
||||
sys.modules.setdefault("simple_syrup", _simple_syrup_package)
|
||||
|
||||
from .simple_syrup.runtime.attention_region_prompt_handler import ( # noqa: E402
|
||||
register_attention_region_prompt_handler,
|
||||
)
|
||||
from .simple_syrup.runtime.comfy_safetensors_dtypes import ( # noqa: E402
|
||||
register_comfy_safetensors_dtypes,
|
||||
)
|
||||
@@ -52,6 +55,7 @@ register_comfy_safetensors_dtypes()
|
||||
register_quant_cache_routes()
|
||||
register_external_llm_routes()
|
||||
register_mask_batch_preview_routes()
|
||||
register_attention_region_prompt_handler()
|
||||
|
||||
__all__ = [
|
||||
"WEB_DIRECTORY",
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Parse explicit attention concepts without interpreting prompt language."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def parse_attention_concepts(value: str) -> tuple[str, ...]:
|
||||
"""Return canonical concepts separated only by vertical bars."""
|
||||
|
||||
if not isinstance(value, str):
|
||||
raise TypeError("Attention concepts must be text.")
|
||||
return tuple(part.strip() for part in value.split("|") if part.strip())
|
||||
@@ -0,0 +1,30 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Resolve explicit two-dimensional geometry for flattened attention maps."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
|
||||
def factor_spatial_geometry(
|
||||
token_count: int, *, target_aspect: float
|
||||
) -> tuple[int, int]:
|
||||
"""Return the token factor pair nearest the supplied positive aspect ratio."""
|
||||
|
||||
if type(token_count) is not int or token_count < 1:
|
||||
raise ValueError("Attention spatial token count must be positive.")
|
||||
if target_aspect <= 0.0:
|
||||
raise ValueError("Attention target aspect ratio must be positive.")
|
||||
candidates: list[tuple[float, int, int]] = []
|
||||
for height in range(1, math.isqrt(token_count) + 1):
|
||||
if token_count % height:
|
||||
continue
|
||||
width = token_count // height
|
||||
for candidate_height, candidate_width in ((height, width), (width, height)):
|
||||
error = abs(math.log((candidate_width / candidate_height) / target_aspect))
|
||||
candidates.append((error, candidate_height, candidate_width))
|
||||
_error, height, width = min(candidates)
|
||||
return height, width
|
||||
@@ -0,0 +1,157 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define immutable requests and plans for attention-region capture."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
from .attention_spatial_transform import AttentionSpatialTransform
|
||||
from .graph_provenance import GraphLink
|
||||
|
||||
|
||||
class AttentionRegionRequestKind(StrEnum):
|
||||
"""Identify one public attention-region operation."""
|
||||
|
||||
CONCEPT_SEGS = "concept_segs"
|
||||
ALL_PROMPT_SEGS = "all_prompt_segs"
|
||||
REGION_MASK = "region_mask"
|
||||
MASKED_CONDITIONING = "masked_conditioning"
|
||||
|
||||
|
||||
class AttentionCaptureProfile(StrEnum):
|
||||
"""Select the density of attention observations retained during sampling."""
|
||||
|
||||
FAST = "fast"
|
||||
BALANCED = "balanced"
|
||||
EXHAUSTIVE = "exhaustive"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionRegionControls:
|
||||
"""Hold validated attention-native capture and region-shaping controls."""
|
||||
|
||||
capture_start: float
|
||||
capture_end: float
|
||||
minimum_strength: float
|
||||
minimum_consensus: float
|
||||
split_sensitivity: float
|
||||
minimum_region_size: int
|
||||
profile: AttentionCaptureProfile
|
||||
keep_only: int = 0
|
||||
keep_by: str = "largest size"
|
||||
combine_segs: bool = False
|
||||
matte_solidity: float = 0.0
|
||||
edge_feather: int = 8
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require normalized ranges and a non-empty capture interval."""
|
||||
|
||||
normalized = (
|
||||
self.capture_start,
|
||||
self.capture_end,
|
||||
self.minimum_strength,
|
||||
self.minimum_consensus,
|
||||
self.split_sensitivity,
|
||||
self.matte_solidity,
|
||||
)
|
||||
if any(
|
||||
isinstance(value, bool) or not isinstance(value, int | float)
|
||||
for value in normalized
|
||||
):
|
||||
raise TypeError("Attention-region controls must be real numbers.")
|
||||
if not 0.0 <= self.capture_start < self.capture_end <= 1.0:
|
||||
raise ValueError("Attention capture start must be below end within 0..1.")
|
||||
if not 0.0 <= self.minimum_strength <= 1.0:
|
||||
raise ValueError("Minimum attention strength must be within 0..1.")
|
||||
if not 0.0 <= self.minimum_consensus <= 1.0:
|
||||
raise ValueError("Minimum attention consensus must be within 0..1.")
|
||||
if not 0.0 <= self.split_sensitivity <= 1.0:
|
||||
raise ValueError("Attention split sensitivity must be within 0..1.")
|
||||
if type(self.minimum_region_size) is not int or self.minimum_region_size < 1:
|
||||
raise ValueError("Minimum attention region size must be positive.")
|
||||
if type(self.keep_only) is not int or self.keep_only < 0:
|
||||
raise ValueError("Attention keep_only must be non-negative.")
|
||||
if self.keep_by not in ("largest size", "highest confidence"):
|
||||
raise ValueError("Attention keep_by has an invalid policy.")
|
||||
if type(self.combine_segs) is not bool:
|
||||
raise TypeError("Attention combine_segs must be boolean.")
|
||||
if not 0.0 <= self.matte_solidity <= 1.0:
|
||||
raise ValueError("Attention matte solidity must be within 0..1.")
|
||||
if type(self.edge_feather) is not int or self.edge_feather < 0:
|
||||
raise ValueError("Attention edge feather must be non-negative.")
|
||||
if not isinstance(self.profile, AttentionCaptureProfile):
|
||||
raise TypeError("Attention capture profile has an invalid type.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionRegionRequest:
|
||||
"""Bind one public node request to its queries and capture controls."""
|
||||
|
||||
node_id: str
|
||||
kind: AttentionRegionRequestKind
|
||||
queries: tuple[str, ...]
|
||||
controls: AttentionRegionControls
|
||||
sampler_stage: int = 1
|
||||
spatial_transforms: tuple[AttentionSpatialTransform, ...] = ()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require stable node identity and canonical non-empty query strings."""
|
||||
|
||||
if not self.node_id.strip():
|
||||
raise ValueError("Attention-region request node id cannot be empty.")
|
||||
if not isinstance(self.kind, AttentionRegionRequestKind):
|
||||
raise TypeError("Attention-region request kind has an invalid type.")
|
||||
if any(not query or query != query.strip() for query in self.queries):
|
||||
raise ValueError("Attention-region queries must be canonical strings.")
|
||||
if self.kind is AttentionRegionRequestKind.ALL_PROMPT_SEGS:
|
||||
if self.queries:
|
||||
raise ValueError(
|
||||
"All-prompt attention requests cannot contain queries."
|
||||
)
|
||||
elif not self.queries:
|
||||
raise ValueError("Concept and mask attention requests require concepts.")
|
||||
if type(self.sampler_stage) is not int or self.sampler_stage < -1:
|
||||
raise ValueError("Attention sampler stage must be -1 or greater.")
|
||||
if any(
|
||||
not isinstance(transform, AttentionSpatialTransform)
|
||||
for transform in self.spatial_transforms
|
||||
):
|
||||
raise TypeError("Attention request spatial transforms are invalid.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionCapturePlan:
|
||||
"""Describe one coalesced sampler capture and its graph rewrite authority."""
|
||||
|
||||
sampler_node_id: str
|
||||
model_owner_node_id: str
|
||||
model_input_name: str
|
||||
model_link: GraphLink
|
||||
positive_link: GraphLink
|
||||
requests: tuple[AttentionRegionRequest, ...]
|
||||
prompt_text: str | None = None
|
||||
clip_link: GraphLink | None = None
|
||||
source_aspect: float | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require canonical unique requests and complete graph-edge identity."""
|
||||
|
||||
if not self.sampler_node_id or not self.model_owner_node_id:
|
||||
raise ValueError("Attention capture plan node ids cannot be empty.")
|
||||
if not self.model_input_name:
|
||||
raise ValueError("Attention capture plan model input cannot be empty.")
|
||||
if self.source_aspect is not None and self.source_aspect <= 0.0:
|
||||
raise ValueError("Attention capture source aspect must be positive.")
|
||||
request_ids = tuple(request.node_id for request in self.requests)
|
||||
if not request_ids or request_ids != tuple(sorted(set(request_ids))):
|
||||
raise ValueError("Attention capture requests must be unique and ordered.")
|
||||
|
||||
@property
|
||||
def capture_node_id(self) -> str:
|
||||
"""Return a collision-resistant deterministic injected node id."""
|
||||
|
||||
return f"__simple_syrup_attention_capture__{self.sampler_node_id}"
|
||||
@@ -0,0 +1,164 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Model prompt-token spans and compact captured attention observations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from .attention_spatial_transform import AttentionSpatialTransform
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionTokenSpan:
|
||||
"""Bind a readable prompt occurrence to exact conditioning token positions."""
|
||||
|
||||
label: str
|
||||
occurrence: int
|
||||
token_indices: tuple[int, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require canonical labels and ordered non-negative token indices."""
|
||||
|
||||
if not self.label or self.label != self.label.strip():
|
||||
raise ValueError("Attention token span label must be canonical.")
|
||||
if type(self.occurrence) is not int or self.occurrence < 1:
|
||||
raise ValueError("Attention token span occurrence must be positive.")
|
||||
if (
|
||||
not self.token_indices
|
||||
or self.token_indices != tuple(sorted(set(self.token_indices)))
|
||||
or self.token_indices[0] < 0
|
||||
):
|
||||
raise ValueError("Attention token indices must be ordered and unique.")
|
||||
|
||||
@property
|
||||
def display_label(self) -> str:
|
||||
"""Disambiguate repeated concepts while keeping first labels concise."""
|
||||
|
||||
return (
|
||||
self.label if self.occurrence == 1 else f"{self.label} #{self.occurrence}"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionTokenCatalog:
|
||||
"""Hold readable prompt spans and conditioning sequence length."""
|
||||
|
||||
sequence_length: int
|
||||
spans: tuple[AttentionTokenSpan, ...]
|
||||
token_ids: tuple[object, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require all spans to fit the captured conditioning sequence."""
|
||||
|
||||
if type(self.sequence_length) is not int or self.sequence_length < 1:
|
||||
raise ValueError("Attention token sequence length must be positive.")
|
||||
if len(self.token_ids) != self.sequence_length:
|
||||
raise ValueError("Attention token ids must match the sequence length.")
|
||||
if any(
|
||||
index >= self.sequence_length
|
||||
for span in self.spans
|
||||
for index in span.token_indices
|
||||
):
|
||||
raise ValueError("Attention token span exceeds its conditioning sequence.")
|
||||
|
||||
def exact_matches(self, query: str) -> tuple[AttentionTokenSpan, ...]:
|
||||
"""Return every prompt occurrence whose normalized label equals a query."""
|
||||
|
||||
normalized = _normalized_label(query)
|
||||
return tuple(
|
||||
span for span in self.spans if _normalized_label(span.label) == normalized
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CapturedAttentionMap:
|
||||
"""Store one head-aggregated token map and its denoising observation identity."""
|
||||
|
||||
label: str
|
||||
values: torch.Tensor
|
||||
progress: float
|
||||
layer_key: str
|
||||
batch_index: int = 0
|
||||
confidence: float = 1.0
|
||||
spatial_height: int | None = None
|
||||
spatial_width: int | None = None
|
||||
spatial_transforms: tuple[AttentionSpatialTransform, ...] = ()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require a finite CPU spatial vector and normalized progress."""
|
||||
|
||||
if not self.label:
|
||||
raise ValueError("Captured attention map label cannot be empty.")
|
||||
if (
|
||||
not isinstance(self.values, torch.Tensor)
|
||||
or self.values.device.type != "cpu"
|
||||
or self.values.ndim != 1
|
||||
or self.values.numel() < 1
|
||||
or not self.values.is_floating_point()
|
||||
or not torch.isfinite(self.values).all().item()
|
||||
):
|
||||
raise ValueError("Captured attention values must be a finite CPU vector.")
|
||||
if not 0.0 <= self.progress <= 1.0:
|
||||
raise ValueError("Captured attention progress must be within 0..1.")
|
||||
if not self.layer_key:
|
||||
raise ValueError("Captured attention layer key cannot be empty.")
|
||||
if type(self.batch_index) is not int or self.batch_index < 0:
|
||||
raise ValueError("Captured attention batch index must be non-negative.")
|
||||
if not 0.0 <= self.confidence <= 1.0:
|
||||
raise ValueError("Captured attention confidence must be within 0..1.")
|
||||
if (self.spatial_height is None) != (self.spatial_width is None):
|
||||
raise ValueError("Captured attention geometry must be complete or absent.")
|
||||
if self.spatial_height is not None and (
|
||||
type(self.spatial_height) is not int
|
||||
or self.spatial_height < 1
|
||||
or type(self.spatial_width) is not int
|
||||
or self.spatial_width < 1
|
||||
or self.spatial_height * self.spatial_width != int(self.values.numel())
|
||||
):
|
||||
raise ValueError("Captured attention geometry must match its values.")
|
||||
if any(
|
||||
not isinstance(transform, AttentionSpatialTransform)
|
||||
for transform in self.spatial_transforms
|
||||
):
|
||||
raise TypeError("Captured attention spatial transforms are invalid.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OpenVocabularyContext:
|
||||
"""Hold one query's encoded SDXL context and semantic token positions."""
|
||||
|
||||
label: str
|
||||
values: torch.Tensor
|
||||
token_indices: tuple[int, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require one finite CPU context with valid unique token positions."""
|
||||
|
||||
if not self.label.strip() or self.label != self.label.strip():
|
||||
raise ValueError("Open-vocabulary labels must be canonical strings.")
|
||||
if self.values.ndim != 3 or int(self.values.shape[0]) != 1:
|
||||
raise ValueError("Open-vocabulary context must have shape 1xTxC.")
|
||||
if self.values.device.type != "cpu" or not torch.isfinite(self.values).all():
|
||||
raise ValueError("Open-vocabulary context must be finite CPU storage.")
|
||||
if not self.token_indices or self.token_indices != tuple(
|
||||
sorted(set(self.token_indices))
|
||||
):
|
||||
raise ValueError(
|
||||
"Open-vocabulary token positions must be unique and ordered."
|
||||
)
|
||||
if any(
|
||||
index < 0 or index >= int(self.values.shape[1])
|
||||
for index in self.token_indices
|
||||
):
|
||||
raise ValueError("Open-vocabulary token positions exceed their context.")
|
||||
|
||||
|
||||
def _normalized_label(value: str) -> str:
|
||||
"""Normalize human prompt labels without changing tokenizer semantics."""
|
||||
|
||||
return " ".join(value.casefold().replace("_", " ").split())
|
||||
@@ -0,0 +1,78 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Model ordered sampling stages along one spatial provenance path."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .attention_spatial_transform import AttentionSpatialTransform
|
||||
from .graph_provenance import GraphLink
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSamplerStage:
|
||||
"""Describe one sampler's patch and conditioning authority."""
|
||||
|
||||
sampler_node_id: str
|
||||
model_owner_node_id: str
|
||||
model_link: GraphLink
|
||||
positive_link: GraphLink
|
||||
upstream_link: GraphLink
|
||||
upstream_kind: str
|
||||
forward_transforms: tuple[AttentionSpatialTransform, ...] = ()
|
||||
source_aspect: float | None = None
|
||||
capture_supported: bool = True
|
||||
unsupported_reason: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require unsupported stages to explain why capture cannot be projected."""
|
||||
|
||||
if self.capture_supported and self.unsupported_reason is not None:
|
||||
raise ValueError("Supported attention sampler stages cannot have a reason.")
|
||||
if not self.capture_supported and not self.unsupported_reason:
|
||||
raise ValueError("Unsupported attention sampler stages require a reason.")
|
||||
if self.source_aspect is not None and self.source_aspect <= 0.0:
|
||||
raise ValueError("Attention sampler source aspect must be positive.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSamplerSelection:
|
||||
"""Bind a selected stage to its one-based chronological position."""
|
||||
|
||||
stage: AttentionSamplerStage
|
||||
stage_number: int
|
||||
stage_count: int
|
||||
was_clamped: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSamplerLineage:
|
||||
"""Hold sampling stages ordered from oldest to direct provenance."""
|
||||
|
||||
stages: tuple[AttentionSamplerStage, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require at least one uniquely identified stage."""
|
||||
|
||||
identities = tuple(stage.sampler_node_id for stage in self.stages)
|
||||
if not identities or len(set(identities)) != len(identities):
|
||||
raise ValueError("Attention sampler lineage must contain unique stages.")
|
||||
|
||||
def select(self, requested_stage: int) -> AttentionSamplerSelection:
|
||||
"""Resolve one-based selection with 0/-1 aliases for direct provenance."""
|
||||
|
||||
if type(requested_stage) is not int or requested_stage < -1:
|
||||
raise ValueError("Attention sampler stage must be -1 or greater.")
|
||||
count = len(self.stages)
|
||||
if requested_stage in (-1, 0):
|
||||
return AttentionSamplerSelection(self.stages[-1], count, count, False)
|
||||
selected_number = min(requested_stage, count)
|
||||
return AttentionSamplerSelection(
|
||||
self.stages[selected_number - 1],
|
||||
selected_number,
|
||||
count,
|
||||
requested_stage > count,
|
||||
)
|
||||
@@ -0,0 +1,70 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Describe graph-visible full-canvas transformations for attention masks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
_ANCHORS = frozenset(
|
||||
{
|
||||
"center",
|
||||
"top-left",
|
||||
"top",
|
||||
"top-right",
|
||||
"left",
|
||||
"right",
|
||||
"bottom-left",
|
||||
"bottom",
|
||||
"bottom-right",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class AttentionSpatialTransformKind(StrEnum):
|
||||
"""Identify a supported mask-coordinate transformation."""
|
||||
|
||||
RESIZE = "resize"
|
||||
FIT_RESIZE = "fit_resize"
|
||||
SCALE = "scale"
|
||||
COVER_CROP = "cover_crop"
|
||||
FIT_PAD = "fit_pad"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSpatialTransform:
|
||||
"""Hold validated resize, crop, or pad parameters from one graph node."""
|
||||
|
||||
kind: AttentionSpatialTransformKind
|
||||
width: int | None = None
|
||||
height: int | None = None
|
||||
scale: float | None = None
|
||||
anchor: str = "center"
|
||||
divisible_by: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require complete parameters for the selected transformation kind."""
|
||||
|
||||
if self.anchor not in _ANCHORS:
|
||||
raise ValueError("Attention spatial transform anchor is invalid.")
|
||||
if self.kind is AttentionSpatialTransformKind.SCALE:
|
||||
if self.scale is None or self.scale <= 0.0:
|
||||
raise ValueError("Attention scale transform requires a positive scale.")
|
||||
if self.width is not None or self.height is not None:
|
||||
raise ValueError("Attention scale transform cannot contain a size.")
|
||||
if self.divisible_by != 1:
|
||||
raise ValueError("Attention scale transform cannot set divisibility.")
|
||||
return
|
||||
if (
|
||||
type(self.width) is not int
|
||||
or self.width < 1
|
||||
or type(self.height) is not int
|
||||
or self.height < 1
|
||||
or self.scale is not None
|
||||
):
|
||||
raise ValueError("Attention spatial transform requires a positive size.")
|
||||
if type(self.divisible_by) is not int or self.divisible_by < 1:
|
||||
raise ValueError("Attention spatial divisibility must be positive.")
|
||||
@@ -2,18 +2,19 @@
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Connected-component helpers for binary mask regions."""
|
||||
"""Extract deterministic connected components from binary mask regions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from importlib import import_module
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.segs import BoundingBox
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MaskComponent:
|
||||
"""Represent one connected mask component and its full-image bbox."""
|
||||
|
||||
@@ -22,62 +23,30 @@ class MaskComponent:
|
||||
|
||||
|
||||
def connected_mask_components(active_mask: torch.Tensor) -> tuple[MaskComponent, ...]:
|
||||
"""Return 8-connected components from an HW active-pixel mask."""
|
||||
"""Return spatially ordered 8-connected components from one HW mask."""
|
||||
|
||||
if active_mask.ndim != 2:
|
||||
if not isinstance(active_mask, torch.Tensor) or active_mask.ndim != 2:
|
||||
raise ValueError("active_mask must be an HW tensor.")
|
||||
|
||||
active = active_mask.detach().to(device="cpu", dtype=torch.bool)
|
||||
height = int(active.shape[0])
|
||||
width = int(active.shape[1])
|
||||
visited = torch.zeros((height, width), dtype=torch.bool)
|
||||
active = active_mask.detach().to(device="cpu", dtype=torch.uint8).numpy()
|
||||
if not active.any():
|
||||
return ()
|
||||
cv2 = import_module("cv2")
|
||||
component_count, labels, stats, _centroids = cv2.connectedComponentsWithStats(
|
||||
active,
|
||||
connectivity=8,
|
||||
)
|
||||
components: list[MaskComponent] = []
|
||||
|
||||
for top in range(height):
|
||||
for left in range(width):
|
||||
if visited[top, left].item() or not active[top, left].item():
|
||||
continue
|
||||
components.append(_trace_component(active, visited, left, top))
|
||||
|
||||
return tuple(components)
|
||||
|
||||
|
||||
def _trace_component(
|
||||
active: torch.Tensor,
|
||||
visited: torch.Tensor,
|
||||
start_left: int,
|
||||
start_top: int,
|
||||
) -> MaskComponent:
|
||||
"""Trace one 8-connected component from its first active pixel."""
|
||||
|
||||
height = int(active.shape[0])
|
||||
width = int(active.shape[1])
|
||||
queue: list[tuple[int, int]] = [(start_top, start_left)]
|
||||
visited[start_top, start_left] = True
|
||||
pixels: list[tuple[int, int]] = []
|
||||
index = 0
|
||||
|
||||
while index < len(queue):
|
||||
top, left = queue[index]
|
||||
index += 1
|
||||
pixels.append((top, left))
|
||||
for neighbor_top in range(max(0, top - 1), min(height, top + 2)):
|
||||
for neighbor_left in range(max(0, left - 1), min(width, left + 2)):
|
||||
if visited[neighbor_top, neighbor_left].item():
|
||||
continue
|
||||
visited[neighbor_top, neighbor_left] = True
|
||||
if active[neighbor_top, neighbor_left].item():
|
||||
queue.append((neighbor_top, neighbor_left))
|
||||
|
||||
top_values = [top for top, _left in pixels]
|
||||
left_values = [left for _top, left in pixels]
|
||||
bbox = BoundingBox(
|
||||
min(left_values),
|
||||
min(top_values),
|
||||
max(left_values) + 1,
|
||||
max(top_values) + 1,
|
||||
for label in range(1, int(component_count)):
|
||||
left = int(stats[label, cv2.CC_STAT_LEFT])
|
||||
top = int(stats[label, cv2.CC_STAT_TOP])
|
||||
width = int(stats[label, cv2.CC_STAT_WIDTH])
|
||||
height = int(stats[label, cv2.CC_STAT_HEIGHT])
|
||||
components.append(
|
||||
MaskComponent(
|
||||
bbox=BoundingBox(left, top, left + width, top + height),
|
||||
mask=torch.from_numpy(labels == label),
|
||||
)
|
||||
)
|
||||
return tuple(
|
||||
sorted(components, key=lambda value: (value.bbox.top, value.bbox.left))
|
||||
)
|
||||
component_mask = torch.zeros((height, width), dtype=torch.bool)
|
||||
for top, left in pixels:
|
||||
component_mask[top, left] = True
|
||||
return MaskComponent(bbox=bbox, mask=component_mask)
|
||||
|
||||
@@ -12,9 +12,14 @@ from ..runtime.prompt_control_availability import prompt_control_is_available
|
||||
def get_nodes() -> list[type[object]]:
|
||||
"""Return v3 nodes that can be advertised in this environment."""
|
||||
|
||||
from .all_prompt_attention_segs import AllPromptAttentionSEGSV3
|
||||
from .attention_capture_model import AttentionCaptureModelV3
|
||||
from .attention_masked_conditioning import AttentionMaskedConditioningV3
|
||||
from .attention_region_mask import AttentionRegionMaskV3
|
||||
from .batch_region_conditioning import BatchRegionConditioningV3
|
||||
from .batch_segs import BatchSEGSV3
|
||||
from .compose_regional_conditioning import ComposeRegionalConditioningV3
|
||||
from .concept_attention_segs import ConceptAttentionSEGSV3
|
||||
from .external_llm_prompt import ExternalLLMPromptV3
|
||||
from .ksampler_attention_coupling import KSamplerAttentionCouplingV3
|
||||
from .ksampler_contextual_attention_coupling import (
|
||||
@@ -69,6 +74,10 @@ def get_nodes() -> list[type[object]]:
|
||||
from .wd14_tagger_loader import WD14TaggerLoaderV3
|
||||
|
||||
nodes: list[type[object]] = [
|
||||
AllPromptAttentionSEGSV3,
|
||||
AttentionCaptureModelV3,
|
||||
AttentionMaskedConditioningV3,
|
||||
AttentionRegionMaskV3,
|
||||
BatchRegionConditioningV3,
|
||||
BatchSEGSV3,
|
||||
ConditioningBatchAppendV3,
|
||||
@@ -103,6 +112,7 @@ def get_nodes() -> list[type[object]]:
|
||||
SAMModelLoaderV3,
|
||||
SEGSFromSAMOutputV3,
|
||||
ScaleFactorV3,
|
||||
ConceptAttentionSEGSV3,
|
||||
SeedV3,
|
||||
SimpleLoadAnimaV3,
|
||||
SimplePreviewSEGSV3,
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for exposing every mapped prompt attention region."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AllPromptAttentionSEGSV3(_ComfyNodeBase):
|
||||
"""Expose separate overlapping regions for all readable positive concepts."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "image"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the all-prompt downstream attention-SEGS contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AllPromptAttentionSEGS",
|
||||
display_name="All Prompt Attention SEGS",
|
||||
category="SimpleSyrup/Detection",
|
||||
description=(
|
||||
"Returns separate overlapping soft SEGS for every mapped concept "
|
||||
"in the positive prompt that produced the image."
|
||||
),
|
||||
search_aliases=[
|
||||
"all attention heatmaps",
|
||||
"prompt insight",
|
||||
"token regions",
|
||||
],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Image.Input(
|
||||
"image",
|
||||
tooltip=(
|
||||
"Image whose graph provenance identifies the upstream sampler; "
|
||||
"the image is returned unchanged."
|
||||
),
|
||||
),
|
||||
*attention_region_control_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Image.Output("image", tooltip="Unchanged connected image."),
|
||||
_comfy_io.SEGS.Output(
|
||||
"segs",
|
||||
tooltip="Separate labeled overlapping SEGS for prompt concepts.",
|
||||
is_output_list=True,
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask",
|
||||
tooltip="Soft union of every retained prompt attention region.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
image: object,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
split_sensitivity: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
) -> tuple[object, object, object]:
|
||||
"""Consume all prompt maps and render separate overlapping SEGS."""
|
||||
|
||||
del sampler_stage
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_image(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
image=image,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
split_sensitivity=split_sensitivity,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
),
|
||||
)
|
||||
return result.image, result.segs, result.mask
|
||||
@@ -0,0 +1,89 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Internal Comfy v3 MODEL derivation node injected by prompt provenance."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..services.attention_capture_model_service import ATTENTION_CAPTURE_MODEL_SERVICE
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AttentionCaptureModelV3(_ComfyNodeBase):
|
||||
"""Derive an observation-only MODEL from a prompt-injected capture plan."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the dev-only internal capture node contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AttentionCaptureModel",
|
||||
display_name="Attention Capture Model (Internal)",
|
||||
category="SimpleSyrup/Internal",
|
||||
description=(
|
||||
"Internal prompt-scoped MODEL observer used by downstream "
|
||||
"attention-region nodes."
|
||||
),
|
||||
is_dev_only=True,
|
||||
inputs=[
|
||||
_comfy_io.Model.Input(
|
||||
"model",
|
||||
tooltip="Upstream MODEL cloned for observation-only capture.",
|
||||
),
|
||||
_comfy_io.Conditioning.Input(
|
||||
"positive",
|
||||
tooltip="Positive conditioning whose prompt tokens are mapped.",
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"plan_json",
|
||||
multiline=False,
|
||||
tooltip="Prompt-injected capture plan for the target sampler.",
|
||||
),
|
||||
_comfy_io.Clip.Input(
|
||||
"clip",
|
||||
optional=True,
|
||||
tooltip="Graph-visible CLIP used for exact prompt token mapping.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Model.Output(
|
||||
"model",
|
||||
tooltip="MODEL carrying one observation-only attention observer.",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: object,
|
||||
positive: object,
|
||||
plan_json: str,
|
||||
clip: object | None = None,
|
||||
) -> tuple[object]:
|
||||
"""Prepare and publish capture state before the target sampler runs."""
|
||||
|
||||
del positive
|
||||
return (
|
||||
ATTENTION_CAPTURE_MODEL_SERVICE.prepare(
|
||||
model=model,
|
||||
plan_json=plan_json,
|
||||
clip=clip,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,151 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for attention-derived regional conditioning of a later sampler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.attention_concepts import parse_attention_concepts
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AttentionMaskedConditioningV3(_ComfyNodeBase):
|
||||
"""Mask supplied conditioning with attention captured from an earlier sampler."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "latent"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare later-pass attention-masked conditioning inputs and outputs."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AttentionMaskedConditioning",
|
||||
display_name="Attention Masked Conditioning",
|
||||
category="SimpleSyrup/Conditioning",
|
||||
description=(
|
||||
"Masks supplied conditioning to regions discovered from an earlier "
|
||||
"sampler so it can control a later sampling pass."
|
||||
),
|
||||
search_aliases=[
|
||||
"regional conditioning",
|
||||
"attention conditioning",
|
||||
"masked prompt",
|
||||
],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Latent.Input(
|
||||
"latent",
|
||||
tooltip=(
|
||||
"Latent produced by the sampler used for localization; it is "
|
||||
"returned unchanged for the later sampler."
|
||||
),
|
||||
),
|
||||
_comfy_io.Conditioning.Input(
|
||||
"conditioning",
|
||||
tooltip="Conditioning to apply only inside the discovered regions.",
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"concepts",
|
||||
multiline=True,
|
||||
default="subject",
|
||||
tooltip=("Concepts used to discover regions, separated by |."),
|
||||
),
|
||||
_comfy_io.Float.Input(
|
||||
"conditioning_strength",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=10.0,
|
||||
step=0.05,
|
||||
tooltip="Strength of the supplied conditioning inside the mask.",
|
||||
),
|
||||
*attention_region_control_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent", tooltip="Unchanged localization latent."
|
||||
),
|
||||
_comfy_io.Conditioning.Output(
|
||||
"conditioning",
|
||||
tooltip=(
|
||||
"Supplied conditioning carrying the attention-derived mask."
|
||||
),
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask",
|
||||
tooltip="Soft latent-resolution mask attached to the conditioning.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
latent: object,
|
||||
conditioning: object,
|
||||
concepts: str,
|
||||
conditioning_strength: float,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
split_sensitivity: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
) -> tuple[object, object, object]:
|
||||
"""Return later-pass conditioning masked by earlier-sampler attention."""
|
||||
|
||||
del sampler_stage
|
||||
if not parse_attention_concepts(concepts):
|
||||
raise ValueError("Attention Masked Conditioning requires a concept.")
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_latent(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
latent=latent,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
split_sensitivity=split_sensitivity,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
),
|
||||
)
|
||||
masked = ATTENTION_REGION_NODE_SERVICE.mask_conditioning(
|
||||
conditioning,
|
||||
result.mask,
|
||||
conditioning_strength,
|
||||
)
|
||||
return result.latent, masked, result.mask
|
||||
@@ -0,0 +1,177 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Build shared Comfy v3 inputs and domain controls for attention-region nodes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..domain.attention_region_capture import (
|
||||
AttentionCaptureProfile,
|
||||
AttentionRegionControls,
|
||||
)
|
||||
|
||||
|
||||
def attention_region_control_inputs(io: Any) -> list[object]:
|
||||
"""Return the complete shared attention-native control schema."""
|
||||
|
||||
return [
|
||||
io.Int.Input(
|
||||
"sampler_stage",
|
||||
default=1,
|
||||
min=-1,
|
||||
max=1024,
|
||||
step=1,
|
||||
tooltip=(
|
||||
"Selects the connected sampling stage: 1 is first, 2 is second, "
|
||||
"and 0 or -1 selects the last; oversized values use the last."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"capture_start",
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
max=0.99,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Start of denoising evidence to include; later values ignore more "
|
||||
"of the initial composition phase."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"capture_end",
|
||||
default=1.0,
|
||||
min=0.01,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"End of denoising evidence to include; earlier values ignore more "
|
||||
"late refinement attention."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"minimum_strength",
|
||||
default=0.15,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Minimum normalized attention association retained in a region; "
|
||||
"higher values narrow the silhouette toward its semantic core."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"minimum_consensus",
|
||||
default=0.25,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Fraction of selected observations that must support a pixel; "
|
||||
"higher values keep more persistent regions."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"split_sensitivity",
|
||||
default=0.35,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Sensitivity to separate peaks into instances without shrinking "
|
||||
"their combined silhouette; higher values cut stronger bridges."
|
||||
),
|
||||
),
|
||||
io.Int.Input(
|
||||
"minimum_region_size",
|
||||
default=512,
|
||||
min=1,
|
||||
max=1048576,
|
||||
step=1,
|
||||
tooltip="Discard attention components smaller than this many pixels.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"keep_only",
|
||||
default=1,
|
||||
min=0,
|
||||
max=1024,
|
||||
step=1,
|
||||
tooltip=(
|
||||
"Keep the best N instances per concept; 1 keeps the largest and "
|
||||
"0 keeps all."
|
||||
),
|
||||
),
|
||||
io.Combo.Input(
|
||||
"keep_by",
|
||||
options=["largest size", "highest confidence"],
|
||||
default="largest size",
|
||||
tooltip="Ranks retained instances by area or attention confidence.",
|
||||
),
|
||||
io.Boolean.Input(
|
||||
"combine_segs",
|
||||
default=False,
|
||||
tooltip="Combines retained instances of each concept into one SEG.",
|
||||
),
|
||||
io.Float.Input(
|
||||
"matte_solidity",
|
||||
default=0.75,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Higher values flatten accepted interiors toward fully opaque alpha."
|
||||
),
|
||||
),
|
||||
io.Int.Input(
|
||||
"edge_feather",
|
||||
default=8,
|
||||
min=0,
|
||||
max=4096,
|
||||
step=1,
|
||||
tooltip="Width in output pixels of the matte boundary transition.",
|
||||
),
|
||||
io.Combo.Input(
|
||||
"capture_profile",
|
||||
options=[profile.value for profile in AttentionCaptureProfile],
|
||||
default=AttentionCaptureProfile.FAST.value,
|
||||
tooltip=(
|
||||
"Controls observation density: fast minimizes overhead, balanced "
|
||||
"adds temporal evidence, and exhaustive retains every eligible call."
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def attention_region_controls(
|
||||
*,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
split_sensitivity: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
) -> AttentionRegionControls:
|
||||
"""Build validated domain controls from public node inputs."""
|
||||
|
||||
return AttentionRegionControls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
split_sensitivity=split_sensitivity,
|
||||
minimum_region_size=minimum_region_size,
|
||||
profile=AttentionCaptureProfile(capture_profile),
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
)
|
||||
@@ -0,0 +1,124 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for querying upstream attention as a latent-resolution mask."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.attention_concepts import parse_attention_concepts
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AttentionRegionMaskV3(_ComfyNodeBase):
|
||||
"""Return a reusable mask derived from an upstream sampler's attention."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "latent"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the latent attention-region mask contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AttentionRegionMask",
|
||||
display_name="Attention Region Mask",
|
||||
category="SimpleSyrup/Masking",
|
||||
description=(
|
||||
"Returns a reusable latent-resolution mask for concepts attended "
|
||||
"by the sampler that produced the connected latent."
|
||||
),
|
||||
search_aliases=["latent attention mask", "prompt mask", "region mask"],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Latent.Input(
|
||||
"latent",
|
||||
tooltip=(
|
||||
"Latent whose graph provenance identifies the upstream "
|
||||
"sampler; the latent is returned unchanged."
|
||||
),
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"concepts",
|
||||
multiline=True,
|
||||
default="subject",
|
||||
tooltip=(
|
||||
"Concepts whose attention becomes the mask, separated by |."
|
||||
),
|
||||
),
|
||||
*attention_region_control_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent", tooltip="Unchanged connected latent."
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask",
|
||||
tooltip="Soft latent-resolution union of matched regions.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
latent: object,
|
||||
concepts: str,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
split_sensitivity: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
) -> tuple[object, object]:
|
||||
"""Consume matching maps and return a reusable latent-resolution mask."""
|
||||
|
||||
del sampler_stage
|
||||
if not parse_attention_concepts(concepts):
|
||||
raise ValueError("Attention Region Mask requires at least one concept.")
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_latent(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
latent=latent,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
split_sensitivity=split_sensitivity,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
),
|
||||
)
|
||||
return result.latent, result.mask
|
||||
@@ -0,0 +1,129 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Expose explicitly requested upstream attention concepts as SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.attention_concepts import parse_attention_concepts
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class ConceptAttentionSEGSV3(_ComfyNodeBase):
|
||||
"""Return attention-derived instances for explicitly named concepts."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "image"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the downstream concept-attention SEGS contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.ConceptAttentionSEGS",
|
||||
display_name="Concept Attention SEGS",
|
||||
category="SimpleSyrup/Detection",
|
||||
description=(
|
||||
"Returns regions associated with explicit concepts from a selected "
|
||||
"sampler in the connected image's provenance."
|
||||
),
|
||||
search_aliases=[
|
||||
"attention heatmap",
|
||||
"prompt segmentation",
|
||||
"attention mask",
|
||||
],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Image.Input(
|
||||
"image",
|
||||
tooltip=(
|
||||
"Image whose graph provenance identifies the sampling chain; "
|
||||
"the image is returned unchanged."
|
||||
),
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"concepts",
|
||||
multiline=True,
|
||||
default="subject",
|
||||
tooltip=(
|
||||
"Enter one or more concepts separated by |, such as "
|
||||
"girl | pink hair | cat."
|
||||
),
|
||||
),
|
||||
*attention_region_control_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Image.Output("image", tooltip="Unchanged connected image."),
|
||||
_comfy_io.SEGS.Output(
|
||||
"segs",
|
||||
tooltip="Labeled attention-derived instances for the concepts.",
|
||||
is_output_list=True,
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask", tooltip="Union of all retained concept instances."
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
image: object,
|
||||
concepts: str,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
split_sensitivity: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
) -> tuple[object, object, object]:
|
||||
"""Consume shared capture evidence and render requested concept SEGS."""
|
||||
|
||||
del sampler_stage
|
||||
if not parse_attention_concepts(concepts):
|
||||
raise ValueError("Concept Attention SEGS requires at least one concept.")
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_image(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
image=image,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
split_sensitivity=split_sensitivity,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
),
|
||||
)
|
||||
return result.image, result.segs, result.mask
|
||||
@@ -0,0 +1,262 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Calculate and materialize compact selected-token attention affinities."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.attention_region_maps import AttentionTokenSpan, CapturedAttentionMap
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PendingAttentionMap:
|
||||
"""Retain one compact device-local map until sampling has completed."""
|
||||
|
||||
label: str
|
||||
values: torch.Tensor
|
||||
progress: float
|
||||
layer_key: str
|
||||
batch_index: int
|
||||
spatial_height: int
|
||||
spatial_width: int
|
||||
|
||||
|
||||
class AttentionAffinityCalculator:
|
||||
"""Own selected-key probability estimation and deferred device transfer."""
|
||||
|
||||
def head_tensors(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
heads: int,
|
||||
skip_reshape: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Normalize Comfy's standard and pre-shaped attention tensor layouts."""
|
||||
|
||||
if type(heads) is not int or heads < 1:
|
||||
raise ValueError("Attention capture head count must be positive.")
|
||||
if skip_reshape and query.ndim == 4 and key.ndim == 4:
|
||||
return query, key
|
||||
if query.ndim != 3 or key.ndim != 3:
|
||||
raise ValueError("Attention capture received unsupported Q/K tensor ranks.")
|
||||
if int(query.shape[-1]) % heads or int(key.shape[-1]) % heads:
|
||||
raise ValueError("Attention Q/K channels must divide evenly across heads.")
|
||||
q = query.view(query.shape[0], query.shape[1], heads, -1).permute(0, 2, 1, 3)
|
||||
k = key.view(key.shape[0], key.shape[1], heads, -1).permute(0, 2, 1, 3)
|
||||
return q, k
|
||||
|
||||
def positive_rows(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
raw_branches: object,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Select positive CFG rows while retaining ordinary batch members."""
|
||||
|
||||
if not isinstance(raw_branches, list) or not raw_branches:
|
||||
return query, key
|
||||
if any(type(branch) is not int for branch in raw_branches):
|
||||
raise TypeError("Attention cond_or_uncond entries must be integers.")
|
||||
if int(query.shape[0]) % len(raw_branches):
|
||||
raise ValueError("Attention batch cannot be partitioned by cond_or_uncond.")
|
||||
rows_per_branch = int(query.shape[0]) // len(raw_branches)
|
||||
row_indices = [
|
||||
index
|
||||
for branch_index, branch in enumerate(raw_branches)
|
||||
if branch == 0
|
||||
for index in range(
|
||||
branch_index * rows_per_branch,
|
||||
(branch_index + 1) * rows_per_branch,
|
||||
)
|
||||
]
|
||||
if not row_indices:
|
||||
return query[:0], key[:0]
|
||||
indices = torch.tensor(row_indices, device=query.device, dtype=torch.int64)
|
||||
return query.index_select(0, indices), key.index_select(0, indices)
|
||||
|
||||
def log_denominator(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
token_budget: int,
|
||||
) -> torch.Tensor:
|
||||
"""Estimate softmax normalization from a deterministic bounded key sample."""
|
||||
|
||||
scale = math.sqrt(float(query.shape[-1]))
|
||||
context_tokens = int(key.shape[-2])
|
||||
sample_count = min(context_tokens, token_budget)
|
||||
if sample_count < context_tokens:
|
||||
indices = (
|
||||
torch.linspace(
|
||||
0,
|
||||
context_tokens - 1,
|
||||
sample_count,
|
||||
device=key.device,
|
||||
)
|
||||
.round()
|
||||
.to(dtype=torch.int64)
|
||||
)
|
||||
key = key.index_select(-2, indices)
|
||||
denominator: torch.Tensor | None = None
|
||||
for key_chunk in torch.split(key.float(), 32, dim=-2):
|
||||
logits = torch.einsum("bhqd,bhkd->bhqk", query.float(), key_chunk) / scale
|
||||
chunk_denominator = torch.logsumexp(logits, dim=-1)
|
||||
denominator = (
|
||||
chunk_denominator
|
||||
if denominator is None
|
||||
else torch.logaddexp(denominator, chunk_denominator)
|
||||
)
|
||||
if denominator is None:
|
||||
raise ValueError("Attention context cannot be empty.")
|
||||
correction = math.log(float(context_tokens) / float(sample_count))
|
||||
return denominator + correction
|
||||
|
||||
def capture_spans(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
spans: tuple[AttentionTokenSpan, ...],
|
||||
progress: float,
|
||||
layer_key: str,
|
||||
denominator: torch.Tensor,
|
||||
spatial_height: int,
|
||||
spatial_width: int,
|
||||
) -> tuple[PendingAttentionMap, ...]:
|
||||
"""Compute all native span maps from one unioned selected-key projection."""
|
||||
|
||||
union_indices = tuple(
|
||||
sorted({index for span in spans for index in span.token_indices})
|
||||
)
|
||||
if not union_indices:
|
||||
return ()
|
||||
indices = torch.tensor(union_indices, device=key.device, dtype=torch.int64)
|
||||
selected = key.index_select(-2, indices)
|
||||
logits = torch.einsum("bhqd,bhkd->bhqk", query.float(), selected.float())
|
||||
logits = logits / math.sqrt(float(query.shape[-1]))
|
||||
log_probability = (logits - denominator.unsqueeze(-1)).clamp_max(0.0)
|
||||
probability = torch.exp(log_probability)
|
||||
union_positions = {
|
||||
token_index: index for index, token_index in enumerate(union_indices)
|
||||
}
|
||||
pending: list[PendingAttentionMap] = []
|
||||
for span in spans:
|
||||
positions = torch.tensor(
|
||||
tuple(union_positions[index] for index in span.token_indices),
|
||||
device=probability.device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
compact_values = (
|
||||
probability.index_select(-1, positions)
|
||||
.mean(dim=(1, 3))
|
||||
.detach()
|
||||
.to(dtype=torch.float16)
|
||||
)
|
||||
pending.extend(
|
||||
PendingAttentionMap(
|
||||
span.display_label,
|
||||
values,
|
||||
progress,
|
||||
layer_key,
|
||||
batch_index,
|
||||
spatial_height,
|
||||
spatial_width,
|
||||
)
|
||||
for batch_index, values in enumerate(compact_values)
|
||||
)
|
||||
return tuple(pending)
|
||||
|
||||
def capture_open_vocabulary(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
projected_key: torch.Tensor,
|
||||
heads: int,
|
||||
label: str,
|
||||
progress: float,
|
||||
layer_key: str,
|
||||
native_key: torch.Tensor,
|
||||
denominator_token_budget: int,
|
||||
spatial_height: int,
|
||||
spatial_width: int,
|
||||
) -> tuple[PendingAttentionMap, ...]:
|
||||
"""Normalize one layer-specific OVAM key against native projected queries."""
|
||||
|
||||
_unused_query, key_heads = self.head_tensors(
|
||||
projected_key,
|
||||
projected_key,
|
||||
heads,
|
||||
False,
|
||||
)
|
||||
del _unused_query
|
||||
if int(key_heads.shape[0]) == 1 and int(query.shape[0]) > 1:
|
||||
key_heads = key_heads.expand(int(query.shape[0]), -1, -1, -1)
|
||||
combined_key = torch.cat((native_key, key_heads), dim=-2)
|
||||
denominator = self.log_denominator(
|
||||
query,
|
||||
combined_key,
|
||||
denominator_token_budget,
|
||||
)
|
||||
span = AttentionTokenSpan(label, 1, tuple(range(int(key_heads.shape[-2]))))
|
||||
return self.capture_spans(
|
||||
query,
|
||||
key_heads,
|
||||
(span,),
|
||||
progress,
|
||||
layer_key,
|
||||
denominator,
|
||||
spatial_height,
|
||||
spatial_width,
|
||||
)
|
||||
|
||||
def materialize(
|
||||
self,
|
||||
pending: tuple[PendingAttentionMap, ...],
|
||||
) -> tuple[CapturedAttentionMap, ...]:
|
||||
"""Transfer compatible maps together and calculate concentration on CPU."""
|
||||
|
||||
grouped: dict[
|
||||
tuple[torch.device, torch.dtype, tuple[int, ...]],
|
||||
list[tuple[int, PendingAttentionMap]],
|
||||
] = defaultdict(list)
|
||||
for index, attention_map in enumerate(pending):
|
||||
key = (
|
||||
attention_map.values.device,
|
||||
attention_map.values.dtype,
|
||||
tuple(attention_map.values.shape),
|
||||
)
|
||||
grouped[key].append((index, attention_map))
|
||||
|
||||
ordered: list[CapturedAttentionMap | None] = [None] * len(pending)
|
||||
for group in grouped.values():
|
||||
cpu_values = torch.stack(tuple(value.values for _index, value in group)).to(
|
||||
device="cpu", dtype=torch.float16
|
||||
)
|
||||
maxima = cpu_values.float().amax(dim=1)
|
||||
means = cpu_values.float().mean(dim=1)
|
||||
peak_ratios = maxima / means.clamp_min(1e-12)
|
||||
confidences = (
|
||||
1.0 - torch.exp(-(peak_ratios - 1.0).clamp_min(0.0) / 8.0)
|
||||
).clamp(0.0, 1.0)
|
||||
for group_index, (original_index, attention_map) in enumerate(group):
|
||||
ordered[original_index] = CapturedAttentionMap(
|
||||
attention_map.label,
|
||||
cpu_values[group_index],
|
||||
attention_map.progress,
|
||||
attention_map.layer_key,
|
||||
attention_map.batch_index,
|
||||
float(confidences[group_index].item()),
|
||||
attention_map.spatial_height,
|
||||
attention_map.spatial_width,
|
||||
)
|
||||
if any(value is None for value in ordered):
|
||||
raise RuntimeError("Attention map materialization lost an observation.")
|
||||
return tuple(value for value in ordered if value is not None)
|
||||
|
||||
|
||||
ATTENTION_AFFINITY_CALCULATOR = AttentionAffinityCalculator()
|
||||
@@ -0,0 +1,403 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Observe selected cross-attention affinities without replacing denoising."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, replace
|
||||
from threading import RLock
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.attention_geometry import factor_spatial_geometry
|
||||
from ..domain.attention_region_capture import (
|
||||
AttentionCapturePlan,
|
||||
AttentionCaptureProfile,
|
||||
AttentionRegionRequest,
|
||||
)
|
||||
from ..domain.attention_region_maps import (
|
||||
AttentionTokenCatalog,
|
||||
AttentionTokenSpan,
|
||||
CapturedAttentionMap,
|
||||
OpenVocabularyContext,
|
||||
)
|
||||
from ..domain.regional_model_capabilities import RegionalModelFamily
|
||||
from .attention_coupling.anima_context import ANIMA_CONTEXT_SEQUENCE_LENGTH
|
||||
from .attention_region_affinity import (
|
||||
ATTENTION_AFFINITY_CALCULATOR,
|
||||
PendingAttentionMap,
|
||||
)
|
||||
from .denoising_progress import DENOISING_PROGRESS_RESOLVER
|
||||
from .regional_attention_model_call_values import uniform_model_call_sigma
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RequestTargets:
|
||||
"""Bind one request to exact prompt spans selected for capture."""
|
||||
|
||||
request: AttentionRegionRequest
|
||||
spans: tuple[AttentionTokenSpan, ...]
|
||||
|
||||
|
||||
class AttentionRegionCaptureSession:
|
||||
"""Aggregate compact selected-token affinity maps for one sampler run."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
plan: AttentionCapturePlan,
|
||||
model_family: RegionalModelFamily,
|
||||
token_catalog: AttentionTokenCatalog,
|
||||
request_spans: Mapping[str, tuple[AttentionTokenSpan, ...]],
|
||||
open_vocabulary_contexts: tuple[OpenVocabularyContext, ...] = (),
|
||||
) -> None:
|
||||
"""Resolve request targets and initialize prompt-scoped observations."""
|
||||
|
||||
self.plan = plan
|
||||
self.model_family = model_family
|
||||
self.token_catalog = token_catalog
|
||||
self.open_vocabulary_contexts = open_vocabulary_contexts
|
||||
self._targets = tuple(
|
||||
_RequestTargets(request, request_spans.get(request.node_id, ()))
|
||||
for request in plan.requests
|
||||
)
|
||||
self._pending_maps: list[PendingAttentionMap] = []
|
||||
self._materialized_maps: tuple[CapturedAttentionMap, ...] | None = None
|
||||
self._pending_open_keys: tuple[tuple[str, torch.Tensor], ...] = ()
|
||||
self._call_index = 0
|
||||
self._geometry_logged = False
|
||||
self._geometry_unavailable = False
|
||||
self._lock = RLock()
|
||||
profiles = tuple(target.request.controls.profile for target in self._targets)
|
||||
self._stride = min(_profile_stride(profile) for profile in profiles)
|
||||
self._denominator_token_budget = max(
|
||||
_profile_denominator_tokens(profile) for profile in profiles
|
||||
)
|
||||
self._capture_start = min(
|
||||
target.request.controls.capture_start for target in self._targets
|
||||
)
|
||||
self._capture_end = max(
|
||||
target.request.controls.capture_end for target in self._targets
|
||||
)
|
||||
|
||||
@property
|
||||
def status_message(self) -> str:
|
||||
"""Return a concise model-family capture status for public nodes."""
|
||||
|
||||
if self._geometry_unavailable:
|
||||
return (
|
||||
"No attention maps: the selected Anima sampler input dimensions "
|
||||
"are not graph-visible"
|
||||
)
|
||||
return f"Captured {self.model_family.value} attention"
|
||||
|
||||
def observe(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
heads: int,
|
||||
transformer_options: Mapping[str, object],
|
||||
*,
|
||||
skip_reshape: bool,
|
||||
) -> None:
|
||||
"""Capture selected positive-branch affinities for one attention call."""
|
||||
|
||||
if not any(target.spans for target in self._targets) and not (
|
||||
self.open_vocabulary_contexts
|
||||
):
|
||||
return
|
||||
if not _has_expected_cross_attention_context(
|
||||
key,
|
||||
self.model_family,
|
||||
self.token_catalog.sequence_length,
|
||||
skip_reshape,
|
||||
):
|
||||
return
|
||||
open_keys = self._take_open_vocabulary_keys()
|
||||
progress = _progress(transformer_options)
|
||||
if progress < self._capture_start or progress > self._capture_end:
|
||||
return
|
||||
with self._lock:
|
||||
call_index = self._call_index
|
||||
self._call_index += 1
|
||||
if call_index % self._stride != 0:
|
||||
return
|
||||
if not self._has_resolvable_geometry(transformer_options):
|
||||
return
|
||||
q_heads, k_heads = ATTENTION_AFFINITY_CALCULATOR.head_tensors(
|
||||
query, key, heads, skip_reshape
|
||||
)
|
||||
q_positive, k_positive = ATTENTION_AFFINITY_CALCULATOR.positive_rows(
|
||||
q_heads,
|
||||
k_heads,
|
||||
transformer_options.get("cond_or_uncond"),
|
||||
)
|
||||
layer_key = _layer_key(transformer_options)
|
||||
spans = _unique_spans(self._targets)
|
||||
denominator = ATTENTION_AFFINITY_CALCULATOR.log_denominator(
|
||||
q_positive,
|
||||
k_positive,
|
||||
self._denominator_token_budget,
|
||||
)
|
||||
spatial_height, spatial_width = _spatial_geometry(
|
||||
int(q_positive.shape[-2]),
|
||||
transformer_options,
|
||||
fallback_aspect=self.plan.source_aspect,
|
||||
)
|
||||
self._log_geometry_once(
|
||||
token_count=int(q_positive.shape[-2]),
|
||||
spatial_height=spatial_height,
|
||||
spatial_width=spatial_width,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
captured = ATTENTION_AFFINITY_CALCULATOR.capture_spans(
|
||||
q_positive,
|
||||
k_positive,
|
||||
spans,
|
||||
progress,
|
||||
layer_key,
|
||||
denominator,
|
||||
spatial_height,
|
||||
spatial_width,
|
||||
)
|
||||
open_captured = tuple(
|
||||
captured_map
|
||||
for label, projected_key in open_keys
|
||||
for captured_map in ATTENTION_AFFINITY_CALCULATOR.capture_open_vocabulary(
|
||||
q_positive,
|
||||
projected_key,
|
||||
heads,
|
||||
label,
|
||||
progress,
|
||||
layer_key,
|
||||
k_positive,
|
||||
self._denominator_token_budget,
|
||||
spatial_height,
|
||||
spatial_width,
|
||||
)
|
||||
)
|
||||
with self._lock:
|
||||
self._pending_maps.extend((*captured, *open_captured))
|
||||
|
||||
def _has_resolvable_geometry(
|
||||
self,
|
||||
transformer_options: Mapping[str, object],
|
||||
) -> bool:
|
||||
"""Fail closed when an Anima spatial grid has no trustworthy orientation."""
|
||||
|
||||
if self.model_family is not RegionalModelFamily.ANIMA:
|
||||
return True
|
||||
if (
|
||||
self.plan.source_aspect is not None
|
||||
or _runtime_aspect(transformer_options) is not None
|
||||
):
|
||||
return True
|
||||
with self._lock:
|
||||
first_failure = not self._geometry_unavailable
|
||||
self._geometry_unavailable = True
|
||||
if first_failure:
|
||||
LOGGER.warning(
|
||||
"Skipped Anima attention capture because the selected sampler "
|
||||
"input dimensions are not graph-visible",
|
||||
extra={"sampler_node_id": self.plan.sampler_node_id},
|
||||
)
|
||||
return False
|
||||
|
||||
def _log_geometry_once(
|
||||
self,
|
||||
*,
|
||||
token_count: int,
|
||||
spatial_height: int,
|
||||
spatial_width: int,
|
||||
transformer_options: Mapping[str, object],
|
||||
) -> None:
|
||||
"""Log one representative geometry decision for runtime diagnosis."""
|
||||
|
||||
with self._lock:
|
||||
if self._geometry_logged:
|
||||
return
|
||||
self._geometry_logged = True
|
||||
original_shape = transformer_options.get("original_shape")
|
||||
activations_shape = transformer_options.get("activations_shape")
|
||||
LOGGER.info(
|
||||
"Resolved %s attention capture geometry: %d tokens -> %dx%d; "
|
||||
"original_shape=%r activations_shape=%r",
|
||||
self.model_family.value,
|
||||
token_count,
|
||||
spatial_height,
|
||||
spatial_width,
|
||||
original_shape,
|
||||
activations_shape,
|
||||
extra={
|
||||
"model_family": self.model_family.value,
|
||||
"spatial_token_count": token_count,
|
||||
"spatial_height": spatial_height,
|
||||
"spatial_width": spatial_width,
|
||||
"original_shape": original_shape,
|
||||
"activations_shape": activations_shape,
|
||||
},
|
||||
)
|
||||
|
||||
def stage_open_vocabulary_keys(
|
||||
self,
|
||||
keys: tuple[tuple[str, torch.Tensor], ...],
|
||||
) -> None:
|
||||
"""Stage keys produced by the current layer's clone-local wrapper."""
|
||||
|
||||
with self._lock:
|
||||
self._pending_open_keys = keys
|
||||
|
||||
def _take_open_vocabulary_keys(self) -> tuple[tuple[str, torch.Tensor], ...]:
|
||||
"""Consume only keys paired with the immediately following attention call."""
|
||||
|
||||
with self._lock:
|
||||
keys = self._pending_open_keys
|
||||
self._pending_open_keys = ()
|
||||
return keys
|
||||
|
||||
def maps_for(self, request_node_id: str) -> tuple[CapturedAttentionMap, ...]:
|
||||
"""Return immutable maps whose labels belong to one planned request."""
|
||||
|
||||
target = next(
|
||||
(
|
||||
candidate
|
||||
for candidate in self._targets
|
||||
if candidate.request.node_id == request_node_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if target is None:
|
||||
raise KeyError(f"Unknown attention-region request node: {request_node_id}.")
|
||||
labels = {span.display_label for span in target.spans}
|
||||
labels.update(
|
||||
context.label
|
||||
for context in self.open_vocabulary_contexts
|
||||
if context.label in target.request.queries
|
||||
)
|
||||
maps = self._materialize_maps()
|
||||
return tuple(
|
||||
replace(value, spatial_transforms=target.request.spatial_transforms)
|
||||
for value in maps
|
||||
if value.label in labels
|
||||
)
|
||||
|
||||
def _materialize_maps(self) -> tuple[CapturedAttentionMap, ...]:
|
||||
"""Batch compact device transfers once after denoising has completed."""
|
||||
|
||||
with self._lock:
|
||||
if self._materialized_maps is not None:
|
||||
return self._materialized_maps
|
||||
pending = tuple(self._pending_maps)
|
||||
self._pending_maps.clear()
|
||||
materialized = ATTENTION_AFFINITY_CALCULATOR.materialize(pending)
|
||||
with self._lock:
|
||||
self._materialized_maps = materialized
|
||||
return materialized
|
||||
|
||||
|
||||
def _profile_stride(profile: AttentionCaptureProfile) -> int:
|
||||
"""Return deterministic attention-call subsampling for one capture profile."""
|
||||
|
||||
if profile is AttentionCaptureProfile.FAST:
|
||||
return 16
|
||||
if profile is AttentionCaptureProfile.BALANCED:
|
||||
return 4
|
||||
return 1
|
||||
|
||||
|
||||
def _profile_denominator_tokens(profile: AttentionCaptureProfile) -> int:
|
||||
"""Bound probability calibration work according to the selected profile."""
|
||||
|
||||
if profile is AttentionCaptureProfile.FAST:
|
||||
return 8
|
||||
if profile is AttentionCaptureProfile.BALANCED:
|
||||
return 16
|
||||
return 512
|
||||
|
||||
|
||||
def _progress(options: Mapping[str, object]) -> float:
|
||||
"""Resolve exact denoising progress from Comfy's sampler metadata."""
|
||||
|
||||
sample_sigmas = options.get("sample_sigmas")
|
||||
current_sigmas = options.get("sigmas")
|
||||
if not isinstance(sample_sigmas, torch.Tensor):
|
||||
raise TypeError("Attention capture requires tensor sample_sigmas.")
|
||||
if not isinstance(current_sigmas, torch.Tensor):
|
||||
raise TypeError("Attention capture requires tensor sigmas.")
|
||||
return DENOISING_PROGRESS_RESOLVER.resolve(
|
||||
sample_sigmas,
|
||||
uniform_model_call_sigma(current_sigmas),
|
||||
)
|
||||
|
||||
|
||||
def _unique_spans(
|
||||
targets: tuple[_RequestTargets, ...],
|
||||
) -> tuple[AttentionTokenSpan, ...]:
|
||||
"""Deduplicate coalesced target spans while preserving prompt order."""
|
||||
|
||||
return tuple(dict.fromkeys(span for target in targets for span in target.spans))
|
||||
|
||||
|
||||
def _layer_key(options: Mapping[str, object]) -> str:
|
||||
"""Return a stable diagnostic key from Comfy transformer block metadata."""
|
||||
|
||||
block = options.get("block", "unknown")
|
||||
block_index = options.get("block_index", 0)
|
||||
transformer_index = options.get("transformer_index", 0)
|
||||
return f"{block!r}:{block_index!r}:{transformer_index!r}"
|
||||
|
||||
|
||||
def _spatial_geometry(
|
||||
token_count: int,
|
||||
options: Mapping[str, object],
|
||||
*,
|
||||
fallback_aspect: float | None,
|
||||
) -> tuple[int, int]:
|
||||
"""Resolve the captured attention grid from Comfy's source tensor shape."""
|
||||
|
||||
target_aspect = _runtime_aspect(options) or fallback_aspect or 1.0
|
||||
return factor_spatial_geometry(token_count, target_aspect=target_aspect)
|
||||
|
||||
|
||||
def _runtime_aspect(options: Mapping[str, object]) -> float | None:
|
||||
"""Return a trustworthy aspect from installed transformer metadata."""
|
||||
|
||||
original_shape = options.get("original_shape")
|
||||
if not isinstance(original_shape, (tuple, list)) or len(original_shape) < 2:
|
||||
return None
|
||||
raw_height, raw_width = original_shape[-2:]
|
||||
if (
|
||||
type(raw_height) is not int
|
||||
or raw_height < 1
|
||||
or type(raw_width) is not int
|
||||
or raw_width < 1
|
||||
):
|
||||
return None
|
||||
return raw_width / raw_height
|
||||
|
||||
|
||||
def _has_expected_cross_attention_context(
|
||||
key: torch.Tensor,
|
||||
model_family: RegionalModelFamily,
|
||||
mapped_token_count: int,
|
||||
skip_reshape: bool,
|
||||
) -> bool:
|
||||
"""Reject self-attention calls before capture cadence and affinity work."""
|
||||
|
||||
if key.ndim not in (3, 4):
|
||||
return False
|
||||
context_tokens = int(
|
||||
key.shape[-2] if skip_reshape and key.ndim == 4 else key.shape[1]
|
||||
)
|
||||
expected = (
|
||||
ANIMA_CONTEXT_SEQUENCE_LENGTH
|
||||
if model_family is RegionalModelFamily.ANIMA
|
||||
else mapped_token_count
|
||||
)
|
||||
return context_tokens == expected
|
||||
@@ -0,0 +1,116 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Install composable attention observation on a clone-local Comfy MODEL."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from .attention_region_capture import AttentionRegionCaptureSession
|
||||
from .attention_region_open_vocabulary import OpenVocabularyProjectionMutation
|
||||
from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
|
||||
|
||||
AttentionFunction = Callable[..., torch.Tensor]
|
||||
|
||||
|
||||
class OptimizedAttentionCaptureOverride:
|
||||
"""Compose one observer around Comfy's configured optimized attention call."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
session: AttentionRegionCaptureSession,
|
||||
previous_override: Callable[..., torch.Tensor] | None,
|
||||
) -> None:
|
||||
"""Store capture state and any upstream optimized-attention owner."""
|
||||
|
||||
self._session = session
|
||||
self._previous_override = previous_override
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
original: AttentionFunction,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
heads: int,
|
||||
*args: object,
|
||||
**kwargs: object,
|
||||
) -> torch.Tensor:
|
||||
"""Observe projected Q/K, then delegate output computation unchanged."""
|
||||
|
||||
transformer_options = kwargs.get("transformer_options")
|
||||
if isinstance(transformer_options, Mapping):
|
||||
self._session.observe(
|
||||
query,
|
||||
key,
|
||||
heads,
|
||||
transformer_options,
|
||||
skip_reshape=bool(kwargs.get("skip_reshape", False)),
|
||||
)
|
||||
delegate = self._previous_override or original
|
||||
return (
|
||||
delegate(original, query, key, value, heads, *args, **kwargs)
|
||||
if self._previous_override
|
||||
else delegate(query, key, value, heads, *args, **kwargs)
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionCaptureOverrideMutation:
|
||||
"""Install a composable optimized-attention observer on a cloned MODEL."""
|
||||
|
||||
session: AttentionRegionCaptureSession
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Replace only the clone's override option while preserving its predecessor."""
|
||||
|
||||
model_options = getattr(model, "model_options", None)
|
||||
if not isinstance(model_options, dict):
|
||||
raise TypeError("Attention capture MODEL options must be a dictionary.")
|
||||
transformer_options = model_options.get("transformer_options")
|
||||
if not isinstance(transformer_options, dict):
|
||||
raise TypeError(
|
||||
"Attention capture transformer options must be a dictionary."
|
||||
)
|
||||
previous = transformer_options.get("optimized_attention_override")
|
||||
if previous is not None and not callable(previous):
|
||||
raise TypeError("Existing optimized attention override must be callable.")
|
||||
transformer_options["optimized_attention_override"] = (
|
||||
OptimizedAttentionCaptureOverride(self.session, previous)
|
||||
)
|
||||
|
||||
|
||||
class AttentionRegionCaptureBackend:
|
||||
"""Derive one collision-safe observation-only MODEL for a capture session."""
|
||||
|
||||
def derive(
|
||||
self,
|
||||
model: object,
|
||||
session: AttentionRegionCaptureSession,
|
||||
) -> object:
|
||||
"""Clone MODEL state and install exactly one optimized-attention observer."""
|
||||
|
||||
mutations: tuple[ModelMutation, ...] = (
|
||||
AttentionCaptureOverrideMutation(session),
|
||||
)
|
||||
if session.open_vocabulary_contexts:
|
||||
mutations = (
|
||||
*mutations,
|
||||
OpenVocabularyProjectionMutation(
|
||||
session.open_vocabulary_contexts,
|
||||
session,
|
||||
),
|
||||
)
|
||||
return PATCHER_LIFECYCLE.derive_model(
|
||||
model,
|
||||
mutations,
|
||||
operation="attention-region capture",
|
||||
)
|
||||
|
||||
|
||||
ATTENTION_REGION_CAPTURE_BACKEND = AttentionRegionCaptureBackend()
|
||||
@@ -0,0 +1,306 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Plan and inject prompt-scoped attention capture from downstream graph intent."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from collections.abc import Mapping, MutableMapping
|
||||
from dataclasses import asdict, replace
|
||||
from typing import Any
|
||||
|
||||
from ..domain.attention_concepts import parse_attention_concepts
|
||||
from ..domain.attention_region_capture import (
|
||||
AttentionCapturePlan,
|
||||
AttentionCaptureProfile,
|
||||
AttentionRegionControls,
|
||||
AttentionRegionRequest,
|
||||
AttentionRegionRequestKind,
|
||||
)
|
||||
from ..domain.attention_sampler_lineage import AttentionSamplerStage
|
||||
from ..domain.graph_provenance import BrokenProvenance, GraphLink
|
||||
from .attention_sampler_lineage import ATTENTION_SAMPLER_LINEAGE_RESOLVER
|
||||
from .comfy_graph_provenance import (
|
||||
NodeRegistry,
|
||||
parse_graph_link,
|
||||
)
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
INTERNAL_CAPTURE_NODE_ID = "SimpleSyrup.AttentionCaptureModel"
|
||||
REQUEST_NODE_KINDS: dict[str, AttentionRegionRequestKind] = {
|
||||
"SimpleSyrup.ConceptAttentionSEGS": AttentionRegionRequestKind.CONCEPT_SEGS,
|
||||
"SimpleSyrup.AllPromptAttentionSEGS": AttentionRegionRequestKind.ALL_PROMPT_SEGS,
|
||||
"SimpleSyrup.AttentionRegionMask": AttentionRegionRequestKind.REGION_MASK,
|
||||
"SimpleSyrup.AttentionMaskedConditioning": (
|
||||
AttentionRegionRequestKind.MASKED_CONDITIONING
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class AttentionRegionPromptPlanner:
|
||||
"""Resolve public attention nodes into one immutable plan per sampler."""
|
||||
|
||||
def build(
|
||||
self,
|
||||
prompt: Mapping[str, Any],
|
||||
node_registry: NodeRegistry,
|
||||
) -> tuple[AttentionCapturePlan, ...]:
|
||||
"""Return ordered coalesced plans for every recoverable request node."""
|
||||
|
||||
grouped: dict[
|
||||
str, list[tuple[AttentionRegionRequest, AttentionSamplerStage]]
|
||||
] = defaultdict(list)
|
||||
for node_id in sorted(prompt):
|
||||
node = _node(prompt, node_id)
|
||||
if node is None:
|
||||
continue
|
||||
kind = REQUEST_NODE_KINDS.get(_class_type(node))
|
||||
if kind is None:
|
||||
continue
|
||||
request = _request(node_id, kind, node)
|
||||
source = self._resolve_source(prompt, node, node_registry)
|
||||
if isinstance(source, BrokenProvenance):
|
||||
LOGGER.warning(
|
||||
"Attention-region graph request cannot recover a sampler",
|
||||
extra={"request_node_id": node_id, "reason": source.reason},
|
||||
)
|
||||
continue
|
||||
request = replace(request, spatial_transforms=source.forward_transforms)
|
||||
grouped[source.sampler_node_id].append((request, source))
|
||||
|
||||
plans: list[AttentionCapturePlan] = []
|
||||
for sampler_id in sorted(grouped):
|
||||
pairs = grouped[sampler_id]
|
||||
authority = pairs[0][1]
|
||||
if any(
|
||||
(
|
||||
source.model_owner_node_id,
|
||||
source.model_link,
|
||||
source.positive_link,
|
||||
)
|
||||
!= (
|
||||
authority.model_owner_node_id,
|
||||
authority.model_link,
|
||||
authority.positive_link,
|
||||
)
|
||||
for _request_value, source in pairs[1:]
|
||||
):
|
||||
LOGGER.warning(
|
||||
"Attention-region requests disagree about sampler authority",
|
||||
extra={"sampler_node_id": sampler_id},
|
||||
)
|
||||
continue
|
||||
requests = tuple(
|
||||
sorted(
|
||||
(request for request, _source in pairs),
|
||||
key=lambda value: value.node_id,
|
||||
)
|
||||
)
|
||||
prompt_text, clip_link = _trace_prompt_text(prompt, authority.positive_link)
|
||||
plans.append(
|
||||
AttentionCapturePlan(
|
||||
sampler_node_id=authority.sampler_node_id,
|
||||
model_owner_node_id=authority.model_owner_node_id,
|
||||
model_input_name="model",
|
||||
model_link=authority.model_link,
|
||||
positive_link=authority.positive_link,
|
||||
requests=requests,
|
||||
prompt_text=prompt_text,
|
||||
clip_link=clip_link,
|
||||
source_aspect=authority.source_aspect,
|
||||
)
|
||||
)
|
||||
return tuple(plans)
|
||||
|
||||
def _resolve_source(
|
||||
self,
|
||||
prompt: Mapping[str, Any],
|
||||
request_node: Mapping[str, Any],
|
||||
node_registry: NodeRegistry,
|
||||
) -> AttentionSamplerStage | BrokenProvenance:
|
||||
"""Select one stage from the request's complete sampler lineage."""
|
||||
|
||||
inputs = _inputs(request_node)
|
||||
source_name = "image" if "image" in inputs else "latent"
|
||||
source_link = parse_graph_link(inputs.get(source_name))
|
||||
if source_link is None:
|
||||
return BrokenProvenance(f"{source_name} input is not a graph link")
|
||||
lineage = ATTENTION_SAMPLER_LINEAGE_RESOLVER.resolve(
|
||||
prompt=prompt,
|
||||
start_link=source_link,
|
||||
source_kind=source_name,
|
||||
node_registry=node_registry,
|
||||
)
|
||||
if isinstance(lineage, BrokenProvenance):
|
||||
return lineage
|
||||
requested_stage = int(inputs.get("sampler_stage", 1))
|
||||
selection = lineage.select(requested_stage)
|
||||
if selection.was_clamped:
|
||||
LOGGER.info(
|
||||
"Attention sampler stage was clamped to direct provenance",
|
||||
extra={
|
||||
"requested_stage": requested_stage,
|
||||
"selected_stage": selection.stage_number,
|
||||
"stage_count": selection.stage_count,
|
||||
},
|
||||
)
|
||||
if not selection.stage.capture_supported:
|
||||
return BrokenProvenance(
|
||||
selection.stage.unsupported_reason
|
||||
or "selected sampler stage cannot project attention to its output",
|
||||
node_id=selection.stage.sampler_node_id,
|
||||
)
|
||||
return selection.stage
|
||||
|
||||
|
||||
class AttentionRegionPromptRewriter:
|
||||
"""Inject one internal MODEL patch node for each planned sampler."""
|
||||
|
||||
def rewrite(
|
||||
self,
|
||||
prompt: MutableMapping[str, Any],
|
||||
plans: tuple[AttentionCapturePlan, ...],
|
||||
) -> None:
|
||||
"""Mutate the prompt graph without changing unrelated node inputs."""
|
||||
|
||||
for plan in plans:
|
||||
if plan.capture_node_id in prompt:
|
||||
raise ValueError(
|
||||
"Attention capture node id collides with the prompt graph."
|
||||
)
|
||||
owner = _node(prompt, plan.model_owner_node_id)
|
||||
if owner is None:
|
||||
raise ValueError(
|
||||
"Attention capture model owner disappeared during rewrite."
|
||||
)
|
||||
owner_inputs = _mutable_inputs(owner)
|
||||
current_model = parse_graph_link(owner_inputs.get(plan.model_input_name))
|
||||
if current_model != plan.model_link:
|
||||
raise ValueError(
|
||||
"Attention capture model edge changed during planning."
|
||||
)
|
||||
capture_inputs: dict[str, Any] = {
|
||||
"model": [plan.model_link[0], plan.model_link[1]],
|
||||
"positive": [plan.positive_link[0], plan.positive_link[1]],
|
||||
"plan_json": _serialize_plan(plan),
|
||||
}
|
||||
if plan.clip_link is not None:
|
||||
capture_inputs["clip"] = [plan.clip_link[0], plan.clip_link[1]]
|
||||
prompt[plan.capture_node_id] = {
|
||||
"class_type": INTERNAL_CAPTURE_NODE_ID,
|
||||
"inputs": capture_inputs,
|
||||
}
|
||||
owner_inputs[plan.model_input_name] = [plan.capture_node_id, 0]
|
||||
|
||||
|
||||
def _request(
|
||||
node_id: str,
|
||||
kind: AttentionRegionRequestKind,
|
||||
node: Mapping[str, Any],
|
||||
) -> AttentionRegionRequest:
|
||||
"""Parse one public node request from serialized prompt inputs."""
|
||||
|
||||
inputs = _inputs(node)
|
||||
raw_concepts = inputs.get("concepts", "")
|
||||
queries = parse_attention_concepts(
|
||||
raw_concepts if isinstance(raw_concepts, str) else str(raw_concepts)
|
||||
)
|
||||
controls = AttentionRegionControls(
|
||||
capture_start=float(inputs.get("capture_start", 0.0)),
|
||||
capture_end=float(inputs.get("capture_end", 1.0)),
|
||||
minimum_strength=float(inputs.get("minimum_strength", 0.15)),
|
||||
minimum_consensus=float(inputs.get("minimum_consensus", 0.25)),
|
||||
split_sensitivity=float(inputs.get("split_sensitivity", 0.35)),
|
||||
minimum_region_size=int(inputs.get("minimum_region_size", 512)),
|
||||
keep_only=int(inputs.get("keep_only", 1)),
|
||||
keep_by=str(inputs.get("keep_by", "largest size")),
|
||||
combine_segs=bool(inputs.get("combine_segs", False)),
|
||||
matte_solidity=float(inputs.get("matte_solidity", 0.75)),
|
||||
edge_feather=int(inputs.get("edge_feather", 8)),
|
||||
profile=AttentionCaptureProfile(str(inputs.get("capture_profile", "fast"))),
|
||||
)
|
||||
return AttentionRegionRequest(
|
||||
node_id,
|
||||
kind,
|
||||
queries,
|
||||
controls,
|
||||
int(inputs.get("sampler_stage", 1)),
|
||||
)
|
||||
|
||||
|
||||
def _trace_prompt_text(
|
||||
prompt: Mapping[str, Any],
|
||||
positive_link: GraphLink,
|
||||
) -> tuple[str | None, GraphLink | None]:
|
||||
"""Recover direct CLIP text and CLIP provenance when graph-visible."""
|
||||
|
||||
node = _node(prompt, positive_link[0])
|
||||
if node is None or _class_type(node) != "CLIPTextEncode":
|
||||
return None, None
|
||||
inputs = _inputs(node)
|
||||
text = inputs.get("text")
|
||||
clip_link = parse_graph_link(inputs.get("clip"))
|
||||
return (text if isinstance(text, str) else None), clip_link
|
||||
|
||||
|
||||
def _serialize_plan(plan: AttentionCapturePlan) -> str:
|
||||
"""Serialize only capture-affecting controls into the upstream MODEL node."""
|
||||
|
||||
capture_requests = tuple(
|
||||
replace(
|
||||
request,
|
||||
controls=AttentionRegionControls(
|
||||
capture_start=request.controls.capture_start,
|
||||
capture_end=request.controls.capture_end,
|
||||
minimum_strength=0.0,
|
||||
minimum_consensus=0.0,
|
||||
split_sensitivity=0.0,
|
||||
minimum_region_size=1,
|
||||
profile=request.controls.profile,
|
||||
),
|
||||
)
|
||||
for request in plan.requests
|
||||
)
|
||||
payload = asdict(replace(plan, requests=capture_requests))
|
||||
payload["model_link"] = list(plan.model_link)
|
||||
payload["positive_link"] = list(plan.positive_link)
|
||||
payload["clip_link"] = list(plan.clip_link) if plan.clip_link is not None else None
|
||||
return json.dumps(payload, separators=(",", ":"), sort_keys=True)
|
||||
|
||||
|
||||
def _node(prompt: Mapping[str, Any], node_id: str) -> Mapping[str, Any] | None:
|
||||
"""Return one valid serialized prompt node."""
|
||||
|
||||
value = prompt.get(node_id)
|
||||
return value if isinstance(value, Mapping) else None
|
||||
|
||||
|
||||
def _class_type(node: Mapping[str, Any]) -> str:
|
||||
"""Return a serialized class type or an empty unsupported value."""
|
||||
|
||||
value = node.get("class_type")
|
||||
return value if isinstance(value, str) else ""
|
||||
|
||||
|
||||
def _inputs(node: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
"""Return immutable serialized inputs or an empty mapping."""
|
||||
|
||||
value = node.get("inputs")
|
||||
return value if isinstance(value, Mapping) else {}
|
||||
|
||||
|
||||
def _mutable_inputs(node: Mapping[str, Any]) -> MutableMapping[str, Any]:
|
||||
"""Return mutable serialized inputs required for prompt rewriting."""
|
||||
|
||||
value = node.get("inputs")
|
||||
if not isinstance(value, MutableMapping):
|
||||
raise TypeError("Prompt node inputs must be mutable for attention capture.")
|
||||
return value
|
||||
|
||||
|
||||
ATTENTION_REGION_PROMPT_PLANNER = AttentionRegionPromptPlanner()
|
||||
ATTENTION_REGION_PROMPT_REWRITER = AttentionRegionPromptRewriter()
|
||||
@@ -0,0 +1,119 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Project absent SDXL query contexts through live cross-attention key layers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol, cast
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ..domain.attention_region_maps import OpenVocabularyContext
|
||||
from .model_object_patch_batch import (
|
||||
ExactModelObjectReplacement,
|
||||
ModelObjectPatchBatchMutation,
|
||||
)
|
||||
|
||||
|
||||
class OpenVocabularyProjectionSink(Protocol):
|
||||
"""Receive query keys projected for the immediately following attention call."""
|
||||
|
||||
def stage_open_vocabulary_keys(
|
||||
self,
|
||||
keys: tuple[tuple[str, torch.Tensor], ...],
|
||||
) -> None:
|
||||
"""Stage projected key tensors under stable query labels."""
|
||||
|
||||
|
||||
class OpenVocabularyKeyProjector(nn.Module):
|
||||
"""Delegate native key projection and side-project compact query contexts."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
original: nn.Module,
|
||||
contexts: tuple[OpenVocabularyContext, ...],
|
||||
sink: OpenVocabularyProjectionSink,
|
||||
) -> None:
|
||||
"""Retain one live key module and immutable CPU query contexts."""
|
||||
|
||||
super().__init__()
|
||||
self.original = original
|
||||
self._contexts = contexts
|
||||
self._sink = sink
|
||||
self._projected_cache: dict[
|
||||
tuple[torch.device, torch.dtype], tuple[tuple[str, torch.Tensor], ...]
|
||||
] = {}
|
||||
|
||||
def forward(self, value: torch.Tensor) -> torch.Tensor:
|
||||
"""Return native keys unchanged after staging selected query keys."""
|
||||
|
||||
native = cast(torch.Tensor, self.original(value))
|
||||
cache_key = (value.device, value.dtype)
|
||||
projected = self._projected_cache.get(cache_key)
|
||||
if projected is None:
|
||||
projected = self._project(value)
|
||||
self._projected_cache[cache_key] = projected
|
||||
self._sink.stage_open_vocabulary_keys(projected)
|
||||
return native
|
||||
|
||||
def _project(
|
||||
self,
|
||||
reference: torch.Tensor,
|
||||
) -> tuple[tuple[str, torch.Tensor], ...]:
|
||||
"""Project immutable query contexts once for one execution device and dtype."""
|
||||
|
||||
projected: list[tuple[str, torch.Tensor]] = []
|
||||
with torch.no_grad():
|
||||
for context in self._contexts:
|
||||
encoded = context.values.to(
|
||||
device=reference.device,
|
||||
dtype=reference.dtype,
|
||||
)
|
||||
query_keys = cast(torch.Tensor, self.original(encoded))
|
||||
indices = torch.tensor(
|
||||
context.token_indices,
|
||||
device=query_keys.device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
projected.append((context.label, query_keys.index_select(1, indices)))
|
||||
return tuple(projected)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OpenVocabularyProjectionMutation:
|
||||
"""Install clone-local wrappers on every standard-UNet cross-attention key."""
|
||||
|
||||
contexts: tuple[OpenVocabularyContext, ...]
|
||||
sink: OpenVocabularyProjectionSink
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Discover canonical attn2 key paths and patch them as one batch."""
|
||||
|
||||
from comfy.ldm.modules.attention import BasicTransformerBlock
|
||||
|
||||
base_model = getattr(model, "model", None)
|
||||
if not isinstance(base_model, nn.Module):
|
||||
raise TypeError("Open-vocabulary MODEL root must be a torch module.")
|
||||
replacements: list[ExactModelObjectReplacement] = []
|
||||
for path, module in base_model.named_modules():
|
||||
if not isinstance(module, BasicTransformerBlock) or module.attn2 is None:
|
||||
continue
|
||||
to_k = module.attn2.to_k
|
||||
if not isinstance(to_k, nn.Module):
|
||||
raise TypeError("Open-vocabulary attn2 key projection is invalid.")
|
||||
replacements.append(
|
||||
ExactModelObjectReplacement(
|
||||
f"{path}.attn2.to_k",
|
||||
to_k,
|
||||
OpenVocabularyKeyProjector(to_k, self.contexts, self.sink),
|
||||
)
|
||||
)
|
||||
if not replacements:
|
||||
raise ValueError(
|
||||
"Open-vocabulary SDXL cross-attention layers were not found."
|
||||
)
|
||||
ModelObjectPatchBatchMutation(tuple(replacements)).apply(model)
|
||||
@@ -0,0 +1,208 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Decode validated prompt-injected attention capture plans at runtime."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from ..domain.attention_region_capture import (
|
||||
AttentionCapturePlan,
|
||||
AttentionCaptureProfile,
|
||||
AttentionRegionControls,
|
||||
AttentionRegionRequest,
|
||||
AttentionRegionRequestKind,
|
||||
)
|
||||
from ..domain.attention_spatial_transform import (
|
||||
AttentionSpatialTransform,
|
||||
AttentionSpatialTransformKind,
|
||||
)
|
||||
from .comfy_graph_provenance import parse_graph_link
|
||||
|
||||
|
||||
class AttentionCapturePlanCodec:
|
||||
"""Own strict deserialization of internal prompt-plan JSON."""
|
||||
|
||||
def decode(self, value: str) -> AttentionCapturePlan:
|
||||
"""Return one validated domain plan from a serialized prompt input."""
|
||||
|
||||
if not isinstance(value, str) or not value:
|
||||
raise ValueError("Attention capture plan JSON cannot be empty.")
|
||||
payload = json.loads(value)
|
||||
if not isinstance(payload, Mapping):
|
||||
raise ValueError("Attention capture plan JSON must contain an object.")
|
||||
requests_value = payload.get("requests")
|
||||
if not isinstance(requests_value, list):
|
||||
raise ValueError("Attention capture plan requests must be a list.")
|
||||
requests = tuple(_request(item) for item in requests_value)
|
||||
return AttentionCapturePlan(
|
||||
sampler_node_id=_string(payload, "sampler_node_id"),
|
||||
model_owner_node_id=_string(payload, "model_owner_node_id"),
|
||||
model_input_name=_string(payload, "model_input_name"),
|
||||
model_link=_link(payload, "model_link"),
|
||||
positive_link=_link(payload, "positive_link"),
|
||||
requests=requests,
|
||||
prompt_text=_optional_string(payload, "prompt_text"),
|
||||
clip_link=_optional_link(payload, "clip_link"),
|
||||
source_aspect=_optional_float(payload, "source_aspect"),
|
||||
)
|
||||
|
||||
|
||||
def _request(value: object) -> AttentionRegionRequest:
|
||||
"""Decode one request and its validated controls."""
|
||||
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("Attention capture request must be an object.")
|
||||
controls_value = value.get("controls")
|
||||
if not isinstance(controls_value, Mapping):
|
||||
raise ValueError("Attention capture request controls must be an object.")
|
||||
queries_value = value.get("queries")
|
||||
if not isinstance(queries_value, list) or any(
|
||||
not isinstance(query, str) for query in queries_value
|
||||
):
|
||||
raise ValueError("Attention capture request queries must be strings.")
|
||||
controls = AttentionRegionControls(
|
||||
capture_start=_float(controls_value, "capture_start"),
|
||||
capture_end=_float(controls_value, "capture_end"),
|
||||
minimum_strength=_float(controls_value, "minimum_strength"),
|
||||
minimum_consensus=_float(controls_value, "minimum_consensus"),
|
||||
split_sensitivity=_float(controls_value, "split_sensitivity"),
|
||||
minimum_region_size=_integer(controls_value, "minimum_region_size"),
|
||||
keep_only=_integer(controls_value, "keep_only"),
|
||||
keep_by=_string(controls_value, "keep_by"),
|
||||
combine_segs=_boolean(controls_value, "combine_segs"),
|
||||
matte_solidity=_float(controls_value, "matte_solidity"),
|
||||
edge_feather=_integer(controls_value, "edge_feather"),
|
||||
profile=AttentionCaptureProfile(_string(controls_value, "profile")),
|
||||
)
|
||||
return AttentionRegionRequest(
|
||||
node_id=_string(value, "node_id"),
|
||||
kind=AttentionRegionRequestKind(_string(value, "kind")),
|
||||
queries=tuple(queries_value),
|
||||
controls=controls,
|
||||
sampler_stage=_integer(value, "sampler_stage"),
|
||||
spatial_transforms=_spatial_transforms(value.get("spatial_transforms")),
|
||||
)
|
||||
|
||||
|
||||
def _string(mapping: Mapping[str, Any], key: str) -> str:
|
||||
"""Return one required non-empty string field."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if not isinstance(value, str) or not value:
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be a string.")
|
||||
return value
|
||||
|
||||
|
||||
def _optional_string(mapping: Mapping[str, Any], key: str) -> str | None:
|
||||
"""Return one optional string field without coercion."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if value is None:
|
||||
return None
|
||||
if not isinstance(value, str):
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be a string.")
|
||||
return value
|
||||
|
||||
|
||||
def _float(mapping: Mapping[str, Any], key: str) -> float:
|
||||
"""Return one required real field as float."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if isinstance(value, bool) or not isinstance(value, int | float):
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be numeric.")
|
||||
return float(value)
|
||||
|
||||
|
||||
def _integer(mapping: Mapping[str, Any], key: str) -> int:
|
||||
"""Return one required integer field."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if type(value) is not int:
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be an integer.")
|
||||
return value
|
||||
|
||||
|
||||
def _boolean(mapping: Mapping[str, Any], key: str) -> bool:
|
||||
"""Return one required strict boolean field."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if type(value) is not bool:
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be boolean.")
|
||||
return value
|
||||
|
||||
|
||||
def _spatial_transforms(value: object) -> tuple[AttentionSpatialTransform, ...]:
|
||||
"""Decode one request's ordered graph-visible spatial transforms."""
|
||||
|
||||
if not isinstance(value, list):
|
||||
raise ValueError("Attention spatial transforms must be a list.")
|
||||
transforms: list[AttentionSpatialTransform] = []
|
||||
for item in value:
|
||||
if not isinstance(item, Mapping):
|
||||
raise ValueError("Attention spatial transform must be an object.")
|
||||
raw_kind = _string(item, "kind")
|
||||
transforms.append(
|
||||
AttentionSpatialTransform(
|
||||
kind=AttentionSpatialTransformKind(raw_kind),
|
||||
width=_optional_integer(item, "width"),
|
||||
height=_optional_integer(item, "height"),
|
||||
scale=_optional_float(item, "scale"),
|
||||
anchor=str(item.get("anchor", "center")),
|
||||
divisible_by=_integer(item, "divisible_by"),
|
||||
)
|
||||
)
|
||||
return tuple(transforms)
|
||||
|
||||
|
||||
def _optional_integer(mapping: Mapping[str, Any], key: str) -> int | None:
|
||||
"""Return one optional strict integer field."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if value is None:
|
||||
return None
|
||||
if type(value) is not int:
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be an integer.")
|
||||
return value
|
||||
|
||||
|
||||
def _optional_float(mapping: Mapping[str, Any], key: str) -> float | None:
|
||||
"""Return one optional real field."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool) or not isinstance(value, int | float):
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be numeric.")
|
||||
return float(value)
|
||||
|
||||
|
||||
def _link(mapping: Mapping[str, Any], key: str) -> tuple[str, int]:
|
||||
"""Return one required serialized graph link."""
|
||||
|
||||
link = parse_graph_link(mapping.get(key))
|
||||
if link is None:
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be a graph link.")
|
||||
return link
|
||||
|
||||
|
||||
def _optional_link(
|
||||
mapping: Mapping[str, Any],
|
||||
key: str,
|
||||
) -> tuple[str, int] | None:
|
||||
"""Return one optional serialized graph link."""
|
||||
|
||||
value = mapping.get(key)
|
||||
if value is None:
|
||||
return None
|
||||
link = parse_graph_link(value)
|
||||
if link is None:
|
||||
raise ValueError(f"Attention capture plan field '{key}' must be a graph link.")
|
||||
return link
|
||||
|
||||
|
||||
ATTENTION_CAPTURE_PLAN_CODEC = AttentionCapturePlanCodec()
|
||||
@@ -0,0 +1,73 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Register pre-validation prompt rewriting for downstream attention nodes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import MutableMapping
|
||||
from importlib import import_module as _import_module
|
||||
from typing import Any, cast
|
||||
|
||||
from .attention_region_graph import (
|
||||
ATTENTION_REGION_PROMPT_PLANNER,
|
||||
ATTENTION_REGION_PROMPT_REWRITER,
|
||||
)
|
||||
from .attention_region_store import ATTENTION_REGION_CAPTURE_STORE
|
||||
from .comfy_graph_provenance import NodeRegistry
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
_REGISTRATION_MARKER = "_simple_syrup_attention_region_prompt_handler"
|
||||
|
||||
|
||||
def register_attention_region_prompt_handler() -> None:
|
||||
"""Register exactly one graph-intent handler when PromptServer is available."""
|
||||
|
||||
try:
|
||||
server_module = _import_module("server")
|
||||
except ImportError:
|
||||
return
|
||||
prompt_server = getattr(server_module, "PromptServer", None)
|
||||
instance = getattr(prompt_server, "instance", None)
|
||||
if instance is None or getattr(instance, _REGISTRATION_MARKER, False):
|
||||
return
|
||||
registrar = getattr(instance, "add_on_prompt_handler", None)
|
||||
if not callable(registrar):
|
||||
return
|
||||
registrar(_rewrite_attention_region_prompt)
|
||||
setattr(instance, _REGISTRATION_MARKER, True)
|
||||
|
||||
|
||||
def _rewrite_attention_region_prompt(json_data: object) -> object:
|
||||
"""Plan and inject attention capture while preserving invalid submissions."""
|
||||
|
||||
if not isinstance(json_data, MutableMapping):
|
||||
return json_data
|
||||
prompt = json_data.get("prompt")
|
||||
if not isinstance(prompt, MutableMapping):
|
||||
return json_data
|
||||
try:
|
||||
ATTENTION_REGION_CAPTURE_STORE.clear()
|
||||
plans = ATTENTION_REGION_PROMPT_PLANNER.build(
|
||||
cast(dict[str, Any], prompt),
|
||||
_node_registry(),
|
||||
)
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(
|
||||
cast(dict[str, Any], prompt),
|
||||
plans,
|
||||
)
|
||||
except (TypeError, ValueError, RuntimeError):
|
||||
LOGGER.exception("Attention-region prompt rewrite failed")
|
||||
return json_data
|
||||
|
||||
|
||||
def _node_registry() -> NodeRegistry:
|
||||
"""Return Comfy's active node registry after extension registration."""
|
||||
|
||||
nodes_module = _import_module("nodes")
|
||||
registry = getattr(nodes_module, "NODE_CLASS_MAPPINGS", None)
|
||||
if not isinstance(registry, dict):
|
||||
raise TypeError("Comfy node registry is unavailable for attention tracing.")
|
||||
return cast(NodeRegistry, registry)
|
||||
@@ -0,0 +1,38 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Publish non-blocking attention-region status through Comfy's progress channel."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from importlib import import_module
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AttentionRegionStatusPublisher:
|
||||
"""Send concise node status without making UI transport a runtime dependency."""
|
||||
|
||||
def publish(self, node_id: str, message: str) -> None:
|
||||
"""Send progress text when Comfy's PromptServer is available."""
|
||||
|
||||
if not node_id or not message:
|
||||
return
|
||||
try:
|
||||
server_module = import_module("server")
|
||||
prompt_server = getattr(server_module, "PromptServer", None)
|
||||
instance = getattr(prompt_server, "instance", None)
|
||||
sender = getattr(instance, "send_progress_text", None)
|
||||
if callable(sender):
|
||||
sender(message, node_id)
|
||||
except (ImportError, AttributeError, RuntimeError):
|
||||
LOGGER.debug(
|
||||
"Attention-region status transport is unavailable",
|
||||
extra={"node_id": node_id, "status": message},
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
ATTENTION_REGION_STATUS_PUBLISHER = AttentionRegionStatusPublisher()
|
||||
@@ -0,0 +1,76 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own prompt-scoped attention capture sessions until public nodes consume them."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from threading import RLock
|
||||
from typing import Protocol
|
||||
|
||||
from ..domain.attention_region_maps import CapturedAttentionMap
|
||||
|
||||
|
||||
class AttentionCaptureSession(Protocol):
|
||||
"""Expose the read surface required by downstream attention nodes."""
|
||||
|
||||
@property
|
||||
def status_message(self) -> str:
|
||||
"""Return one concise user-facing capture status."""
|
||||
|
||||
def maps_for(
|
||||
self,
|
||||
request_node_id: str,
|
||||
) -> tuple[CapturedAttentionMap, ...]:
|
||||
"""Return captured maps associated with one public request node."""
|
||||
|
||||
|
||||
class AttentionRegionCaptureStore:
|
||||
"""Publish and consume request sessions with deterministic cleanup."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create an empty thread-safe session index."""
|
||||
|
||||
self._sessions: dict[str, AttentionCaptureSession] = {}
|
||||
self._lock = RLock()
|
||||
|
||||
def publish(
|
||||
self,
|
||||
request_node_ids: tuple[str, ...],
|
||||
session: AttentionCaptureSession,
|
||||
) -> None:
|
||||
"""Bind every unique request id to one shared capture session."""
|
||||
|
||||
if not request_node_ids or request_node_ids != tuple(
|
||||
sorted(set(request_node_ids))
|
||||
):
|
||||
raise ValueError(
|
||||
"Attention capture request ids must be unique and ordered."
|
||||
)
|
||||
with self._lock:
|
||||
collisions = tuple(
|
||||
node_id for node_id in request_node_ids if node_id in self._sessions
|
||||
)
|
||||
if collisions:
|
||||
raise RuntimeError(
|
||||
"Attention capture request state is already active: "
|
||||
+ ", ".join(collisions)
|
||||
)
|
||||
for node_id in request_node_ids:
|
||||
self._sessions[node_id] = session
|
||||
|
||||
def consume(self, request_node_id: str) -> AttentionCaptureSession | None:
|
||||
"""Remove and return one request binding without affecting its siblings."""
|
||||
|
||||
with self._lock:
|
||||
return self._sessions.pop(request_node_id, None)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Remove all sessions after an interrupted or completed prompt."""
|
||||
|
||||
with self._lock:
|
||||
self._sessions.clear()
|
||||
|
||||
|
||||
ATTENTION_REGION_CAPTURE_STORE = AttentionRegionCaptureStore()
|
||||
@@ -0,0 +1,254 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Map readable prompt concepts to exact SDXL and Anima token positions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections import Counter
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.attention_region_maps import (
|
||||
AttentionTokenCatalog,
|
||||
AttentionTokenSpan,
|
||||
OpenVocabularyContext,
|
||||
)
|
||||
from ..domain.regional_model_capabilities import RegionalModelFamily
|
||||
|
||||
|
||||
class AttentionPromptTokenizer:
|
||||
"""Build model-family token catalogs through a live Comfy CLIP value."""
|
||||
|
||||
def build_catalog(
|
||||
self,
|
||||
*,
|
||||
clip: object,
|
||||
prompt_text: str,
|
||||
model_family: RegionalModelFamily,
|
||||
) -> AttentionTokenCatalog:
|
||||
"""Return comma-concept spans aligned to the denoiser context sequence."""
|
||||
|
||||
tokenize = getattr(clip, "tokenize", None)
|
||||
if not callable(tokenize):
|
||||
raise TypeError(
|
||||
"Attention prompt mapping requires a tokenizable CLIP value."
|
||||
)
|
||||
tokenized = tokenize(prompt_text, return_word_ids=True)
|
||||
family_tokens = _family_tokens(tokenized, model_family)
|
||||
full = _flatten_tokens(family_tokens)
|
||||
if not full:
|
||||
raise ValueError("Attention prompt tokenization returned no tokens.")
|
||||
|
||||
concepts = tuple(
|
||||
part.strip() for part in prompt_text.split(",") if part.strip()
|
||||
)
|
||||
spans: list[AttentionTokenSpan] = []
|
||||
occurrences: Counter[str] = Counter()
|
||||
search_start = 0
|
||||
for concept in concepts:
|
||||
match = _find_concept_tokens(
|
||||
tokenize=tokenize,
|
||||
concept=concept,
|
||||
model_family=model_family,
|
||||
full=full,
|
||||
search_start=search_start,
|
||||
)
|
||||
if match is None:
|
||||
continue
|
||||
indices, search_start = match
|
||||
normalized = " ".join(concept.casefold().replace("_", " ").split())
|
||||
occurrences[normalized] += 1
|
||||
spans.append(AttentionTokenSpan(concept, occurrences[normalized], indices))
|
||||
return AttentionTokenCatalog(
|
||||
len(full),
|
||||
tuple(spans),
|
||||
tuple(item[0] for item in full),
|
||||
)
|
||||
|
||||
def resolve_query(
|
||||
self,
|
||||
*,
|
||||
clip: object,
|
||||
query: str,
|
||||
catalog: AttentionTokenCatalog,
|
||||
model_family: RegionalModelFamily,
|
||||
) -> tuple[AttentionTokenSpan, ...]:
|
||||
"""Resolve an exact label or complete distributed native token coverage."""
|
||||
|
||||
exact = catalog.exact_matches(query)
|
||||
if exact:
|
||||
return exact
|
||||
tokenize = getattr(clip, "tokenize", None)
|
||||
if not callable(tokenize):
|
||||
raise TypeError("Attention query matching requires a tokenizable CLIP.")
|
||||
words = tuple(
|
||||
match.group(0).strip("()[]{}\"'")
|
||||
for match in re.finditer(r"[^\s,;|]+", query)
|
||||
if match.group(0).strip("()[]{}\"'")
|
||||
)
|
||||
if not words:
|
||||
return ()
|
||||
used_indices: set[int] = set()
|
||||
resolved_indices: list[int] = []
|
||||
word_index = 0
|
||||
while word_index < len(words):
|
||||
match: tuple[tuple[int, ...], int] | None = None
|
||||
consumed_words = 0
|
||||
for word_count in range(len(words) - word_index, 0, -1):
|
||||
phrase = " ".join(words[word_index : word_index + word_count])
|
||||
candidate = _find_unused_concept_tokens(
|
||||
tokenize=tokenize,
|
||||
concept=phrase,
|
||||
model_family=model_family,
|
||||
full_ids=catalog.token_ids,
|
||||
used_indices=used_indices,
|
||||
)
|
||||
if candidate is not None:
|
||||
match = candidate
|
||||
consumed_words = word_count
|
||||
break
|
||||
if match is None:
|
||||
return ()
|
||||
indices, _end = match
|
||||
used_indices.update(indices)
|
||||
resolved_indices.extend(indices)
|
||||
word_index += consumed_words
|
||||
return (AttentionTokenSpan(query, 1, tuple(sorted(set(resolved_indices)))),)
|
||||
|
||||
def encode_open_vocabulary_query(
|
||||
self,
|
||||
*,
|
||||
clip: object,
|
||||
query: str,
|
||||
) -> OpenVocabularyContext:
|
||||
"""Encode one absent SDXL phrase without changing sampler conditioning."""
|
||||
|
||||
tokenize = getattr(clip, "tokenize", None)
|
||||
encode = getattr(clip, "encode_from_tokens", None)
|
||||
if not callable(tokenize) or not callable(encode):
|
||||
raise TypeError("Open-vocabulary search requires an encodable CLIP value.")
|
||||
tokens = tokenize(query, return_word_ids=True)
|
||||
stream = _flatten_tokens(
|
||||
_family_tokens(tokens, RegionalModelFamily.STANDARD_UNET)
|
||||
)
|
||||
indices = tuple(index for index, item in enumerate(stream) if item[2] != 0)
|
||||
if not indices:
|
||||
raise ValueError(f"Open-vocabulary query {query!r} has no semantic tokens.")
|
||||
encoded = encode(tokens)
|
||||
if not isinstance(encoded, torch.Tensor):
|
||||
raise TypeError("Open-vocabulary CLIP encoding did not return a tensor.")
|
||||
return OpenVocabularyContext(
|
||||
query,
|
||||
encoded.detach().to(device="cpu", dtype=torch.float16),
|
||||
indices,
|
||||
)
|
||||
|
||||
|
||||
def _family_tokens(tokenized: object, family: RegionalModelFamily) -> object:
|
||||
"""Select the token stream that feeds each supported denoiser context."""
|
||||
|
||||
if not isinstance(tokenized, Mapping):
|
||||
raise TypeError("Comfy CLIP tokenization must return a mapping.")
|
||||
key = "t5xxl" if family is RegionalModelFamily.ANIMA else "g"
|
||||
value = tokenized.get(key)
|
||||
if value is None and family is RegionalModelFamily.STANDARD_UNET:
|
||||
value = tokenized.get("l")
|
||||
if value is None:
|
||||
raise ValueError(f"Attention prompt tokenization has no '{key}' stream.")
|
||||
return value
|
||||
|
||||
|
||||
def _flatten_tokens(value: object) -> tuple[tuple[object, float, int], ...]:
|
||||
"""Flatten Comfy token chunks while retaining special-token positions."""
|
||||
|
||||
if not isinstance(value, Sequence) or isinstance(value, str | bytes):
|
||||
raise TypeError("Attention token stream must contain token chunks.")
|
||||
flattened: list[tuple[object, float, int]] = []
|
||||
for chunk in value:
|
||||
if not isinstance(chunk, Sequence) or isinstance(chunk, str | bytes):
|
||||
raise TypeError("Attention token chunk must be a sequence.")
|
||||
for item in chunk:
|
||||
if not isinstance(item, Sequence) or len(item) < 3:
|
||||
raise ValueError(
|
||||
"Attention token entries require id, weight, and word id."
|
||||
)
|
||||
word_id = item[2]
|
||||
if isinstance(word_id, bool) or not isinstance(word_id, int):
|
||||
raise TypeError("Attention token word ids must be integers.")
|
||||
flattened.append((item[0], float(item[1]), word_id))
|
||||
return tuple(flattened)
|
||||
|
||||
|
||||
def _find_concept_tokens(
|
||||
*,
|
||||
tokenize: Any,
|
||||
concept: str,
|
||||
model_family: RegionalModelFamily,
|
||||
full: tuple[tuple[object, float, int], ...],
|
||||
search_start: int,
|
||||
) -> tuple[tuple[int, ...], int] | None:
|
||||
"""Find one concept occurrence using tokenizer-equivalent prefix variants."""
|
||||
|
||||
full_ids = tuple(item[0] for item in full)
|
||||
for candidate in (concept, f" {concept}", f", {concept}"):
|
||||
candidate_tokens = _flatten_tokens(
|
||||
_family_tokens(tokenize(candidate, return_word_ids=True), model_family)
|
||||
)
|
||||
candidate_ids = tuple(item[0] for item in candidate_tokens if item[2] != 0)
|
||||
if not candidate_ids:
|
||||
continue
|
||||
start = _find_subsequence(full_ids, candidate_ids, search_start)
|
||||
if start is not None:
|
||||
indices = tuple(range(start, start + len(candidate_ids)))
|
||||
return indices, start + len(candidate_ids)
|
||||
return None
|
||||
|
||||
|
||||
def _find_subsequence(
|
||||
values: tuple[object, ...],
|
||||
target: tuple[object, ...],
|
||||
start: int,
|
||||
) -> int | None:
|
||||
"""Return the first target occurrence at or after a stable cursor."""
|
||||
|
||||
final_start = len(values) - len(target)
|
||||
for index in range(max(0, start), final_start + 1):
|
||||
if values[index : index + len(target)] == target:
|
||||
return index
|
||||
return None
|
||||
|
||||
|
||||
def _find_unused_concept_tokens(
|
||||
*,
|
||||
tokenize: Any,
|
||||
concept: str,
|
||||
model_family: RegionalModelFamily,
|
||||
full_ids: tuple[object, ...],
|
||||
used_indices: set[int],
|
||||
) -> tuple[tuple[int, ...], int] | None:
|
||||
"""Find one tokenizer-equivalent phrase occurrence not already consumed."""
|
||||
|
||||
for candidate in (concept, f" {concept}", f", {concept}"):
|
||||
candidate_tokens = _flatten_tokens(
|
||||
_family_tokens(tokenize(candidate, return_word_ids=True), model_family)
|
||||
)
|
||||
candidate_ids = tuple(item[0] for item in candidate_tokens if item[2] != 0)
|
||||
if not candidate_ids:
|
||||
continue
|
||||
search_start = 0
|
||||
while (
|
||||
start := _find_subsequence(full_ids, candidate_ids, search_start)
|
||||
) is not None:
|
||||
indices = tuple(range(start, start + len(candidate_ids)))
|
||||
if used_indices.isdisjoint(indices):
|
||||
return indices, start + len(candidate_ids)
|
||||
search_start = start + 1
|
||||
return None
|
||||
|
||||
|
||||
ATTENTION_PROMPT_TOKENIZER = AttentionPromptTokenizer()
|
||||
@@ -0,0 +1,510 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Trace ordered sampler lineage through known spatial graph boundaries."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Any
|
||||
|
||||
from ..domain.attention_sampler_lineage import (
|
||||
AttentionSamplerLineage,
|
||||
AttentionSamplerStage,
|
||||
)
|
||||
from ..domain.attention_spatial_transform import (
|
||||
AttentionSpatialTransform,
|
||||
AttentionSpatialTransformKind,
|
||||
)
|
||||
from ..domain.graph_provenance import BrokenProvenance, GraphLink
|
||||
from .comfy_graph_provenance import MAX_PROVENANCE_HOPS, NodeRegistry, parse_graph_link
|
||||
|
||||
_DIRECT_SAMPLERS = frozenset({"KSampler", "KSamplerAdvanced", "SamplerCustom"})
|
||||
_GUIDER_SAMPLER = "SamplerCustomAdvanced"
|
||||
_FULL_CANVAS_DETAILERS = frozenset({"SimpleSyrup.DetailSEGSAsRegions"})
|
||||
_CROP_LOCAL_DETAILERS = frozenset(
|
||||
{
|
||||
"SimpleSyrup.DetailSEGSByScaleFactor",
|
||||
"SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion",
|
||||
"DetailerForEach",
|
||||
}
|
||||
)
|
||||
_IMAGE_TO_LATENT = {
|
||||
"VAEDecode": "samples",
|
||||
"SimpleSyrup.VAEDecodeOptions": "samples",
|
||||
}
|
||||
_LATENT_TO_IMAGE = {
|
||||
"VAEEncode": "pixels",
|
||||
"VAEEncodeForInpaint": "pixels",
|
||||
"SimpleSyrup.VAEEncodeOptions": "pixels",
|
||||
"SimpleSyrup.SimpleVAEEncode": "image",
|
||||
}
|
||||
_IMAGE_TRANSFORMS = {
|
||||
"ImageScale": "image",
|
||||
"ImageScaleBy": "image",
|
||||
"ImageUpscaleWithModel": "image",
|
||||
"SimpleSyrup.ResizeImageToTarget": "image",
|
||||
}
|
||||
_LATENT_TRANSFORMS = {
|
||||
"LatentUpscale": "samples",
|
||||
"LatentUpscaleBy": "samples",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Cursor:
|
||||
"""Track one graph link and its spatial value kind."""
|
||||
|
||||
link: GraphLink
|
||||
kind: str
|
||||
|
||||
|
||||
class AttentionSamplerLineageResolver:
|
||||
"""Discover every recognized sampler on one connected spatial ancestry."""
|
||||
|
||||
def resolve(
|
||||
self,
|
||||
*,
|
||||
prompt: Mapping[str, Any],
|
||||
start_link: GraphLink,
|
||||
source_kind: str,
|
||||
node_registry: NodeRegistry,
|
||||
) -> AttentionSamplerLineage | BrokenProvenance:
|
||||
"""Return chronological stages or an exact provenance failure."""
|
||||
|
||||
current = _Cursor(start_link, source_kind)
|
||||
visited: set[GraphLink] = set()
|
||||
reverse_stages: list[AttentionSamplerStage] = []
|
||||
reverse_transforms: list[AttentionSpatialTransform] = []
|
||||
for _hop in range(MAX_PROVENANCE_HOPS):
|
||||
if current.link in visited:
|
||||
return BrokenProvenance(
|
||||
"sampler lineage contains a cycle", node_id=current.link[0]
|
||||
)
|
||||
visited.add(current.link)
|
||||
node = _node(prompt, current.link[0])
|
||||
if node is None:
|
||||
return BrokenProvenance(
|
||||
"sampler lineage source node is missing", node_id=current.link[0]
|
||||
)
|
||||
class_type = _class_type(node)
|
||||
inputs = _inputs(node)
|
||||
stage = _sampling_stage(
|
||||
prompt,
|
||||
current,
|
||||
class_type,
|
||||
inputs,
|
||||
node_registry,
|
||||
)
|
||||
if isinstance(stage, BrokenProvenance):
|
||||
return stage
|
||||
if stage is not None:
|
||||
stage = replace(
|
||||
stage,
|
||||
forward_transforms=tuple(reversed(reverse_transforms)),
|
||||
)
|
||||
reverse_stages.append(stage)
|
||||
current = _Cursor(stage.upstream_link, stage.upstream_kind)
|
||||
continue
|
||||
next_cursor = _spatial_predecessor(
|
||||
current, class_type, inputs, node_registry
|
||||
)
|
||||
if isinstance(next_cursor, BrokenProvenance):
|
||||
if reverse_stages and _is_spatial_terminal(inputs):
|
||||
break
|
||||
return next_cursor
|
||||
transform = _spatial_transform(class_type, inputs)
|
||||
if isinstance(transform, BrokenProvenance):
|
||||
return transform
|
||||
if transform is not None:
|
||||
reverse_transforms.append(transform)
|
||||
current = next_cursor
|
||||
else:
|
||||
return BrokenProvenance(
|
||||
"sampler lineage exceeded the hop limit", node_id=current.link[0]
|
||||
)
|
||||
if not reverse_stages:
|
||||
return BrokenProvenance("sampler lineage contains no supported sampler")
|
||||
return AttentionSamplerLineage(tuple(reversed(reverse_stages)))
|
||||
|
||||
|
||||
def _sampling_stage(
|
||||
prompt: Mapping[str, Any],
|
||||
cursor: _Cursor,
|
||||
class_type: str,
|
||||
inputs: Mapping[str, Any],
|
||||
node_registry: NodeRegistry,
|
||||
) -> AttentionSamplerStage | BrokenProvenance | None:
|
||||
"""Resolve recognized direct, guider, and detailer sampling authorities."""
|
||||
|
||||
if class_type in _DIRECT_SAMPLERS:
|
||||
return _direct_stage(
|
||||
prompt,
|
||||
cursor.link[0],
|
||||
inputs,
|
||||
"latent_image",
|
||||
"latent",
|
||||
node_registry,
|
||||
)
|
||||
if class_type == _GUIDER_SAMPLER:
|
||||
guider_link = parse_graph_link(inputs.get("guider"))
|
||||
upstream = parse_graph_link(inputs.get("latent_image"))
|
||||
if guider_link is None or upstream is None:
|
||||
return BrokenProvenance(
|
||||
"advanced sampler guider or latent input is not a graph link",
|
||||
node_id=cursor.link[0],
|
||||
)
|
||||
guider = _node(prompt, guider_link[0])
|
||||
if guider is None:
|
||||
return BrokenProvenance(
|
||||
"advanced sampler guider is missing", guider_link[0]
|
||||
)
|
||||
guider_inputs = _inputs(guider)
|
||||
model = parse_graph_link(guider_inputs.get("model"))
|
||||
positive = parse_graph_link(
|
||||
guider_inputs.get("positive") or guider_inputs.get("conditioning")
|
||||
)
|
||||
if model is None or positive is None:
|
||||
return BrokenProvenance(
|
||||
"advanced guider MODEL or positive is not a graph link",
|
||||
node_id=guider_link[0],
|
||||
)
|
||||
return AttentionSamplerStage(
|
||||
cursor.link[0],
|
||||
guider_link[0],
|
||||
model,
|
||||
positive,
|
||||
upstream,
|
||||
"latent",
|
||||
source_aspect=_source_aspect(
|
||||
prompt,
|
||||
upstream,
|
||||
"latent",
|
||||
node_registry,
|
||||
),
|
||||
)
|
||||
if class_type in _FULL_CANVAS_DETAILERS:
|
||||
return _direct_stage(
|
||||
prompt,
|
||||
cursor.link[0],
|
||||
inputs,
|
||||
"image",
|
||||
"image",
|
||||
node_registry,
|
||||
)
|
||||
if class_type in _CROP_LOCAL_DETAILERS:
|
||||
stage = _direct_stage(
|
||||
prompt,
|
||||
cursor.link[0],
|
||||
inputs,
|
||||
"image",
|
||||
"image",
|
||||
node_registry,
|
||||
)
|
||||
if isinstance(stage, BrokenProvenance):
|
||||
return stage
|
||||
return replace(
|
||||
stage,
|
||||
capture_supported=False,
|
||||
unsupported_reason=(
|
||||
"the sampler runs on crop-local images whose placement is not "
|
||||
"available to attention capture"
|
||||
),
|
||||
)
|
||||
generic = _generic_latent_stage(
|
||||
prompt,
|
||||
cursor.link[0],
|
||||
inputs,
|
||||
node_registry,
|
||||
)
|
||||
if generic is not None:
|
||||
return generic
|
||||
return None
|
||||
|
||||
|
||||
def _generic_latent_stage(
|
||||
prompt: Mapping[str, Any],
|
||||
node_id: str,
|
||||
inputs: Mapping[str, Any],
|
||||
node_registry: NodeRegistry,
|
||||
) -> AttentionSamplerStage | None:
|
||||
"""Recognize sampler-compatible nodes by their standard graph inputs."""
|
||||
|
||||
if not all(name in inputs for name in ("model", "positive", "latent_image")):
|
||||
return None
|
||||
stage = _direct_stage(
|
||||
prompt,
|
||||
node_id,
|
||||
inputs,
|
||||
"latent_image",
|
||||
"latent",
|
||||
node_registry,
|
||||
)
|
||||
return None if isinstance(stage, BrokenProvenance) else stage
|
||||
|
||||
|
||||
def _direct_stage(
|
||||
prompt: Mapping[str, Any],
|
||||
node_id: str,
|
||||
inputs: Mapping[str, Any],
|
||||
upstream_name: str,
|
||||
upstream_kind: str,
|
||||
node_registry: NodeRegistry,
|
||||
) -> AttentionSamplerStage | BrokenProvenance:
|
||||
"""Resolve one node whose MODEL and positive are direct graph inputs."""
|
||||
|
||||
model = parse_graph_link(inputs.get("model"))
|
||||
positive = parse_graph_link(inputs.get("positive"))
|
||||
upstream = parse_graph_link(inputs.get(upstream_name))
|
||||
if model is None or positive is None or upstream is None:
|
||||
return BrokenProvenance(
|
||||
"sampling stage MODEL, positive, or spatial input is not a graph link",
|
||||
node_id=node_id,
|
||||
)
|
||||
return AttentionSamplerStage(
|
||||
node_id,
|
||||
node_id,
|
||||
model,
|
||||
positive,
|
||||
upstream,
|
||||
upstream_kind,
|
||||
source_aspect=_source_aspect(
|
||||
prompt,
|
||||
upstream,
|
||||
upstream_kind,
|
||||
node_registry,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _spatial_predecessor(
|
||||
cursor: _Cursor,
|
||||
class_type: str,
|
||||
inputs: Mapping[str, Any],
|
||||
node_registry: NodeRegistry,
|
||||
) -> _Cursor | BrokenProvenance:
|
||||
"""Follow one known modality boundary, resize, or declared passthrough."""
|
||||
|
||||
mapping: Mapping[str, str]
|
||||
next_kind = cursor.kind
|
||||
if cursor.kind == "image" and class_type in _IMAGE_TO_LATENT:
|
||||
mapping = _IMAGE_TO_LATENT
|
||||
next_kind = "latent"
|
||||
elif cursor.kind == "latent" and class_type in _LATENT_TO_IMAGE:
|
||||
mapping = _LATENT_TO_IMAGE
|
||||
next_kind = "image"
|
||||
elif cursor.kind == "image" and class_type in _IMAGE_TRANSFORMS:
|
||||
mapping = _IMAGE_TRANSFORMS
|
||||
elif cursor.kind == "latent" and class_type in _LATENT_TRANSFORMS:
|
||||
mapping = _LATENT_TRANSFORMS
|
||||
else:
|
||||
class_def = node_registry.get(class_type)
|
||||
rules = (
|
||||
getattr(class_def, "GRAPH_PASSTHROUGH_OUTPUTS", None) if class_def else None
|
||||
)
|
||||
input_name = rules.get(cursor.link[1]) if isinstance(rules, Mapping) else None
|
||||
if not isinstance(input_name, str):
|
||||
return BrokenProvenance(
|
||||
"spatial source has no recognized provenance adapter",
|
||||
node_id=cursor.link[0],
|
||||
class_type=class_type,
|
||||
)
|
||||
mapping = {class_type: input_name}
|
||||
next_link = parse_graph_link(inputs.get(mapping[class_type]))
|
||||
if next_link is None:
|
||||
return BrokenProvenance(
|
||||
"spatial predecessor input is not a graph link",
|
||||
node_id=cursor.link[0],
|
||||
class_type=class_type,
|
||||
)
|
||||
return _Cursor(next_link, next_kind)
|
||||
|
||||
|
||||
def _is_spatial_terminal(inputs: Mapping[str, Any]) -> bool:
|
||||
"""Return whether a source exposes no earlier image or latent graph link."""
|
||||
|
||||
return not any(
|
||||
parse_graph_link(inputs.get(name)) is not None
|
||||
for name in ("image", "pixels", "samples", "latent", "latent_image")
|
||||
)
|
||||
|
||||
|
||||
def _source_aspect(
|
||||
prompt: Mapping[str, Any],
|
||||
start_link: GraphLink,
|
||||
source_kind: str,
|
||||
node_registry: NodeRegistry,
|
||||
) -> float | None:
|
||||
"""Resolve a sampler input aspect from graph-visible spatial ancestry."""
|
||||
|
||||
current = _Cursor(start_link, source_kind)
|
||||
visited: set[GraphLink] = set()
|
||||
for _hop in range(MAX_PROVENANCE_HOPS):
|
||||
if current.link in visited:
|
||||
return None
|
||||
visited.add(current.link)
|
||||
node = _node(prompt, current.link[0])
|
||||
if node is None:
|
||||
return None
|
||||
class_type = _class_type(node)
|
||||
inputs = _inputs(node)
|
||||
explicit = _explicit_output_aspect(class_type, inputs)
|
||||
if explicit is not None:
|
||||
return explicit
|
||||
sampling_input = _sampling_spatial_input(class_type, inputs)
|
||||
if sampling_input is not None:
|
||||
current = sampling_input
|
||||
continue
|
||||
predecessor = _spatial_predecessor(
|
||||
current,
|
||||
class_type,
|
||||
inputs,
|
||||
node_registry,
|
||||
)
|
||||
if isinstance(predecessor, BrokenProvenance):
|
||||
return None
|
||||
current = predecessor
|
||||
return None
|
||||
|
||||
|
||||
def _explicit_output_aspect(
|
||||
class_type: str,
|
||||
inputs: Mapping[str, Any],
|
||||
) -> float | None:
|
||||
"""Return dimensions declared by a terminal or sized transform."""
|
||||
|
||||
sized_transform = class_type in {
|
||||
"ImageScale",
|
||||
"LatentUpscale",
|
||||
"SimpleSyrup.ResizeImageToTarget",
|
||||
}
|
||||
if not sized_transform and not _is_spatial_terminal(inputs):
|
||||
return None
|
||||
width = inputs.get("width")
|
||||
height = inputs.get("height")
|
||||
if type(width) is not int or width < 1 or type(height) is not int or height < 1:
|
||||
return None
|
||||
return width / height
|
||||
|
||||
|
||||
def _sampling_spatial_input(
|
||||
class_type: str,
|
||||
inputs: Mapping[str, Any],
|
||||
) -> _Cursor | None:
|
||||
"""Follow through a sampler while resolving the aspect of a later stage."""
|
||||
|
||||
if class_type in _DIRECT_SAMPLERS or class_type == _GUIDER_SAMPLER:
|
||||
link = parse_graph_link(inputs.get("latent_image"))
|
||||
return _Cursor(link, "latent") if link is not None else None
|
||||
if class_type in _FULL_CANVAS_DETAILERS | _CROP_LOCAL_DETAILERS:
|
||||
link = parse_graph_link(inputs.get("image"))
|
||||
return _Cursor(link, "image") if link is not None else None
|
||||
if all(name in inputs for name in ("model", "positive", "latent_image")):
|
||||
link = parse_graph_link(inputs.get("latent_image"))
|
||||
return _Cursor(link, "latent") if link is not None else None
|
||||
return None
|
||||
|
||||
|
||||
def _node(prompt: Mapping[str, Any], node_id: str) -> Mapping[str, Any] | None:
|
||||
"""Return one serialized prompt node."""
|
||||
|
||||
value = prompt.get(node_id)
|
||||
return value if isinstance(value, Mapping) else None
|
||||
|
||||
|
||||
def _class_type(node: Mapping[str, Any]) -> str:
|
||||
"""Return one serialized class type."""
|
||||
|
||||
value = node.get("class_type")
|
||||
return value if isinstance(value, str) else ""
|
||||
|
||||
|
||||
def _inputs(node: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
"""Return serialized node inputs."""
|
||||
|
||||
value = node.get("inputs")
|
||||
return value if isinstance(value, Mapping) else {}
|
||||
|
||||
|
||||
ATTENTION_SAMPLER_LINEAGE_RESOLVER = AttentionSamplerLineageResolver()
|
||||
|
||||
|
||||
def _spatial_transform(
|
||||
class_type: str, inputs: Mapping[str, Any]
|
||||
) -> AttentionSpatialTransform | BrokenProvenance | None:
|
||||
"""Decode one recognized full-canvas transform from serialized inputs."""
|
||||
|
||||
if class_type in {"ImageScaleBy", "LatentUpscaleBy"}:
|
||||
return _scale_transform(class_type, inputs.get("scale_by"))
|
||||
if class_type in {"ImageScale", "LatentUpscale"}:
|
||||
return _sized_transform(
|
||||
class_type,
|
||||
inputs,
|
||||
AttentionSpatialTransformKind.RESIZE,
|
||||
)
|
||||
if class_type != "SimpleSyrup.ResizeImageToTarget":
|
||||
return None
|
||||
raw_mode = inputs.get("resize_mode", "Keep AR")
|
||||
mode = str(raw_mode).casefold()
|
||||
kind = (
|
||||
AttentionSpatialTransformKind.FIT_RESIZE
|
||||
if "keep ar" in mode
|
||||
else AttentionSpatialTransformKind.RESIZE
|
||||
)
|
||||
if "crop" in mode:
|
||||
kind = AttentionSpatialTransformKind.COVER_CROP
|
||||
elif "pad" in mode:
|
||||
kind = AttentionSpatialTransformKind.FIT_PAD
|
||||
result = _sized_transform(class_type, inputs, kind)
|
||||
if isinstance(result, BrokenProvenance):
|
||||
return result
|
||||
divisible_by = inputs.get("divisible_by", 1)
|
||||
if type(divisible_by) is not int or divisible_by < 1:
|
||||
return BrokenProvenance(
|
||||
"spatial divisibility is not graph-visible",
|
||||
class_type=class_type,
|
||||
)
|
||||
return replace(
|
||||
result,
|
||||
anchor=str(inputs.get("crop_position", "center")),
|
||||
divisible_by=divisible_by,
|
||||
)
|
||||
|
||||
|
||||
def _scale_transform(
|
||||
class_type: str, raw_scale: object
|
||||
) -> AttentionSpatialTransform | BrokenProvenance:
|
||||
"""Decode one positive numeric scale factor."""
|
||||
|
||||
if isinstance(raw_scale, bool) or not isinstance(raw_scale, int | float):
|
||||
return BrokenProvenance(
|
||||
"spatial scale is not a graph-visible number", class_type=class_type
|
||||
)
|
||||
try:
|
||||
return AttentionSpatialTransform(
|
||||
AttentionSpatialTransformKind.SCALE,
|
||||
scale=float(raw_scale),
|
||||
)
|
||||
except ValueError as exc:
|
||||
return BrokenProvenance(str(exc), class_type=class_type)
|
||||
|
||||
|
||||
def _sized_transform(
|
||||
class_type: str,
|
||||
inputs: Mapping[str, Any],
|
||||
kind: AttentionSpatialTransformKind,
|
||||
) -> AttentionSpatialTransform | BrokenProvenance:
|
||||
"""Decode one graph-visible positive target size."""
|
||||
|
||||
width = inputs.get("width")
|
||||
height = inputs.get("height")
|
||||
if type(width) is not int or type(height) is not int:
|
||||
return BrokenProvenance(
|
||||
"spatial target size is not graph-visible", class_type=class_type
|
||||
)
|
||||
try:
|
||||
return AttentionSpatialTransform(kind, width=width, height=height)
|
||||
except ValueError as exc:
|
||||
return BrokenProvenance(str(exc), class_type=class_type)
|
||||
@@ -0,0 +1,167 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Prepare an observation-only MODEL from one prompt-injected capture plan."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ..domain.attention_region_capture import (
|
||||
AttentionCapturePlan,
|
||||
AttentionRegionRequestKind,
|
||||
)
|
||||
from ..domain.attention_region_maps import (
|
||||
AttentionTokenCatalog,
|
||||
AttentionTokenSpan,
|
||||
CapturedAttentionMap,
|
||||
OpenVocabularyContext,
|
||||
)
|
||||
from ..domain.regional_model_capabilities import RegionalModelFamily
|
||||
from ..runtime.attention_region_capture import AttentionRegionCaptureSession
|
||||
from ..runtime.attention_region_capture_backend import (
|
||||
ATTENTION_REGION_CAPTURE_BACKEND,
|
||||
)
|
||||
from ..runtime.attention_region_plan_codec import ATTENTION_CAPTURE_PLAN_CODEC
|
||||
from ..runtime.attention_region_store import (
|
||||
ATTENTION_REGION_CAPTURE_STORE,
|
||||
AttentionCaptureSession,
|
||||
)
|
||||
from ..runtime.attention_region_tokens import ATTENTION_PROMPT_TOKENIZER
|
||||
from ..runtime.regional_model_capabilities import RegionalModelCapabilityRegistry
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EmptyAttentionCaptureSession:
|
||||
"""Represent an observable unsupported capture without fabricated maps."""
|
||||
|
||||
request_node_ids: tuple[str, ...]
|
||||
status_message: str
|
||||
|
||||
def maps_for(self, request_node_id: str) -> tuple[CapturedAttentionMap, ...]:
|
||||
"""Return no maps after validating request ownership."""
|
||||
|
||||
if request_node_id not in self.request_node_ids:
|
||||
raise KeyError(f"Unknown attention-region request node: {request_node_id}.")
|
||||
return ()
|
||||
|
||||
|
||||
class AttentionCaptureModelService:
|
||||
"""Select the model adapter, publish session state, and derive MODEL."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create an isolated model capability registry."""
|
||||
|
||||
self._capabilities = RegionalModelCapabilityRegistry()
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
*,
|
||||
model: object,
|
||||
plan_json: str,
|
||||
clip: object | None,
|
||||
) -> object:
|
||||
"""Return a capture MODEL or the original MODEL for an observable no-op."""
|
||||
|
||||
plan = ATTENTION_CAPTURE_PLAN_CODEC.decode(plan_json)
|
||||
request_ids = tuple(request.node_id for request in plan.requests)
|
||||
try:
|
||||
capabilities = self._capabilities.capabilities_for(model)
|
||||
if clip is None or plan.prompt_text is None:
|
||||
raise ValueError("prompt text or CLIP provenance is not graph-visible")
|
||||
catalog = ATTENTION_PROMPT_TOKENIZER.build_catalog(
|
||||
clip=clip,
|
||||
prompt_text=plan.prompt_text,
|
||||
model_family=capabilities.model_family,
|
||||
)
|
||||
request_spans, unresolved = _resolve_request_spans(
|
||||
plan=plan,
|
||||
clip=clip,
|
||||
catalog=catalog,
|
||||
model_family=capabilities.model_family,
|
||||
)
|
||||
open_contexts: tuple[OpenVocabularyContext, ...] = ()
|
||||
if capabilities.model_family.value == "standard_unet":
|
||||
open_contexts = tuple(
|
||||
ATTENTION_PROMPT_TOKENIZER.encode_open_vocabulary_query(
|
||||
clip=clip,
|
||||
query=query,
|
||||
)
|
||||
for query in unresolved
|
||||
)
|
||||
capture_session = AttentionRegionCaptureSession(
|
||||
plan=plan,
|
||||
model_family=capabilities.model_family,
|
||||
token_catalog=catalog,
|
||||
request_spans=request_spans,
|
||||
open_vocabulary_contexts=_unique_open_contexts(open_contexts),
|
||||
)
|
||||
session: AttentionCaptureSession = capture_session
|
||||
derived = ATTENTION_REGION_CAPTURE_BACKEND.derive(model, capture_session)
|
||||
except (TypeError, ValueError, RuntimeError) as exc:
|
||||
reason = f"Attention capture unavailable: {exc}"
|
||||
LOGGER.warning(
|
||||
"Attention-region MODEL preparation no-op",
|
||||
extra={
|
||||
"sampler_node_id": plan.sampler_node_id,
|
||||
"request_node_ids": request_ids,
|
||||
"reason": str(exc),
|
||||
},
|
||||
)
|
||||
session = EmptyAttentionCaptureSession(request_ids, reason)
|
||||
derived = model
|
||||
ATTENTION_REGION_CAPTURE_STORE.publish(request_ids, session)
|
||||
return derived
|
||||
|
||||
|
||||
ATTENTION_CAPTURE_MODEL_SERVICE = AttentionCaptureModelService()
|
||||
|
||||
|
||||
def _resolve_request_spans(
|
||||
*,
|
||||
plan: AttentionCapturePlan,
|
||||
clip: object,
|
||||
catalog: AttentionTokenCatalog,
|
||||
model_family: RegionalModelFamily,
|
||||
) -> tuple[dict[str, tuple[AttentionTokenSpan, ...]], tuple[str, ...]]:
|
||||
"""Resolve and deduplicate native request spans before MODEL patching."""
|
||||
|
||||
resolved: dict[str, tuple[AttentionTokenSpan, ...]] = {}
|
||||
unresolved: list[str] = []
|
||||
for request in plan.requests:
|
||||
if request.kind is AttentionRegionRequestKind.ALL_PROMPT_SEGS:
|
||||
resolved[request.node_id] = catalog.spans
|
||||
continue
|
||||
spans: list[AttentionTokenSpan] = []
|
||||
for query in request.queries:
|
||||
matches = ATTENTION_PROMPT_TOKENIZER.resolve_query(
|
||||
clip=clip,
|
||||
query=query,
|
||||
catalog=catalog,
|
||||
model_family=model_family,
|
||||
)
|
||||
if matches:
|
||||
spans.extend(matches)
|
||||
else:
|
||||
unresolved.append(query)
|
||||
LOGGER.info(
|
||||
"Attention concept is absent from native conditioning",
|
||||
extra={"request_node_id": request.node_id, "concept": query},
|
||||
)
|
||||
resolved[request.node_id] = tuple(dict.fromkeys(spans))
|
||||
return resolved, tuple(dict.fromkeys(unresolved))
|
||||
|
||||
|
||||
def _unique_open_contexts(
|
||||
contexts: tuple[OpenVocabularyContext, ...],
|
||||
) -> tuple[OpenVocabularyContext, ...]:
|
||||
"""Deduplicate query contexts by normalized public label in stable order."""
|
||||
|
||||
unique: dict[str, OpenVocabularyContext] = {}
|
||||
for context in contexts:
|
||||
unique.setdefault(context.label.casefold(), context)
|
||||
return tuple(unique.values())
|
||||
@@ -0,0 +1,99 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Extract, score, and retain attention-region components."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.attention_region_capture import AttentionRegionControls
|
||||
from ..domain.segs import BoundingBox
|
||||
from ..masking.mask_components import connected_mask_components
|
||||
from .attention_region_evidence import AttentionConceptEvidence
|
||||
from .attention_region_instance_splitting import (
|
||||
ATTENTION_INSTANCE_SPLITTING_SERVICE,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionRegionComponent:
|
||||
"""Represent one retained concept instance before matte shaping."""
|
||||
|
||||
label: str
|
||||
alpha: torch.Tensor
|
||||
support: torch.Tensor
|
||||
bbox: BoundingBox
|
||||
area: int
|
||||
confidence: float
|
||||
|
||||
|
||||
class AttentionComponentService:
|
||||
"""Own per-concept component filtering and deterministic retention."""
|
||||
|
||||
def retained(
|
||||
self,
|
||||
evidence: AttentionConceptEvidence,
|
||||
controls: AttentionRegionControls,
|
||||
) -> tuple[AttentionRegionComponent, ...]:
|
||||
"""Return valid components under the requested per-concept policy."""
|
||||
|
||||
candidates: list[AttentionRegionComponent] = []
|
||||
for cohesive_region in connected_mask_components(evidence.support):
|
||||
partitions = ATTENTION_INSTANCE_SPLITTING_SERVICE.partition(
|
||||
alpha=evidence.alpha,
|
||||
support=cohesive_region.mask,
|
||||
minimum_strength=controls.minimum_strength,
|
||||
sensitivity=controls.split_sensitivity,
|
||||
)
|
||||
for partition in partitions:
|
||||
area = int(partition.count_nonzero().item())
|
||||
if area < controls.minimum_region_size:
|
||||
continue
|
||||
confidence = float(evidence.confidence[partition].mean().item())
|
||||
candidates.append(
|
||||
AttentionRegionComponent(
|
||||
evidence.label,
|
||||
evidence.alpha * partition.to(dtype=evidence.alpha.dtype),
|
||||
partition,
|
||||
_active_bounds(partition),
|
||||
area,
|
||||
confidence,
|
||||
)
|
||||
)
|
||||
if controls.keep_only == 0 or len(candidates) <= controls.keep_only:
|
||||
return tuple(candidates)
|
||||
ranked = sorted(
|
||||
candidates, key=lambda value: _rank_key(value, controls.keep_by)
|
||||
)
|
||||
return tuple(ranked[: controls.keep_only])
|
||||
|
||||
|
||||
def _rank_key(
|
||||
component: AttentionRegionComponent, keep_by: str
|
||||
) -> tuple[float, int, int]:
|
||||
"""Return deterministic descending evidence or area ordering."""
|
||||
|
||||
primary = (
|
||||
-component.confidence if keep_by == "highest confidence" else -component.area
|
||||
)
|
||||
return primary, component.bbox.top, component.bbox.left
|
||||
|
||||
|
||||
def _active_bounds(active: torch.Tensor) -> BoundingBox:
|
||||
"""Return the minimal full-image bounds around one active partition."""
|
||||
|
||||
coordinates = active.nonzero(as_tuple=False)
|
||||
if int(coordinates.shape[0]) < 1:
|
||||
raise ValueError("Attention instance partition requires an active pixel.")
|
||||
top = int(coordinates[:, 0].min().item())
|
||||
bottom = int(coordinates[:, 0].max().item()) + 1
|
||||
left = int(coordinates[:, 1].min().item())
|
||||
right = int(coordinates[:, 1].max().item()) + 1
|
||||
return BoundingBox(left, top, right, bottom)
|
||||
|
||||
|
||||
ATTENTION_COMPONENT_SERVICE = AttentionComponentService()
|
||||
@@ -0,0 +1,114 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Aggregate temporal attention observations into concept evidence."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
from ..domain.attention_geometry import factor_spatial_geometry
|
||||
from ..domain.attention_region_capture import AttentionRegionControls
|
||||
from ..domain.attention_region_maps import CapturedAttentionMap
|
||||
from .attention_spatial_projection import ATTENTION_SPATIAL_PROJECTION_SERVICE
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionConceptEvidence:
|
||||
"""Hold one concept's alpha, support, and pre-matte confidence evidence."""
|
||||
|
||||
label: str
|
||||
alpha: torch.Tensor
|
||||
support: torch.Tensor
|
||||
confidence: torch.Tensor
|
||||
|
||||
|
||||
class AttentionEvidenceAggregator:
|
||||
"""Combine selected observations without owning component or matte policy."""
|
||||
|
||||
def aggregate(
|
||||
self,
|
||||
*,
|
||||
maps: tuple[CapturedAttentionMap, ...],
|
||||
controls: AttentionRegionControls,
|
||||
height: int,
|
||||
width: int,
|
||||
batch_index: int,
|
||||
) -> tuple[AttentionConceptEvidence, ...]:
|
||||
"""Return ordered concept evidence for one output batch member."""
|
||||
|
||||
grouped: dict[str, list[CapturedAttentionMap]] = defaultdict(list)
|
||||
for attention_map in maps:
|
||||
if (
|
||||
attention_map.batch_index == batch_index
|
||||
and controls.capture_start
|
||||
<= attention_map.progress
|
||||
<= controls.capture_end
|
||||
):
|
||||
grouped[attention_map.label].append(attention_map)
|
||||
return tuple(
|
||||
_aggregate_group(label, tuple(observations), controls, height, width)
|
||||
for label, observations in grouped.items()
|
||||
)
|
||||
|
||||
|
||||
def _aggregate_group(
|
||||
label: str,
|
||||
observations: tuple[CapturedAttentionMap, ...],
|
||||
controls: AttentionRegionControls,
|
||||
height: int,
|
||||
width: int,
|
||||
) -> AttentionConceptEvidence:
|
||||
"""Combine strength and persistence for one readable concept."""
|
||||
|
||||
resized = torch.stack(
|
||||
tuple(_resize_observation(value, height, width) for value in observations)
|
||||
)
|
||||
normalized = resized / resized.amax(dim=(1, 2), keepdim=True).clamp_min(1e-12)
|
||||
weights = torch.tensor(
|
||||
[value.confidence for value in observations], dtype=resized.dtype
|
||||
).reshape(-1, 1, 1)
|
||||
strength = (resized * weights).sum(dim=0) / weights.sum().clamp_min(1e-6)
|
||||
strength = strength / strength.amax().clamp_min(1e-12)
|
||||
consensus = (normalized >= controls.minimum_strength).float().mean(dim=0)
|
||||
confidence = strength * consensus
|
||||
persistent = torch.where(
|
||||
consensus >= controls.minimum_consensus,
|
||||
confidence,
|
||||
torch.zeros_like(confidence),
|
||||
)
|
||||
maximum = persistent.amax()
|
||||
alpha = persistent if maximum <= 0.0 else persistent / maximum
|
||||
support = (alpha > 0.0) & (alpha >= controls.minimum_strength)
|
||||
return AttentionConceptEvidence(label, alpha * support, support, confidence)
|
||||
|
||||
|
||||
def _resize_observation(
|
||||
observation: CapturedAttentionMap, height: int, width: int
|
||||
) -> torch.Tensor:
|
||||
"""Resize one map using captured geometry or a compatibility fallback."""
|
||||
|
||||
source_height = observation.spatial_height
|
||||
source_width = observation.spatial_width
|
||||
if source_height is None or source_width is None:
|
||||
source_height, source_width = factor_spatial_geometry(
|
||||
int(observation.values.numel()), target_aspect=width / height
|
||||
)
|
||||
source = observation.values.float().reshape(1, 1, source_height, source_width)
|
||||
projected = ATTENTION_SPATIAL_PROJECTION_SERVICE.project(
|
||||
source[0, 0], observation.spatial_transforms
|
||||
)
|
||||
return functional.interpolate(
|
||||
projected.reshape(1, 1, int(projected.shape[0]), int(projected.shape[1])),
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)[0, 0]
|
||||
|
||||
|
||||
ATTENTION_EVIDENCE_AGGREGATOR = AttentionEvidenceAggregator()
|
||||
@@ -0,0 +1,66 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Partition cohesive attention support around independently strong peaks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
import torch
|
||||
|
||||
from ..masking.mask_components import connected_mask_components
|
||||
|
||||
|
||||
class AttentionInstanceSplittingService:
|
||||
"""Separate peak instances without deleting accepted concept support."""
|
||||
|
||||
def partition(
|
||||
self,
|
||||
*,
|
||||
alpha: torch.Tensor,
|
||||
support: torch.Tensor,
|
||||
minimum_strength: float,
|
||||
sensitivity: float,
|
||||
) -> tuple[torch.Tensor, ...]:
|
||||
"""Assign every supported pixel to a nearby peak-derived instance."""
|
||||
|
||||
active = support.detach().to(device="cpu", dtype=torch.bool)
|
||||
if sensitivity <= 0.0:
|
||||
return (active.to(device=support.device),)
|
||||
peak_threshold = minimum_strength + (
|
||||
sensitivity * (1.0 - minimum_strength) * 0.35
|
||||
)
|
||||
peak_support = active & (
|
||||
alpha.detach().to(device="cpu", dtype=torch.float32) >= peak_threshold
|
||||
)
|
||||
peaks = connected_mask_components(peak_support)
|
||||
if len(peaks) <= 1:
|
||||
return (active.to(device=support.device),)
|
||||
|
||||
source = torch.ones_like(active, dtype=torch.uint8)
|
||||
for peak in peaks:
|
||||
source[peak.mask] = 0
|
||||
cv2 = import_module("cv2")
|
||||
_distance, labels = cv2.distanceTransformWithLabels(
|
||||
source.numpy(),
|
||||
cv2.DIST_L2,
|
||||
5,
|
||||
labelType=cv2.DIST_LABEL_CCOMP,
|
||||
)
|
||||
active_labels = torch.from_numpy(labels).to(dtype=torch.int32)
|
||||
partitions: list[torch.Tensor] = []
|
||||
seen_labels: set[int] = set()
|
||||
for peak in peaks:
|
||||
label = int(active_labels[peak.mask][0].item())
|
||||
if label in seen_labels:
|
||||
continue
|
||||
seen_labels.add(label)
|
||||
partition = active & (active_labels == label)
|
||||
if partition.any():
|
||||
partitions.append(partition.to(device=support.device))
|
||||
return tuple(partitions) if partitions else (active.to(device=support.device),)
|
||||
|
||||
|
||||
ATTENTION_INSTANCE_SPLITTING_SERVICE = AttentionInstanceSplittingService()
|
||||
@@ -0,0 +1,86 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Shape retained attention alpha into controllable feathered mattes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
import torch
|
||||
|
||||
from ..masking.mask_components import connected_mask_components
|
||||
from ..masking.segs_mask_ops import feather_mask
|
||||
|
||||
|
||||
class AttentionMatteService:
|
||||
"""Blend raw attention alpha toward a solid, hole-filled matte."""
|
||||
|
||||
def shape(
|
||||
self,
|
||||
*,
|
||||
alpha: torch.Tensor,
|
||||
support: torch.Tensor,
|
||||
solidity: float,
|
||||
edge_feather: int,
|
||||
) -> torch.Tensor:
|
||||
"""Return one finite HW matte with solidity confined to accepted support."""
|
||||
|
||||
raw = alpha.float().clamp(0.0, 1.0)
|
||||
if solidity <= 0.0:
|
||||
return raw
|
||||
topology_radius = _topology_radius(support, edge_feather)
|
||||
cohesive = _close_narrow_channels(support, topology_radius)
|
||||
filled = _fill_enclosed_holes(cohesive).float()
|
||||
solid = feather_mask(filled, edge_feather)
|
||||
return torch.lerp(raw, solid, float(solidity)).clamp(0.0, 1.0)
|
||||
|
||||
|
||||
def _topology_radius(support: torch.Tensor, edge_feather: int) -> int:
|
||||
"""Scale narrow-channel cleanup conservatively with output resolution."""
|
||||
|
||||
height, width = int(support.shape[0]), int(support.shape[1])
|
||||
resolution_radius = round(min(height, width) * 0.03)
|
||||
return max(1, min(32, max(edge_feather, resolution_radius)))
|
||||
|
||||
|
||||
def _close_narrow_channels(support: torch.Tensor, radius: int) -> torch.Tensor:
|
||||
"""Close thin exterior branches without filling their broad source opening."""
|
||||
|
||||
height, width = int(support.shape[0]), int(support.shape[1])
|
||||
diameter = radius * 2 + 1
|
||||
if min(height, width) <= diameter:
|
||||
return support.detach().to(dtype=torch.bool)
|
||||
active = support.detach().to(device="cpu", dtype=torch.uint8).numpy()
|
||||
cv2 = import_module("cv2")
|
||||
kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (diameter, diameter))
|
||||
closed = cv2.morphologyEx(
|
||||
active,
|
||||
cv2.MORPH_CLOSE,
|
||||
kernel,
|
||||
borderType=cv2.BORDER_CONSTANT,
|
||||
borderValue=0,
|
||||
)
|
||||
cohesive = (closed > 0) | (active > 0)
|
||||
return torch.from_numpy(cohesive).to(device=support.device)
|
||||
|
||||
|
||||
def _fill_enclosed_holes(support: torch.Tensor) -> torch.Tensor:
|
||||
"""Fill inverse components that do not touch the image boundary."""
|
||||
|
||||
active = support.detach().to(device="cpu", dtype=torch.bool)
|
||||
inverse = ~active
|
||||
height, width = int(active.shape[0]), int(active.shape[1])
|
||||
filled = active
|
||||
for component in connected_mask_components(inverse):
|
||||
box = component.bbox
|
||||
touches_edge = (
|
||||
box.left == 0 or box.top == 0 or box.right == width or box.bottom == height
|
||||
)
|
||||
if not touches_edge:
|
||||
filled |= component.mask
|
||||
return filled.to(device=support.device)
|
||||
|
||||
|
||||
ATTENTION_MATTE_SERVICE = AttentionMatteService()
|
||||
@@ -0,0 +1,194 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Orchestrate downstream attention results for IMAGE and LATENT public nodes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from importlib import import_module
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.attention_region_capture import AttentionRegionControls
|
||||
from ..domain.attention_region_maps import CapturedAttentionMap
|
||||
from ..domain.segs import ImpactSegs, to_impact_compatible_segs
|
||||
from ..runtime.attention_region_status import ATTENTION_REGION_STATUS_PUBLISHER
|
||||
from ..runtime.attention_region_store import (
|
||||
ATTENTION_REGION_CAPTURE_STORE,
|
||||
AttentionCaptureSession,
|
||||
)
|
||||
from .attention_region_rendering import ATTENTION_REGION_RENDERING_SERVICE
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionImageResult:
|
||||
"""Return an unchanged image, batch SEGS, and batch mask."""
|
||||
|
||||
image: torch.Tensor
|
||||
segs: list[ImpactSegs]
|
||||
mask: torch.Tensor
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionLatentResult:
|
||||
"""Return an unchanged latent and attention-derived mask."""
|
||||
|
||||
latent: dict[str, object]
|
||||
mask: torch.Tensor
|
||||
|
||||
|
||||
class AttentionRegionNodeService:
|
||||
"""Consume one prompt-scoped session and render public node outputs."""
|
||||
|
||||
def for_image(
|
||||
self,
|
||||
*,
|
||||
request_node_id: str,
|
||||
image: object,
|
||||
controls: AttentionRegionControls,
|
||||
) -> AttentionImageResult:
|
||||
"""Render one SEGS payload per image batch member."""
|
||||
|
||||
image_batch = _image_batch(image)
|
||||
session, maps = _consume_maps(request_node_id)
|
||||
height = int(image_batch.shape[1])
|
||||
width = int(image_batch.shape[2])
|
||||
segs: list[ImpactSegs] = []
|
||||
masks: list[torch.Tensor] = []
|
||||
for batch_index in range(int(image_batch.shape[0])):
|
||||
native, mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=controls,
|
||||
height=height,
|
||||
width=width,
|
||||
image=image_batch[batch_index : batch_index + 1],
|
||||
batch_index=batch_index,
|
||||
)
|
||||
segs.append(to_impact_compatible_segs(native))
|
||||
masks.append(mask)
|
||||
_publish_status(request_node_id, session.status_message, maps)
|
||||
return AttentionImageResult(image_batch, segs, torch.cat(masks, dim=0))
|
||||
|
||||
def for_latent(
|
||||
self,
|
||||
*,
|
||||
request_node_id: str,
|
||||
latent: object,
|
||||
controls: AttentionRegionControls,
|
||||
) -> AttentionLatentResult:
|
||||
"""Render one latent-resolution mask per latent batch member."""
|
||||
|
||||
latent_value, samples = _latent_samples(latent)
|
||||
session, maps = _consume_maps(request_node_id)
|
||||
height = int(samples.shape[-2])
|
||||
width = int(samples.shape[-1])
|
||||
masks = tuple(
|
||||
ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=controls,
|
||||
height=height,
|
||||
width=width,
|
||||
batch_index=batch_index,
|
||||
)[1]
|
||||
for batch_index in range(int(samples.shape[0]))
|
||||
)
|
||||
mask = torch.cat(masks, dim=0)
|
||||
_publish_status(request_node_id, session.status_message, maps)
|
||||
return AttentionLatentResult(latent_value, mask)
|
||||
|
||||
def mask_conditioning(
|
||||
self,
|
||||
conditioning: object,
|
||||
mask: torch.Tensor,
|
||||
strength: float,
|
||||
) -> object:
|
||||
"""Attach one reusable mask through Comfy's generic conditioning policy."""
|
||||
|
||||
if isinstance(strength, bool) or not isinstance(strength, int | float):
|
||||
raise TypeError("Attention conditioning strength must be numeric.")
|
||||
if not 0.0 <= float(strength) <= 10.0:
|
||||
raise ValueError("Attention conditioning strength must be within 0..10.")
|
||||
node_helpers = import_module("node_helpers")
|
||||
setter = getattr(node_helpers, "conditioning_set_values", None)
|
||||
if not callable(setter):
|
||||
raise RuntimeError("Comfy conditioning mask support is unavailable.")
|
||||
return setter(
|
||||
conditioning,
|
||||
{
|
||||
"mask": mask,
|
||||
"set_area_to_bounds": False,
|
||||
"mask_strength": float(strength),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _consume_maps(
|
||||
request_node_id: str,
|
||||
) -> tuple[AttentionCaptureSession, tuple[CapturedAttentionMap, ...]]:
|
||||
"""Consume one request binding or create an observable empty session."""
|
||||
|
||||
session = ATTENTION_REGION_CAPTURE_STORE.consume(request_node_id)
|
||||
if session is None:
|
||||
from .attention_capture_model_service import EmptyAttentionCaptureSession
|
||||
|
||||
message = "No upstream attention capture was recovered"
|
||||
session = EmptyAttentionCaptureSession((request_node_id,), message)
|
||||
LOGGER.warning(
|
||||
"Attention-region public node executed without capture state",
|
||||
extra={"request_node_id": request_node_id},
|
||||
)
|
||||
return session, session.maps_for(request_node_id)
|
||||
|
||||
|
||||
def _publish_status(
|
||||
request_node_id: str,
|
||||
status: str,
|
||||
maps: tuple[CapturedAttentionMap, ...],
|
||||
) -> None:
|
||||
"""Publish capture status with an exact observation count."""
|
||||
|
||||
ATTENTION_REGION_STATUS_PUBLISHER.publish(
|
||||
request_node_id,
|
||||
f"{status}: {len(maps)} maps",
|
||||
)
|
||||
|
||||
|
||||
def _image_batch(value: object) -> torch.Tensor:
|
||||
"""Return a finite non-empty BHWC image tensor."""
|
||||
|
||||
if (
|
||||
not isinstance(value, torch.Tensor)
|
||||
or value.ndim != 4
|
||||
or int(value.shape[0]) < 1
|
||||
or int(value.shape[-1]) not in (1, 3, 4)
|
||||
or not torch.isfinite(value).all().item()
|
||||
):
|
||||
raise ValueError("Attention SEGS requires a finite non-empty BHWC IMAGE.")
|
||||
return value
|
||||
|
||||
|
||||
def _latent_samples(value: object) -> tuple[dict[str, object], torch.Tensor]:
|
||||
"""Return a latent mapping and supported BCHW or BCTHW sample tensor."""
|
||||
|
||||
if not isinstance(value, dict):
|
||||
raise TypeError("Attention Region Mask requires a LATENT dictionary.")
|
||||
samples = value.get("samples")
|
||||
if (
|
||||
not isinstance(samples, torch.Tensor)
|
||||
or samples.ndim not in (4, 5)
|
||||
or int(samples.shape[0]) < 1
|
||||
or (samples.ndim == 5 and int(samples.shape[-3]) != 1)
|
||||
):
|
||||
raise ValueError(
|
||||
"Attention Region Mask requires BCHW image latents or singleton-frame "
|
||||
"BCTHW Anima latents."
|
||||
)
|
||||
return value, samples
|
||||
|
||||
|
||||
ATTENTION_REGION_NODE_SERVICE = AttentionRegionNodeService()
|
||||
@@ -0,0 +1,140 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Orchestrate attention evidence, components, mattes, and SEGS packaging."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.attention_region_capture import AttentionRegionControls
|
||||
from ..domain.attention_region_maps import CapturedAttentionMap
|
||||
from ..domain.segs import BoundingBox, CropRegion, NativeSegs, Segment
|
||||
from .attention_region_components import ATTENTION_COMPONENT_SERVICE
|
||||
from .attention_region_evidence import ATTENTION_EVIDENCE_AGGREGATOR
|
||||
from .attention_region_matte import ATTENTION_MATTE_SERVICE
|
||||
|
||||
|
||||
class AttentionRegionRenderingService:
|
||||
"""Render captured maps through independently owned shaping policies."""
|
||||
|
||||
def render(
|
||||
self,
|
||||
*,
|
||||
maps: tuple[CapturedAttentionMap, ...],
|
||||
controls: AttentionRegionControls,
|
||||
height: int,
|
||||
width: int,
|
||||
image: torch.Tensor | None = None,
|
||||
batch_index: int = 0,
|
||||
) -> tuple[NativeSegs, torch.Tensor]:
|
||||
"""Return Impact-compatible SEGS and their full-image union mask."""
|
||||
|
||||
_validate_geometry(height, width, image)
|
||||
evidence_values = ATTENTION_EVIDENCE_AGGREGATOR.aggregate(
|
||||
maps=maps,
|
||||
controls=controls,
|
||||
height=height,
|
||||
width=width,
|
||||
batch_index=batch_index,
|
||||
)
|
||||
segments: list[Segment] = []
|
||||
union = torch.zeros((1, height, width), dtype=torch.float32)
|
||||
for evidence in evidence_values:
|
||||
components = ATTENTION_COMPONENT_SERVICE.retained(evidence, controls)
|
||||
shaped = tuple(
|
||||
(
|
||||
component,
|
||||
ATTENTION_MATTE_SERVICE.shape(
|
||||
alpha=component.alpha,
|
||||
support=component.support,
|
||||
solidity=controls.matte_solidity,
|
||||
edge_feather=controls.edge_feather,
|
||||
),
|
||||
)
|
||||
for component in components
|
||||
)
|
||||
if controls.combine_segs and shaped:
|
||||
combined_mask = torch.stack(tuple(mask for _item, mask in shaped)).amax(
|
||||
dim=0
|
||||
)
|
||||
confidence = max(item.confidence for item, _mask in shaped)
|
||||
segments.append(
|
||||
_segment_from_mask(evidence.label, combined_mask, confidence, image)
|
||||
)
|
||||
union = torch.maximum(union, combined_mask.unsqueeze(0))
|
||||
continue
|
||||
for component, matte in shaped:
|
||||
segments.append(
|
||||
_segment_from_mask(
|
||||
component.label, matte, component.confidence, image
|
||||
)
|
||||
)
|
||||
union = torch.maximum(union, matte.unsqueeze(0))
|
||||
return ((height, width), tuple(segments)), union
|
||||
|
||||
|
||||
def _segment_from_mask(
|
||||
label: str,
|
||||
full_mask: torch.Tensor,
|
||||
confidence: float,
|
||||
image: torch.Tensor | None,
|
||||
) -> Segment:
|
||||
"""Crop one shaped full-image matte into a native SEG."""
|
||||
|
||||
bbox = _active_bounds(full_mask > 0.0)
|
||||
region = CropRegion(bbox.left, bbox.top, bbox.right, bbox.bottom)
|
||||
cropped_mask = full_mask[region.top : region.bottom, region.left : region.right]
|
||||
cropped_image: object | None = None
|
||||
if image is not None:
|
||||
cropped_image = image[
|
||||
0, region.top : region.bottom, region.left : region.right, :
|
||||
]
|
||||
return Segment(
|
||||
cropped_image=cropped_image,
|
||||
cropped_mask=cropped_mask,
|
||||
confidence=max(0.0, min(1.0, confidence)),
|
||||
crop_region=region,
|
||||
bbox=bbox,
|
||||
label=label,
|
||||
)
|
||||
|
||||
|
||||
def _active_bounds(active: torch.Tensor) -> BoundingBox:
|
||||
"""Return the minimal box containing every nonzero matte pixel."""
|
||||
|
||||
coordinates = active.nonzero(as_tuple=False)
|
||||
if int(coordinates.shape[0]) < 1:
|
||||
raise ValueError("Attention-region bounds require an active pixel.")
|
||||
top = int(coordinates[:, 0].min().item())
|
||||
bottom = int(coordinates[:, 0].max().item()) + 1
|
||||
left = int(coordinates[:, 1].min().item())
|
||||
right = int(coordinates[:, 1].max().item()) + 1
|
||||
return BoundingBox(left, top, right, bottom)
|
||||
|
||||
|
||||
def _validate_geometry(
|
||||
height: int,
|
||||
width: int,
|
||||
image: torch.Tensor | None,
|
||||
) -> None:
|
||||
"""Require positive output geometry and an optional matching BHWC image."""
|
||||
|
||||
if type(height) is not int or height < 1 or type(width) is not int or width < 1:
|
||||
raise ValueError("Attention-region output dimensions must be positive.")
|
||||
if image is None:
|
||||
return
|
||||
if (
|
||||
not isinstance(image, torch.Tensor)
|
||||
or image.ndim != 4
|
||||
or int(image.shape[0]) < 1
|
||||
or int(image.shape[1]) != height
|
||||
or int(image.shape[2]) != width
|
||||
):
|
||||
raise ValueError(
|
||||
"Attention-region image must be a matching non-empty BHWC tensor."
|
||||
)
|
||||
|
||||
|
||||
ATTENTION_REGION_RENDERING_SERVICE = AttentionRegionRenderingService()
|
||||
@@ -0,0 +1,92 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Project captured attention through graph-visible full-canvas transforms."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
from ..domain.attention_spatial_transform import (
|
||||
AttentionSpatialTransform,
|
||||
AttentionSpatialTransformKind,
|
||||
)
|
||||
from ..image.resize_geometry import (
|
||||
CropPosition,
|
||||
ResizeMode,
|
||||
ResizeTarget,
|
||||
build_resize_plan,
|
||||
)
|
||||
|
||||
|
||||
class AttentionSpatialProjectionService:
|
||||
"""Apply ordered resize, crop, and pad geometry to one attention map."""
|
||||
|
||||
def project(
|
||||
self,
|
||||
values: torch.Tensor,
|
||||
transforms: tuple[AttentionSpatialTransform, ...],
|
||||
) -> torch.Tensor:
|
||||
"""Return one HW tensor transformed in graph-forward order."""
|
||||
|
||||
result = values.float()
|
||||
for transform in transforms:
|
||||
result = _apply_transform(result, transform)
|
||||
return result
|
||||
|
||||
|
||||
def _apply_transform(
|
||||
values: torch.Tensor, transform: AttentionSpatialTransform
|
||||
) -> torch.Tensor:
|
||||
"""Apply one validated transform using the product resize geometry policy."""
|
||||
|
||||
if transform.kind is AttentionSpatialTransformKind.SCALE:
|
||||
if transform.scale is None:
|
||||
raise ValueError("Attention scale transform is incomplete.")
|
||||
height = max(1, round(int(values.shape[0]) * transform.scale))
|
||||
width = max(1, round(int(values.shape[1]) * transform.scale))
|
||||
return _resize(values, height, width)
|
||||
if transform.width is None or transform.height is None:
|
||||
raise ValueError("Attention sized transform is incomplete.")
|
||||
mode = {
|
||||
AttentionSpatialTransformKind.RESIZE: ResizeMode.STRETCH,
|
||||
AttentionSpatialTransformKind.FIT_RESIZE: ResizeMode.KEEP_AR,
|
||||
AttentionSpatialTransformKind.COVER_CROP: ResizeMode.CROP,
|
||||
AttentionSpatialTransformKind.FIT_PAD: ResizeMode.PAD,
|
||||
}[transform.kind]
|
||||
plan = build_resize_plan(
|
||||
int(values.shape[1]),
|
||||
int(values.shape[0]),
|
||||
ResizeTarget(transform.width, transform.height, transform.divisible_by),
|
||||
mode,
|
||||
CropPosition(transform.anchor),
|
||||
)
|
||||
resized = _resize(values, plan.resize_height, plan.resize_width)
|
||||
if plan.has_crop:
|
||||
resized = resized[
|
||||
plan.crop_y : plan.crop_y + plan.output_height,
|
||||
plan.crop_x : plan.crop_x + plan.output_width,
|
||||
]
|
||||
if plan.has_pad:
|
||||
resized = functional.pad(
|
||||
resized,
|
||||
(plan.pad_left, plan.pad_right, plan.pad_top, plan.pad_bottom),
|
||||
value=0.0,
|
||||
)
|
||||
return resized
|
||||
|
||||
|
||||
def _resize(values: torch.Tensor, height: int, width: int) -> torch.Tensor:
|
||||
"""Resize one HW map with stable bilinear interpolation."""
|
||||
|
||||
return functional.interpolate(
|
||||
values.reshape(1, 1, int(values.shape[0]), int(values.shape[1])),
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)[0, 0]
|
||||
|
||||
|
||||
ATTENTION_SPATIAL_PROJECTION_SERVICE = AttentionSpatialProjectionService()
|
||||
@@ -0,0 +1,440 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test observation-only selected-token attention capture."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.attention_region_capture import (
|
||||
AttentionCapturePlan,
|
||||
AttentionCaptureProfile,
|
||||
AttentionRegionControls,
|
||||
AttentionRegionRequest,
|
||||
AttentionRegionRequestKind,
|
||||
)
|
||||
from simple_syrup.domain.attention_region_maps import (
|
||||
AttentionTokenCatalog,
|
||||
AttentionTokenSpan,
|
||||
OpenVocabularyContext,
|
||||
)
|
||||
from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily
|
||||
from simple_syrup.runtime.attention_region_affinity import (
|
||||
ATTENTION_AFFINITY_CALCULATOR,
|
||||
)
|
||||
from simple_syrup.runtime.attention_region_capture import AttentionRegionCaptureSession
|
||||
from simple_syrup.runtime.attention_region_capture_backend import (
|
||||
ATTENTION_REGION_CAPTURE_BACKEND,
|
||||
OptimizedAttentionCaptureOverride,
|
||||
)
|
||||
from simple_syrup.runtime.attention_region_open_vocabulary import (
|
||||
OpenVocabularyKeyProjector,
|
||||
)
|
||||
|
||||
|
||||
def test_capture_selects_positive_rows_and_exact_prompt_tokens() -> None:
|
||||
"""Capture one native concept without retaining a full attention matrix."""
|
||||
|
||||
session = _session(profile=AttentionCaptureProfile.EXHAUSTIVE)
|
||||
query = torch.tensor(
|
||||
[
|
||||
[[1.0, 0.0], [0.0, 1.0]],
|
||||
[[-1.0, 0.0], [0.0, -1.0]],
|
||||
]
|
||||
)
|
||||
key = torch.tensor(
|
||||
[
|
||||
[[0.0, 0.0], [1.0, 0.0], [0.0, 1.0]],
|
||||
[[0.0, 0.0], [-1.0, 0.0], [0.0, -1.0]],
|
||||
]
|
||||
)
|
||||
|
||||
session.observe(
|
||||
query,
|
||||
key,
|
||||
1,
|
||||
_options(cond_or_uncond=[0, 1]),
|
||||
skip_reshape=False,
|
||||
)
|
||||
|
||||
maps = session.maps_for("search")
|
||||
assert len(maps) == 1
|
||||
assert maps[0].label == "pink hair"
|
||||
assert maps[0].values.device.type == "cpu"
|
||||
assert maps[0].values[0].item() > maps[0].values[1].item()
|
||||
|
||||
|
||||
def test_capture_profile_subsamples_calls_without_changing_attention_output() -> None:
|
||||
"""Delegate every denoising call while retaining only fast-profile samples."""
|
||||
|
||||
session = _session(profile=AttentionCaptureProfile.FAST)
|
||||
override = OptimizedAttentionCaptureOverride(session, None)
|
||||
query = torch.ones(1, 2, 2)
|
||||
key = torch.ones(1, 3, 2)
|
||||
value = torch.arange(6, dtype=torch.float32).reshape(1, 3, 2)
|
||||
|
||||
def original(*args: object, **kwargs: object) -> torch.Tensor:
|
||||
"""Return a sentinel output while accepting Comfy attention arguments."""
|
||||
|
||||
del args, kwargs
|
||||
return torch.full((1, 2, 2), 7.0)
|
||||
|
||||
outputs = tuple(
|
||||
override(
|
||||
original,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
1,
|
||||
transformer_options=_options(cond_or_uncond=[0]),
|
||||
)
|
||||
for _index in range(17)
|
||||
)
|
||||
|
||||
assert all(torch.equal(output, outputs[0]) for output in outputs)
|
||||
assert outputs[0].eq(7.0).all().item()
|
||||
assert len(session.maps_for("search")) == 2
|
||||
|
||||
|
||||
def test_coalesced_requests_share_one_unioned_native_affinity_pass(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""Calculate shared prompt-token affinities once and route maps per request."""
|
||||
|
||||
controls = AttentionRegionControls(
|
||||
0.0,
|
||||
1.0,
|
||||
0.3,
|
||||
0.2,
|
||||
0.5,
|
||||
1,
|
||||
AttentionCaptureProfile.EXHAUSTIVE,
|
||||
)
|
||||
requests = (
|
||||
AttentionRegionRequest(
|
||||
"all",
|
||||
AttentionRegionRequestKind.ALL_PROMPT_SEGS,
|
||||
(),
|
||||
controls,
|
||||
),
|
||||
AttentionRegionRequest(
|
||||
"hair",
|
||||
AttentionRegionRequestKind.CONCEPT_SEGS,
|
||||
("pink hair",),
|
||||
controls,
|
||||
),
|
||||
)
|
||||
plan = AttentionCapturePlan(
|
||||
"sampler",
|
||||
"sampler",
|
||||
"model",
|
||||
("model", 0),
|
||||
("positive", 0),
|
||||
requests,
|
||||
"1girl, pink hair",
|
||||
("loader", 1),
|
||||
)
|
||||
spans = (
|
||||
AttentionTokenSpan("1girl", 1, (1,)),
|
||||
AttentionTokenSpan("pink hair", 1, (2,)),
|
||||
)
|
||||
session = AttentionRegionCaptureSession(
|
||||
plan=plan,
|
||||
model_family=RegionalModelFamily.STANDARD_UNET,
|
||||
token_catalog=AttentionTokenCatalog(4, spans, (0, 1, 2, 3)),
|
||||
request_spans={"hair": (spans[1],), "all": spans},
|
||||
)
|
||||
original = ATTENTION_AFFINITY_CALCULATOR.capture_spans
|
||||
calls: list[tuple[AttentionTokenSpan, ...]] = []
|
||||
|
||||
def record_capture(*args: Any, **kwargs: Any) -> Any:
|
||||
"""Record the unioned spans before delegating to real affinity math."""
|
||||
|
||||
captured_spans = kwargs.get("spans", args[2])
|
||||
calls.append(captured_spans)
|
||||
return original(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(ATTENTION_AFFINITY_CALCULATOR, "capture_spans", record_capture)
|
||||
session.observe(
|
||||
torch.ones(1, 2, 2),
|
||||
torch.ones(1, 4, 2),
|
||||
1,
|
||||
_options(cond_or_uncond=[0]),
|
||||
skip_reshape=False,
|
||||
)
|
||||
|
||||
assert calls == [spans]
|
||||
assert tuple(value.label for value in session.maps_for("hair")) == ("pink hair",)
|
||||
assert tuple(value.label for value in session.maps_for("all")) == (
|
||||
"1girl",
|
||||
"pink hair",
|
||||
)
|
||||
|
||||
|
||||
def test_backend_clones_model_and_composes_existing_attention_override() -> None:
|
||||
"""Keep source options intact while preserving an upstream override."""
|
||||
|
||||
source = _patcher()
|
||||
|
||||
def previous(
|
||||
original: object,
|
||||
*args: object,
|
||||
**kwargs: object,
|
||||
) -> torch.Tensor:
|
||||
"""Return a sentinel from an admitted upstream override."""
|
||||
|
||||
del original, args, kwargs
|
||||
return torch.tensor(3.0)
|
||||
|
||||
source.model_options["transformer_options"]["optimized_attention_override"] = (
|
||||
previous
|
||||
)
|
||||
|
||||
derived: Any = ATTENTION_REGION_CAPTURE_BACKEND.derive(source, _session())
|
||||
|
||||
assert derived.parent is source
|
||||
assert (
|
||||
source.model_options["transformer_options"]["optimized_attention_override"]
|
||||
is previous
|
||||
)
|
||||
installed = derived.model_options["transformer_options"][
|
||||
"optimized_attention_override"
|
||||
]
|
||||
assert isinstance(installed, OptimizedAttentionCaptureOverride)
|
||||
assert (
|
||||
installed(
|
||||
lambda *_args, **_kwargs: torch.tensor(1.0),
|
||||
torch.ones(1, 1, 1),
|
||||
torch.ones(1, 1, 1),
|
||||
torch.ones(1, 1, 1),
|
||||
1,
|
||||
).item()
|
||||
== 3.0
|
||||
)
|
||||
|
||||
|
||||
def test_open_vocabulary_query_projects_side_keys_without_changing_native_keys() -> (
|
||||
None
|
||||
):
|
||||
"""Reuse SDXL spatial queries for an absent phrase without conditioning edits."""
|
||||
|
||||
controls = AttentionRegionControls(
|
||||
0.0,
|
||||
1.0,
|
||||
0.3,
|
||||
0.2,
|
||||
0.5,
|
||||
1,
|
||||
AttentionCaptureProfile.EXHAUSTIVE,
|
||||
)
|
||||
request = AttentionRegionRequest(
|
||||
"search",
|
||||
AttentionRegionRequestKind.CONCEPT_SEGS,
|
||||
("head",),
|
||||
controls,
|
||||
)
|
||||
plan = AttentionCapturePlan(
|
||||
"sampler",
|
||||
"sampler",
|
||||
"model",
|
||||
("model", 0),
|
||||
("positive", 0),
|
||||
(request,),
|
||||
"1girl",
|
||||
("loader", 1),
|
||||
)
|
||||
context = OpenVocabularyContext(
|
||||
"head",
|
||||
torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 0.0]]]),
|
||||
(1,),
|
||||
)
|
||||
session = AttentionRegionCaptureSession(
|
||||
plan=plan,
|
||||
model_family=RegionalModelFamily.STANDARD_UNET,
|
||||
token_catalog=AttentionTokenCatalog(
|
||||
3, (AttentionTokenSpan("1girl", 1, (1,)),), (1, 2, 3)
|
||||
),
|
||||
request_spans={"search": ()},
|
||||
open_vocabulary_contexts=(context,),
|
||||
)
|
||||
projector = OpenVocabularyKeyProjector(torch.nn.Identity(), (context,), session)
|
||||
native_context = torch.tensor([[[2.0, 0.0], [0.0, 2.0], [0.0, 0.0]]])
|
||||
|
||||
native_keys = projector(native_context)
|
||||
session.observe(
|
||||
torch.tensor([[[1.0, 0.0], [0.0, 1.0]]]),
|
||||
native_keys,
|
||||
1,
|
||||
_options(cond_or_uncond=[0]),
|
||||
skip_reshape=False,
|
||||
)
|
||||
|
||||
assert torch.equal(native_keys, native_context)
|
||||
maps = session.maps_for("search")
|
||||
assert len(maps) == 1
|
||||
assert maps[0].label == "head"
|
||||
assert maps[0].values[0].item() > maps[0].values[1].item()
|
||||
|
||||
|
||||
def test_open_vocabulary_query_projection_is_cached_per_device_and_dtype() -> None:
|
||||
"""Avoid repeating absent-query key projections at every denoising call."""
|
||||
|
||||
class CountingProjection(torch.nn.Module):
|
||||
"""Count native and side-projection calls while preserving values."""
|
||||
|
||||
calls: int
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize an unused projection counter."""
|
||||
|
||||
super().__init__()
|
||||
self.calls = 0
|
||||
|
||||
def forward(self, value: torch.Tensor) -> torch.Tensor:
|
||||
"""Return the input after recording the projection."""
|
||||
|
||||
self.calls += 1
|
||||
return value
|
||||
|
||||
context = OpenVocabularyContext(
|
||||
"head",
|
||||
torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 0.0]]]),
|
||||
(1,),
|
||||
)
|
||||
projection = CountingProjection()
|
||||
projector = OpenVocabularyKeyProjector(projection, (context,), _session())
|
||||
native_context = torch.ones(1, 3, 2)
|
||||
|
||||
assert torch.equal(projector(native_context), native_context)
|
||||
assert torch.equal(projector(native_context), native_context)
|
||||
assert projection.calls == 3
|
||||
|
||||
|
||||
def test_sampled_denominator_bounds_unusually_strong_selected_keys() -> None:
|
||||
"""Keep fast-profile maps finite when denominator sampling misses the peak key."""
|
||||
|
||||
session = _session(
|
||||
profile=AttentionCaptureProfile.FAST,
|
||||
sequence_length=40,
|
||||
)
|
||||
query = torch.tensor([[[1000.0, 0.0], [0.0, 1.0]]])
|
||||
key = torch.zeros(1, 40, 2)
|
||||
key[0, 1, 0] = 1000.0
|
||||
|
||||
session.observe(
|
||||
query,
|
||||
key,
|
||||
1,
|
||||
_options(cond_or_uncond=[0]),
|
||||
skip_reshape=False,
|
||||
)
|
||||
|
||||
maps = session.maps_for("search")
|
||||
assert len(maps) == 1
|
||||
assert torch.isfinite(maps[0].values).all().item()
|
||||
assert maps[0].values.max().item() <= 1.0
|
||||
|
||||
|
||||
def test_graph_source_aspect_orients_dit_geometry_without_transformer_metadata() -> (
|
||||
None
|
||||
):
|
||||
"""Use the selected sampler's portrait input when a DiT omits shape metadata."""
|
||||
|
||||
session = _session(source_aspect=0.75)
|
||||
|
||||
session.observe(
|
||||
torch.ones(1, 12, 2),
|
||||
torch.ones(1, 3, 2),
|
||||
1,
|
||||
_options(cond_or_uncond=[0]),
|
||||
skip_reshape=False,
|
||||
)
|
||||
|
||||
attention_map = session.maps_for("search")[0]
|
||||
assert (attention_map.spatial_height, attention_map.spatial_width) == (4, 3)
|
||||
|
||||
|
||||
def test_anima_without_graph_or_runtime_geometry_fails_closed() -> None:
|
||||
"""Return no maps instead of guessing an unprojectable Anima grid orientation."""
|
||||
|
||||
session = _session(
|
||||
sequence_length=512,
|
||||
model_family=RegionalModelFamily.ANIMA,
|
||||
)
|
||||
|
||||
session.observe(
|
||||
torch.ones(1, 12, 2),
|
||||
torch.ones(1, 512, 2),
|
||||
1,
|
||||
_options(cond_or_uncond=[0]),
|
||||
skip_reshape=False,
|
||||
)
|
||||
|
||||
assert session.maps_for("search") == ()
|
||||
assert "not graph-visible" in session.status_message
|
||||
|
||||
|
||||
def _session(
|
||||
profile: AttentionCaptureProfile = AttentionCaptureProfile.EXHAUSTIVE,
|
||||
sequence_length: int = 3,
|
||||
source_aspect: float | None = None,
|
||||
model_family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET,
|
||||
) -> AttentionRegionCaptureSession:
|
||||
"""Return one exact-native-query capture session."""
|
||||
|
||||
controls = AttentionRegionControls(0.0, 1.0, 0.3, 0.2, 0.5, 1, profile)
|
||||
request = AttentionRegionRequest(
|
||||
"search",
|
||||
AttentionRegionRequestKind.CONCEPT_SEGS,
|
||||
("pink hair",),
|
||||
controls,
|
||||
)
|
||||
plan = AttentionCapturePlan(
|
||||
"sampler",
|
||||
"sampler",
|
||||
"model",
|
||||
("model", 0),
|
||||
("positive", 0),
|
||||
(request,),
|
||||
"pink hair",
|
||||
("loader", 1),
|
||||
source_aspect,
|
||||
)
|
||||
catalog = AttentionTokenCatalog(
|
||||
sequence_length,
|
||||
(AttentionTokenSpan("pink hair", 1, (1,)),),
|
||||
tuple(range(sequence_length)),
|
||||
)
|
||||
return AttentionRegionCaptureSession(
|
||||
plan=plan,
|
||||
model_family=model_family,
|
||||
token_catalog=catalog,
|
||||
request_spans={"search": catalog.spans},
|
||||
)
|
||||
|
||||
|
||||
def _options(*, cond_or_uncond: list[int]) -> dict[str, object]:
|
||||
"""Return exact sampler metadata at the beginning of denoising."""
|
||||
|
||||
return {
|
||||
"sample_sigmas": torch.tensor([1.0, 0.5, 0.0]),
|
||||
"sigmas": torch.tensor([1.0]),
|
||||
"cond_or_uncond": cond_or_uncond,
|
||||
"block": ("middle", 0),
|
||||
"block_index": 0,
|
||||
}
|
||||
|
||||
|
||||
def _patcher() -> Any:
|
||||
"""Create a real CPU ModelPatcher with isolated transformer options."""
|
||||
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
base_model = torch.nn.Module()
|
||||
base_model.diffusion_model = torch.nn.Linear(1, 1)
|
||||
device = torch.device("cpu")
|
||||
return ModelPatcher(base_model, load_device=device, offload_device=device)
|
||||
@@ -0,0 +1,368 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test prompt-level attention-region planning and graph rewriting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
from simple_syrup.domain.attention_region_capture import AttentionRegionRequestKind
|
||||
from simple_syrup.domain.attention_spatial_transform import (
|
||||
AttentionSpatialTransformKind,
|
||||
)
|
||||
from simple_syrup.runtime.attention_region_graph import (
|
||||
ATTENTION_REGION_PROMPT_PLANNER,
|
||||
ATTENTION_REGION_PROMPT_REWRITER,
|
||||
INTERNAL_CAPTURE_NODE_ID,
|
||||
)
|
||||
from simple_syrup.runtime.attention_region_plan_codec import (
|
||||
ATTENTION_CAPTURE_PLAN_CODEC,
|
||||
)
|
||||
|
||||
|
||||
class _ImagePass:
|
||||
"""Declare an exact IMAGE passthrough for provenance tests."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "image"}
|
||||
|
||||
|
||||
class _LatentPass:
|
||||
"""Declare an exact LATENT passthrough for provenance tests."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "latent"}
|
||||
|
||||
|
||||
def test_image_requests_coalesce_and_rewrite_one_direct_sampler_model_edge() -> None:
|
||||
"""Plan two downstream requests as one capture on their shared KSampler."""
|
||||
|
||||
prompt = _direct_prompt()
|
||||
plans = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {"ImagePass": _ImagePass})
|
||||
|
||||
assert len(plans) == 1
|
||||
plan = plans[0]
|
||||
assert plan.sampler_node_id == "sampler"
|
||||
assert plan.model_owner_node_id == "sampler"
|
||||
assert plan.model_link == ("loader", 0)
|
||||
assert plan.positive_link == ("positive", 0)
|
||||
assert plan.prompt_text == "1girl, pink hair, twintails, smug, selfie"
|
||||
assert plan.clip_link == ("loader", 1)
|
||||
assert tuple(request.kind for request in plan.requests) == (
|
||||
AttentionRegionRequestKind.ALL_PROMPT_SEGS,
|
||||
AttentionRegionRequestKind.CONCEPT_SEGS,
|
||||
)
|
||||
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(prompt, plans)
|
||||
|
||||
capture_id = plan.capture_node_id
|
||||
assert prompt["sampler"]["inputs"]["model"] == [capture_id, 0]
|
||||
capture = prompt[capture_id]
|
||||
assert capture["class_type"] == INTERNAL_CAPTURE_NODE_ID
|
||||
assert capture["inputs"]["model"] == ["loader", 0]
|
||||
assert capture["inputs"]["positive"] == ["positive", 0]
|
||||
assert capture["inputs"]["clip"] == ["loader", 1]
|
||||
serialized = json.loads(capture["inputs"]["plan_json"])
|
||||
assert [request["node_id"] for request in serialized["requests"]] == [
|
||||
"all",
|
||||
"search",
|
||||
]
|
||||
|
||||
|
||||
def test_latent_request_traces_through_declared_passthrough_to_advanced_guider() -> (
|
||||
None
|
||||
):
|
||||
"""Resolve the MODEL owner behind SamplerCustomAdvanced before execution."""
|
||||
|
||||
prompt: dict[str, Any] = {
|
||||
"model": {"class_type": "LoadModel", "inputs": {}},
|
||||
"positive": {"class_type": "Conditioning", "inputs": {}},
|
||||
"guider": {
|
||||
"class_type": "CFGGuider",
|
||||
"inputs": {
|
||||
"model": ["model", 0],
|
||||
"positive": ["positive", 0],
|
||||
"negative": ["negative", 0],
|
||||
},
|
||||
},
|
||||
"sampler": {
|
||||
"class_type": "SamplerCustomAdvanced",
|
||||
"inputs": {
|
||||
"guider": ["guider", 0],
|
||||
"latent_image": ["empty", 0],
|
||||
},
|
||||
},
|
||||
"empty": {"class_type": "EmptyLatentImage", "inputs": {}},
|
||||
"pass": {"class_type": "LatentPass", "inputs": {"latent": ["sampler", 1]}},
|
||||
"mask": _request_node(
|
||||
"SimpleSyrup.AttentionRegionMask", "latent", ["pass", 0], "head"
|
||||
),
|
||||
}
|
||||
|
||||
plans = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {"LatentPass": _LatentPass})
|
||||
|
||||
assert len(plans) == 1
|
||||
plan = plans[0]
|
||||
assert plan.sampler_node_id == "sampler"
|
||||
assert plan.model_owner_node_id == "guider"
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(prompt, plans)
|
||||
assert prompt["guider"]["inputs"]["model"] == [plan.capture_node_id, 0]
|
||||
|
||||
|
||||
def test_unsupported_provenance_is_observable_and_not_rewritten(caplog: Any) -> None:
|
||||
"""Leave opaque image paths untouched while logging the exact request id."""
|
||||
|
||||
prompt: dict[str, Any] = {
|
||||
"loaded": {"class_type": "LoadImage", "inputs": {}},
|
||||
"search": _request_node(
|
||||
"SimpleSyrup.ConceptAttentionSEGS", "image", ["loaded", 0], "head"
|
||||
),
|
||||
}
|
||||
|
||||
plans = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {})
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(prompt, plans)
|
||||
|
||||
assert plans == ()
|
||||
assert set(prompt) == {"loaded", "search"}
|
||||
assert "cannot recover a sampler" in caplog.text
|
||||
assert caplog.records[0].request_node_id == "search"
|
||||
|
||||
|
||||
def test_post_render_controls_do_not_invalidate_sampler_capture_cache() -> None:
|
||||
"""Keep upstream capture identity stable while users reshape saved maps."""
|
||||
|
||||
original = _direct_prompt()
|
||||
changed_shape = deepcopy(original)
|
||||
shape_inputs = changed_shape["search"]["inputs"]
|
||||
shape_inputs.update(
|
||||
{
|
||||
"minimum_strength": 0.91,
|
||||
"minimum_consensus": 0.82,
|
||||
"split_sensitivity": 0.73,
|
||||
"minimum_region_size": 99,
|
||||
"keep_only": 2,
|
||||
"keep_by": "highest confidence",
|
||||
"combine_segs": True,
|
||||
"matte_solidity": 1.0,
|
||||
"edge_feather": 24,
|
||||
}
|
||||
)
|
||||
changed_capture = deepcopy(original)
|
||||
changed_capture["search"]["inputs"]["capture_start"] = 0.42
|
||||
|
||||
original_plan = ATTENTION_REGION_PROMPT_PLANNER.build(
|
||||
original, {"ImagePass": _ImagePass}
|
||||
)
|
||||
shape_plan = ATTENTION_REGION_PROMPT_PLANNER.build(
|
||||
changed_shape, {"ImagePass": _ImagePass}
|
||||
)
|
||||
capture_plan = ATTENTION_REGION_PROMPT_PLANNER.build(
|
||||
changed_capture, {"ImagePass": _ImagePass}
|
||||
)
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(original, original_plan)
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(changed_shape, shape_plan)
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(changed_capture, capture_plan)
|
||||
|
||||
original_json = original[original_plan[0].capture_node_id]["inputs"]["plan_json"]
|
||||
shape_json = changed_shape[shape_plan[0].capture_node_id]["inputs"]["plan_json"]
|
||||
capture_json = changed_capture[capture_plan[0].capture_node_id]["inputs"][
|
||||
"plan_json"
|
||||
]
|
||||
assert shape_json == original_json
|
||||
assert capture_json != original_json
|
||||
|
||||
|
||||
def _direct_prompt() -> dict[str, Any]:
|
||||
"""Return one decoded-image graph containing two attention requests."""
|
||||
|
||||
prompt: dict[str, Any] = {
|
||||
"loader": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
|
||||
"positive": {
|
||||
"class_type": "CLIPTextEncode",
|
||||
"inputs": {
|
||||
"clip": ["loader", 1],
|
||||
"text": "1girl, pink hair, twintails, smug, selfie",
|
||||
},
|
||||
},
|
||||
"sampler": {
|
||||
"class_type": "KSampler",
|
||||
"inputs": {
|
||||
"model": ["loader", 0],
|
||||
"positive": ["positive", 0],
|
||||
"latent_image": ["empty", 0],
|
||||
},
|
||||
},
|
||||
"empty": {"class_type": "EmptyLatentImage", "inputs": {}},
|
||||
"decode": {
|
||||
"class_type": "VAEDecode",
|
||||
"inputs": {"samples": ["sampler", 0], "vae": ["loader", 2]},
|
||||
},
|
||||
"image_pass": {
|
||||
"class_type": "ImagePass",
|
||||
"inputs": {"image": ["decode", 0]},
|
||||
},
|
||||
"search": _request_node(
|
||||
"SimpleSyrup.ConceptAttentionSEGS",
|
||||
"image",
|
||||
["image_pass", 0],
|
||||
"pink hair | twintails",
|
||||
),
|
||||
"all": _request_node(
|
||||
"SimpleSyrup.AllPromptAttentionSEGS",
|
||||
"image",
|
||||
["decode", 0],
|
||||
"",
|
||||
),
|
||||
}
|
||||
return prompt
|
||||
|
||||
|
||||
def _request_node(
|
||||
class_type: str,
|
||||
source_name: str,
|
||||
source_link: list[object],
|
||||
queries: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Return one serialized public request with representative controls."""
|
||||
|
||||
return {
|
||||
"class_type": class_type,
|
||||
"inputs": {
|
||||
source_name: source_link,
|
||||
"concepts": queries,
|
||||
"sampler_stage": 1,
|
||||
"capture_start": 0.1,
|
||||
"capture_end": 0.8,
|
||||
"minimum_strength": 0.4,
|
||||
"minimum_consensus": 0.3,
|
||||
"split_sensitivity": 0.6,
|
||||
"minimum_region_size": 12,
|
||||
"capture_profile": "balanced",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_sampler_stage_selects_first_direct_or_clamped_lineage_stage() -> None:
|
||||
"""Resolve one-based sampler stages through decode, resize, and encode."""
|
||||
|
||||
prompt = _two_stage_prompt()
|
||||
first = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {})
|
||||
assert tuple(plan.sampler_node_id for plan in first) == ("first",)
|
||||
assert tuple(
|
||||
transform.kind for transform in first[0].requests[0].spatial_transforms
|
||||
) == (AttentionSpatialTransformKind.SCALE,)
|
||||
ATTENTION_REGION_PROMPT_REWRITER.rewrite(prompt, first)
|
||||
serialized = prompt[first[0].capture_node_id]["inputs"]["plan_json"]
|
||||
decoded = ATTENTION_CAPTURE_PLAN_CODEC.decode(serialized)
|
||||
assert decoded.requests[0].spatial_transforms == (
|
||||
first[0].requests[0].spatial_transforms
|
||||
)
|
||||
assert decoded.source_aspect == 768 / 896
|
||||
|
||||
prompt["concept"]["inputs"]["sampler_stage"] = 0
|
||||
direct = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {})
|
||||
assert tuple(plan.sampler_node_id for plan in direct) == ("second",)
|
||||
|
||||
prompt["concept"]["inputs"]["sampler_stage"] = 99
|
||||
clamped = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {})
|
||||
assert tuple(plan.sampler_node_id for plan in clamped) == ("second",)
|
||||
|
||||
|
||||
def test_generic_latent_sampler_inputs_are_compatible_without_class_allowlist() -> None:
|
||||
"""Patch a custom latent sampler that follows Comfy's standard input contract."""
|
||||
|
||||
prompt = _direct_prompt()
|
||||
prompt["sampler"]["class_type"] = "ThirdPartySampler"
|
||||
|
||||
plans = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {"ImagePass": _ImagePass})
|
||||
|
||||
assert tuple(plan.sampler_node_id for plan in plans) == ("sampler",)
|
||||
|
||||
|
||||
def test_crop_local_detailer_remains_in_lineage_but_selected_capture_is_no_op(
|
||||
caplog: Any,
|
||||
) -> None:
|
||||
"""Allow earlier-stage selection through a detailer without misregistered maps."""
|
||||
|
||||
prompt = _detailer_chain_prompt()
|
||||
first = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {})
|
||||
assert tuple(plan.sampler_node_id for plan in first) == ("sampler",)
|
||||
|
||||
prompt["search"]["inputs"]["sampler_stage"] = 2
|
||||
crop_local = ATTENTION_REGION_PROMPT_PLANNER.build(prompt, {})
|
||||
|
||||
assert crop_local == ()
|
||||
assert "crop-local images" in caplog.records[-1].reason
|
||||
|
||||
|
||||
def _two_stage_prompt() -> dict[str, Any]:
|
||||
"""Return a full-canvas hires lineage with two explicit samplers."""
|
||||
|
||||
return {
|
||||
"loader": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
|
||||
"positive": {
|
||||
"class_type": "CLIPTextEncode",
|
||||
"inputs": {"clip": ["loader", 1], "text": "girl, pink hair"},
|
||||
},
|
||||
"empty": {
|
||||
"class_type": "EmptyLatentImage",
|
||||
"inputs": {"width": 768, "height": 896, "batch_size": 1},
|
||||
},
|
||||
"first": {
|
||||
"class_type": "KSampler",
|
||||
"inputs": {
|
||||
"model": ["loader", 0],
|
||||
"positive": ["positive", 0],
|
||||
"latent_image": ["empty", 0],
|
||||
},
|
||||
},
|
||||
"decode_first": {
|
||||
"class_type": "VAEDecode",
|
||||
"inputs": {"samples": ["first", 0], "vae": ["loader", 2]},
|
||||
},
|
||||
"scale": {
|
||||
"class_type": "ImageScaleBy",
|
||||
"inputs": {"image": ["decode_first", 0], "scale_by": 2.0},
|
||||
},
|
||||
"encode": {
|
||||
"class_type": "VAEEncode",
|
||||
"inputs": {"pixels": ["scale", 0], "vae": ["loader", 2]},
|
||||
},
|
||||
"second": {
|
||||
"class_type": "KSamplerAdvanced",
|
||||
"inputs": {
|
||||
"model": ["loader", 0],
|
||||
"positive": ["positive", 0],
|
||||
"latent_image": ["encode", 0],
|
||||
},
|
||||
},
|
||||
"decode_second": {
|
||||
"class_type": "VAEDecode",
|
||||
"inputs": {"samples": ["second", 0], "vae": ["loader", 2]},
|
||||
},
|
||||
"concept": _request_node(
|
||||
"SimpleSyrup.ConceptAttentionSEGS",
|
||||
"image",
|
||||
["decode_second", 0],
|
||||
"pink hair",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _detailer_chain_prompt() -> dict[str, Any]:
|
||||
"""Return one ordinary sampler followed by a crop-local detailer stage."""
|
||||
|
||||
prompt = _direct_prompt()
|
||||
prompt["detailer"] = {
|
||||
"class_type": "SimpleSyrup.DetailSEGSByScaleFactor",
|
||||
"inputs": {
|
||||
"image": ["decode", 0],
|
||||
"model": ["loader", 0],
|
||||
"positive": ["positive", 0],
|
||||
},
|
||||
}
|
||||
prompt["search"]["inputs"]["image"] = ["detailer", 0]
|
||||
del prompt["all"]
|
||||
return prompt
|
||||
@@ -0,0 +1,40 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test safe one-time attention-region prompt handler registration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import simple_syrup.runtime.attention_region_prompt_handler as handler_module
|
||||
|
||||
|
||||
class _PromptServer:
|
||||
"""Record prompt handlers without importing Comfy's HTTP server."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create an empty handler list."""
|
||||
|
||||
self.handlers: list[object] = []
|
||||
|
||||
def add_on_prompt_handler(self, handler: object) -> None:
|
||||
"""Record one registered handler."""
|
||||
|
||||
self.handlers.append(handler)
|
||||
|
||||
|
||||
def test_prompt_handler_registers_once(monkeypatch: Any) -> None:
|
||||
"""Prevent duplicate graph rewrites across extension reload paths."""
|
||||
|
||||
instance = _PromptServer()
|
||||
fake_module = type("ServerModule", (), {})()
|
||||
fake_module.PromptServer = type("PromptServerType", (), {"instance": instance})
|
||||
monkeypatch.setitem(sys.modules, "server", fake_module)
|
||||
|
||||
handler_module.register_attention_region_prompt_handler()
|
||||
handler_module.register_attention_region_prompt_handler()
|
||||
|
||||
assert len(instance.handlers) == 1
|
||||
@@ -0,0 +1,379 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test attention-native temporal shaping and soft SEGS construction."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.attention_region_capture import (
|
||||
AttentionCaptureProfile,
|
||||
AttentionRegionControls,
|
||||
)
|
||||
from simple_syrup.domain.attention_region_maps import CapturedAttentionMap
|
||||
from simple_syrup.services.attention_region_matte import ATTENTION_MATTE_SERVICE
|
||||
from simple_syrup.services.attention_region_rendering import (
|
||||
ATTENTION_REGION_RENDERING_SERVICE,
|
||||
)
|
||||
|
||||
|
||||
def test_strength_and_consensus_shape_soft_regions_before_component_packaging() -> None:
|
||||
"""Reject transient weak pixels while preserving feathered accepted values."""
|
||||
|
||||
maps = (
|
||||
_map("hair", [0.1, 0.8, 0.2, 0.7], 0.2),
|
||||
_map("hair", [0.1, 0.9, 0.1, 0.2], 0.5),
|
||||
_map("hair", [0.1, 0.7, 0.1, 0.1], 0.8),
|
||||
)
|
||||
image = torch.zeros(1, 2, 2, 3)
|
||||
|
||||
segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(strength=0.5, consensus=0.66),
|
||||
height=2,
|
||||
width=2,
|
||||
image=image,
|
||||
)
|
||||
|
||||
assert len(segs[1]) == 1
|
||||
assert segs[1][0].label == "hair"
|
||||
assert isinstance(segs[1][0].cropped_mask, torch.Tensor)
|
||||
assert segs[1][0].cropped_mask.shape == (1, 1)
|
||||
assert mask[0, 0, 1].item() == 1.0
|
||||
assert mask.count_nonzero().item() == 1
|
||||
|
||||
|
||||
def test_temporal_window_changes_region_without_morphology() -> None:
|
||||
"""Select early versus late attention evidence through normalized progress."""
|
||||
|
||||
maps = (
|
||||
_map("composition", [1.0, 1.0, 0.0, 0.0], 0.1),
|
||||
_map("composition", [0.0, 0.0, 1.0, 1.0], 0.9),
|
||||
)
|
||||
|
||||
_early_segs, early = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(start=0.0, end=0.4, strength=0.2, consensus=0.0),
|
||||
height=2,
|
||||
width=2,
|
||||
)
|
||||
_late_segs, late = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(start=0.6, end=1.0, strength=0.2, consensus=0.0),
|
||||
height=2,
|
||||
width=2,
|
||||
)
|
||||
|
||||
assert early[0, 0].sum().item() > early[0, 1].sum().item()
|
||||
assert late[0, 1].sum().item() > late[0, 0].sum().item()
|
||||
|
||||
|
||||
def test_empty_maps_return_correctly_sized_no_op_outputs() -> None:
|
||||
"""Return empty SEGS and a zero mask for unsupported graph/model paths."""
|
||||
|
||||
segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=(),
|
||||
controls=_controls(),
|
||||
height=5,
|
||||
width=7,
|
||||
)
|
||||
|
||||
assert segs == ((5, 7), ())
|
||||
assert mask.shape == (1, 5, 7)
|
||||
assert mask.count_nonzero().item() == 0
|
||||
|
||||
|
||||
def test_disconnected_attention_islands_become_separate_instances() -> None:
|
||||
"""Expose retained disconnected objects as independently controllable SEGS."""
|
||||
|
||||
segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=(_map("outdoors", [1.0, 0.0, 1.0], 0.5),),
|
||||
controls=_controls(strength=0.5, consensus=0.0),
|
||||
height=1,
|
||||
width=3,
|
||||
)
|
||||
|
||||
assert len(segs[1]) == 2
|
||||
assert tuple(segment.label for segment in segs[1]) == ("outdoors", "outdoors")
|
||||
assert mask.count_nonzero().item() == 2
|
||||
|
||||
|
||||
def test_split_sensitivity_preserves_the_complete_concept_union() -> None:
|
||||
"""Change instance separation without deleting moderate concept support."""
|
||||
|
||||
maps = (_map("cat", [1.0, 0.4, 0.4, 0.4, 1.0], 0.5),)
|
||||
_unsplit_segs, unsplit = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(strength=0.3, consensus=0.0, split=0.0),
|
||||
height=1,
|
||||
width=5,
|
||||
)
|
||||
split_segs, split = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(strength=0.3, consensus=0.0, split=1.0),
|
||||
height=1,
|
||||
width=5,
|
||||
)
|
||||
|
||||
assert torch.equal(split, unsplit)
|
||||
assert split.count_nonzero().item() == 5
|
||||
assert len(split_segs[1]) == 2
|
||||
|
||||
|
||||
def test_cohesive_support_retains_a_moderate_body_around_a_strong_core() -> None:
|
||||
"""Keep the complete attended silhouette instead of its sparse core pixels."""
|
||||
|
||||
values = [
|
||||
0.0,
|
||||
0.2,
|
||||
0.2,
|
||||
0.2,
|
||||
0.0,
|
||||
0.2,
|
||||
0.4,
|
||||
0.6,
|
||||
0.4,
|
||||
0.2,
|
||||
0.2,
|
||||
0.6,
|
||||
1.0,
|
||||
0.6,
|
||||
0.2,
|
||||
0.2,
|
||||
0.4,
|
||||
0.6,
|
||||
0.4,
|
||||
0.2,
|
||||
0.0,
|
||||
0.2,
|
||||
0.2,
|
||||
0.2,
|
||||
0.0,
|
||||
]
|
||||
|
||||
_segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=(_map("hair", values, 0.5),),
|
||||
controls=_controls(strength=0.15, consensus=0.0, split=0.75),
|
||||
height=5,
|
||||
width=5,
|
||||
)
|
||||
|
||||
assert mask.count_nonzero().item() == 21
|
||||
assert mask[0, 2, 2].item() == 1.0
|
||||
assert mask[0, 0, 1].item() > 0.0
|
||||
|
||||
|
||||
def test_minimum_region_size_removes_each_pockmark_independently() -> None:
|
||||
"""Discard a tiny island without rejecting or merging the valid component."""
|
||||
|
||||
values = [1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0]
|
||||
segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=(_map("hair", values, 0.5),),
|
||||
controls=_controls(strength=0.5, consensus=0.0, minimum_size=2),
|
||||
height=3,
|
||||
width=3,
|
||||
)
|
||||
|
||||
assert len(segs[1]) == 1
|
||||
assert segs[1][0].bbox == (0, 0, 2, 1)
|
||||
assert mask.count_nonzero().item() == 2
|
||||
|
||||
|
||||
def test_keep_only_and_combine_apply_per_concept_without_changing_union() -> None:
|
||||
"""Retain top instances and optionally package their union as one SEG."""
|
||||
|
||||
maps = (_map("cat", [1.0, 1.0, 0.0, 0.8, 0.8], 0.5),)
|
||||
separate, separate_mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(
|
||||
strength=0.5,
|
||||
consensus=0.0,
|
||||
keep_only=2,
|
||||
combine=False,
|
||||
),
|
||||
height=1,
|
||||
width=5,
|
||||
)
|
||||
combined, combined_mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(
|
||||
strength=0.5,
|
||||
consensus=0.0,
|
||||
keep_only=2,
|
||||
combine=True,
|
||||
),
|
||||
height=1,
|
||||
width=5,
|
||||
)
|
||||
|
||||
assert len(separate[1]) == 2
|
||||
assert len(combined[1]) == 1
|
||||
assert torch.equal(separate_mask, combined_mask)
|
||||
|
||||
|
||||
def test_matte_solidity_flattens_interior_and_edge_feather_softens_boundary() -> None:
|
||||
"""Preserve raw alpha at zero and make a solid feathered matte at one."""
|
||||
|
||||
maps = (_map("dress", [0.0, 0.6, 1.0, 0.7, 0.0], 0.5),)
|
||||
_raw_segs, raw = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(strength=0.2, consensus=0.0, solidity=0.0),
|
||||
height=1,
|
||||
width=5,
|
||||
)
|
||||
_solid_segs, solid = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=maps,
|
||||
controls=_controls(
|
||||
strength=0.2,
|
||||
consensus=0.0,
|
||||
solidity=1.0,
|
||||
feather=1,
|
||||
),
|
||||
height=1,
|
||||
width=5,
|
||||
)
|
||||
|
||||
assert raw[0, 0, 1].item() != raw[0, 0, 2].item()
|
||||
assert solid[0, 0, 2].item() == 1.0
|
||||
assert 0.0 < solid[0, 0, 0].item() < 1.0
|
||||
|
||||
|
||||
def test_full_matte_solidity_fills_only_enclosed_holes() -> None:
|
||||
"""Fill an interior attention gap without filling exterior-connected space."""
|
||||
|
||||
ring = [1.0, 1.0, 1.0, 1.0, 0.0, 1.0, 1.0, 1.0, 1.0]
|
||||
_segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=(_map("hair", ring, 0.5),),
|
||||
controls=_controls(
|
||||
strength=0.5,
|
||||
consensus=0.0,
|
||||
solidity=1.0,
|
||||
feather=0,
|
||||
),
|
||||
height=3,
|
||||
width=3,
|
||||
)
|
||||
|
||||
assert torch.equal(mask, torch.ones_like(mask))
|
||||
|
||||
|
||||
def test_full_matte_solidity_closes_a_narrow_exterior_connected_channel() -> None:
|
||||
"""Remove thin attention squiggles without filling broad exterior space."""
|
||||
|
||||
support = torch.zeros(7, 7, dtype=torch.bool)
|
||||
support[1:6, 1:6] = True
|
||||
support[1:4, 3] = False
|
||||
|
||||
matte = ATTENTION_MATTE_SERVICE.shape(
|
||||
alpha=support.float(),
|
||||
support=support,
|
||||
solidity=1.0,
|
||||
edge_feather=0,
|
||||
)
|
||||
|
||||
assert matte[3, 3].item() == 1.0
|
||||
assert matte[0].count_nonzero().item() == 0
|
||||
assert matte[:, 0].count_nonzero().item() == 0
|
||||
|
||||
|
||||
def test_full_matte_solidity_removes_a_winding_channel_from_a_large_opening() -> None:
|
||||
"""Smooth thin branches while preserving the large excluded exterior area."""
|
||||
|
||||
support = torch.ones(100, 100, dtype=torch.bool)
|
||||
support[55:, 25:75] = False
|
||||
support[35:55, 49:52] = False
|
||||
support[42:45, 42:52] = False
|
||||
|
||||
matte = ATTENTION_MATTE_SERVICE.shape(
|
||||
alpha=support.float(),
|
||||
support=support,
|
||||
solidity=1.0,
|
||||
edge_feather=0,
|
||||
)
|
||||
|
||||
assert matte[38, 50].item() == 1.0
|
||||
assert matte[43, 44].item() == 1.0
|
||||
assert matte[80, 50].item() == 0.0
|
||||
|
||||
|
||||
def test_full_matte_solidity_never_erases_accepted_thin_support() -> None:
|
||||
"""Keep every accepted pixel when topology cleanup adds cohesive support."""
|
||||
|
||||
support = torch.zeros(100, 100, dtype=torch.bool)
|
||||
support[10:90, 50] = True
|
||||
|
||||
matte = ATTENTION_MATTE_SERVICE.shape(
|
||||
alpha=support.float(),
|
||||
support=support,
|
||||
solidity=1.0,
|
||||
edge_feather=0,
|
||||
)
|
||||
|
||||
assert torch.all(matte[support] == 1.0)
|
||||
|
||||
|
||||
def test_highest_confidence_keeps_stronger_component() -> None:
|
||||
"""Rank components from pre-normalized evidence rather than normalized maxima."""
|
||||
|
||||
segs, _mask = ATTENTION_REGION_RENDERING_SERVICE.render(
|
||||
maps=(_map("cat", [1.0, 0.0, 0.55], 0.5),),
|
||||
controls=AttentionRegionControls(
|
||||
0.0,
|
||||
1.0,
|
||||
0.4,
|
||||
0.0,
|
||||
0.0,
|
||||
1,
|
||||
AttentionCaptureProfile.BALANCED,
|
||||
keep_only=1,
|
||||
keep_by="highest confidence",
|
||||
),
|
||||
height=1,
|
||||
width=3,
|
||||
)
|
||||
|
||||
assert len(segs[1]) == 1
|
||||
assert segs[1][0].bbox == (0, 0, 1, 1)
|
||||
|
||||
|
||||
def _map(label: str, values: list[float], progress: float) -> CapturedAttentionMap:
|
||||
"""Create one compact square attention observation."""
|
||||
|
||||
return CapturedAttentionMap(
|
||||
label,
|
||||
torch.tensor(values, dtype=torch.float16),
|
||||
progress,
|
||||
"layer",
|
||||
)
|
||||
|
||||
|
||||
def _controls(
|
||||
*,
|
||||
start: float = 0.0,
|
||||
end: float = 1.0,
|
||||
strength: float = 0.35,
|
||||
consensus: float = 0.25,
|
||||
minimum_size: int = 1,
|
||||
keep_only: int = 0,
|
||||
combine: bool = False,
|
||||
solidity: float = 0.0,
|
||||
feather: int = 8,
|
||||
split: float = 0.0,
|
||||
) -> AttentionRegionControls:
|
||||
"""Return representative balanced rendering controls."""
|
||||
|
||||
return AttentionRegionControls(
|
||||
start,
|
||||
end,
|
||||
strength,
|
||||
consensus,
|
||||
split,
|
||||
minimum_size,
|
||||
AttentionCaptureProfile.BALANCED,
|
||||
keep_only=keep_only,
|
||||
combine_segs=combine,
|
||||
matte_solidity=solidity,
|
||||
edge_feather=feather,
|
||||
)
|
||||
@@ -0,0 +1,53 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test attention capture session publication and cleanup ownership."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.domain.attention_region_maps import CapturedAttentionMap
|
||||
from simple_syrup.runtime.attention_region_store import (
|
||||
AttentionCaptureSession,
|
||||
AttentionRegionCaptureStore,
|
||||
)
|
||||
|
||||
|
||||
class _Session:
|
||||
"""Provide the minimal downstream capture protocol."""
|
||||
|
||||
@property
|
||||
def status_message(self) -> str:
|
||||
"""Return a stable lifecycle-test status."""
|
||||
|
||||
return "test session"
|
||||
|
||||
def maps_for(self, request_node_id: str) -> tuple[CapturedAttentionMap, ...]:
|
||||
"""Return no maps for a lifecycle-only test session."""
|
||||
|
||||
del request_node_id
|
||||
return ()
|
||||
|
||||
|
||||
def test_shared_session_is_consumed_once_per_coalesced_request() -> None:
|
||||
"""Keep sibling request bindings alive until each public node executes."""
|
||||
|
||||
store = AttentionRegionCaptureStore()
|
||||
session: AttentionCaptureSession = _Session()
|
||||
store.publish(("all", "search"), session)
|
||||
|
||||
assert store.consume("search") is session
|
||||
assert store.consume("search") is None
|
||||
assert store.consume("all") is session
|
||||
|
||||
|
||||
def test_store_rejects_overlapping_live_prompt_state() -> None:
|
||||
"""Detect request-id collisions rather than routing maps across executions."""
|
||||
|
||||
store = AttentionRegionCaptureStore()
|
||||
store.publish(("search",), _Session())
|
||||
|
||||
with pytest.raises(RuntimeError, match="already active"):
|
||||
store.publish(("search",), _Session())
|
||||
@@ -0,0 +1,106 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test readable concept alignment to supported denoiser token streams."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from simple_syrup.domain.attention_concepts import parse_attention_concepts
|
||||
from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily
|
||||
from simple_syrup.runtime.attention_region_tokens import ATTENTION_PROMPT_TOKENIZER
|
||||
|
||||
|
||||
class _FakeClip:
|
||||
"""Tokenize comma tags into predictable model-family streams."""
|
||||
|
||||
def tokenize(self, text: str, return_word_ids: bool = False) -> dict[str, object]:
|
||||
"""Return special-delimited integer character tokens for both families."""
|
||||
|
||||
del return_word_ids
|
||||
content = [(ord(character), 1.0, 1) for character in text]
|
||||
stream = [[(1, 1.0, 0), *content, (2, 1.0, 0)]]
|
||||
return {"g": stream, "l": stream, "t5xxl": stream}
|
||||
|
||||
|
||||
def test_sdxl_catalog_preserves_repeated_readable_prompt_occurrences() -> None:
|
||||
"""Keep repeated comma concepts separate while advancing token positions."""
|
||||
|
||||
catalog = ATTENTION_PROMPT_TOKENIZER.build_catalog(
|
||||
clip=_FakeClip(),
|
||||
prompt_text="1girl, pink hair, pink hair",
|
||||
model_family=RegionalModelFamily.STANDARD_UNET,
|
||||
)
|
||||
|
||||
assert tuple(span.display_label for span in catalog.spans) == (
|
||||
"1girl",
|
||||
"pink hair",
|
||||
"pink hair #2",
|
||||
)
|
||||
assert catalog.spans[1].token_indices != catalog.spans[2].token_indices
|
||||
|
||||
|
||||
def test_anima_catalog_selects_adapted_t5_token_positions() -> None:
|
||||
"""Use Anima's adapted T5 stream instead of its Qwen source tokens."""
|
||||
|
||||
catalog = ATTENTION_PROMPT_TOKENIZER.build_catalog(
|
||||
clip=_FakeClip(),
|
||||
prompt_text="smug, selfie",
|
||||
model_family=RegionalModelFamily.ANIMA,
|
||||
)
|
||||
|
||||
assert tuple(span.label for span in catalog.spans) == ("smug", "selfie")
|
||||
assert catalog.sequence_length == len("smug, selfie") + 2
|
||||
|
||||
|
||||
def test_pipe_parser_preserves_commas_and_explicit_concept_order() -> None:
|
||||
"""Treat only vertical bars as concept boundaries."""
|
||||
|
||||
assert parse_attention_concepts(
|
||||
" 1girl, pink hair black dress | cat || smug expression "
|
||||
) == ("1girl, pink hair black dress", "cat", "smug expression")
|
||||
|
||||
|
||||
def test_distributed_query_uses_complete_native_spans_in_any_prompt_order() -> None:
|
||||
"""Cover a composite query through longest native spans across prompt entries."""
|
||||
|
||||
clip = _FakeClip()
|
||||
catalog = ATTENTION_PROMPT_TOKENIZER.build_catalog(
|
||||
clip=clip,
|
||||
prompt_text="black dress, outdoors, 1girl, city, pink hair",
|
||||
model_family=RegionalModelFamily.STANDARD_UNET,
|
||||
)
|
||||
|
||||
spans = ATTENTION_PROMPT_TOKENIZER.resolve_query(
|
||||
clip=clip,
|
||||
query="1girl, pink hair black dress",
|
||||
catalog=catalog,
|
||||
model_family=RegionalModelFamily.STANDARD_UNET,
|
||||
)
|
||||
|
||||
assert tuple(span.label for span in spans) == ("1girl, pink hair black dress",)
|
||||
resolved_ids = tuple(catalog.token_ids[index] for index in spans[0].token_indices)
|
||||
assert all(
|
||||
ord(character) in resolved_ids for character in "1girlpinkhairblackdress"
|
||||
)
|
||||
|
||||
|
||||
def test_distributed_query_requires_complete_native_coverage() -> None:
|
||||
"""Return no native span when any meaningful phrase is absent."""
|
||||
|
||||
clip = _FakeClip()
|
||||
catalog = ATTENTION_PROMPT_TOKENIZER.build_catalog(
|
||||
clip=clip,
|
||||
prompt_text="1girl, pink hair",
|
||||
model_family=RegionalModelFamily.STANDARD_UNET,
|
||||
)
|
||||
|
||||
assert (
|
||||
ATTENTION_PROMPT_TOKENIZER.resolve_query(
|
||||
clip=clip,
|
||||
query="1girl black dress",
|
||||
catalog=catalog,
|
||||
model_family=RegionalModelFamily.STANDARD_UNET,
|
||||
)
|
||||
== ()
|
||||
)
|
||||
@@ -0,0 +1,200 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test public attention-region v3 contracts and shared execution shapes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.attention_region_capture import (
|
||||
AttentionCaptureProfile,
|
||||
AttentionRegionControls,
|
||||
)
|
||||
from simple_syrup.domain.attention_region_maps import CapturedAttentionMap
|
||||
from simple_syrup.nodes_v3 import get_nodes
|
||||
from simple_syrup.runtime.attention_region_store import ATTENTION_REGION_CAPTURE_STORE
|
||||
from simple_syrup.services.attention_capture_model_service import (
|
||||
EmptyAttentionCaptureSession,
|
||||
)
|
||||
from simple_syrup.services.attention_region_node_service import (
|
||||
ATTENTION_REGION_NODE_SERVICE,
|
||||
)
|
||||
|
||||
|
||||
def test_v3_registry_exposes_four_public_nodes_and_internal_capture_node() -> None:
|
||||
"""Pin public identifiers and keep the injected implementation node dev-only."""
|
||||
|
||||
schemas = {
|
||||
cast(Any, node).define_schema().node_id: cast(Any, node).define_schema()
|
||||
for node in get_nodes()
|
||||
}
|
||||
|
||||
assert {
|
||||
"SimpleSyrup.ConceptAttentionSEGS",
|
||||
"SimpleSyrup.AllPromptAttentionSEGS",
|
||||
"SimpleSyrup.AttentionRegionMask",
|
||||
"SimpleSyrup.AttentionMaskedConditioning",
|
||||
"SimpleSyrup.AttentionCaptureModel",
|
||||
} <= schemas.keys()
|
||||
assert schemas["SimpleSyrup.AttentionCaptureModel"].is_dev_only is True
|
||||
for node_id in (
|
||||
"SimpleSyrup.ConceptAttentionSEGS",
|
||||
"SimpleSyrup.AllPromptAttentionSEGS",
|
||||
"SimpleSyrup.AttentionRegionMask",
|
||||
"SimpleSyrup.AttentionMaskedConditioning",
|
||||
):
|
||||
schema = schemas[node_id]
|
||||
assert schema.description
|
||||
assert all(item.tooltip for item in (*schema.inputs, *schema.outputs))
|
||||
|
||||
concept_inputs = tuple(
|
||||
item.id for item in schemas["SimpleSyrup.ConceptAttentionSEGS"].inputs
|
||||
)
|
||||
assert "concepts" in concept_inputs
|
||||
assert "queries" not in concept_inputs
|
||||
assert {
|
||||
"sampler_stage",
|
||||
"keep_only",
|
||||
"keep_by",
|
||||
"combine_segs",
|
||||
"matte_solidity",
|
||||
"edge_feather",
|
||||
} <= set(concept_inputs)
|
||||
concept_defaults = {
|
||||
item.id: item.default
|
||||
for item in schemas["SimpleSyrup.ConceptAttentionSEGS"].inputs
|
||||
if hasattr(item, "default")
|
||||
}
|
||||
assert concept_defaults["minimum_strength"] == 0.15
|
||||
assert concept_defaults["split_sensitivity"] == 0.35
|
||||
assert concept_defaults["minimum_region_size"] == 512
|
||||
assert concept_defaults["keep_only"] == 1
|
||||
assert concept_defaults["matte_solidity"] == 0.75
|
||||
|
||||
|
||||
def test_image_service_preserves_pixels_and_returns_batch_segs_and_soft_masks() -> None:
|
||||
"""Keep BHWC pixels bit-identical while rendering one SEGS payload per image."""
|
||||
|
||||
image = torch.rand(2, 4, 4, 3)
|
||||
session = _Session(
|
||||
(
|
||||
CapturedAttentionMap(
|
||||
"hair",
|
||||
torch.tensor([0.0, 1.0, 0.0, 0.0]),
|
||||
0.5,
|
||||
"layer",
|
||||
0,
|
||||
),
|
||||
CapturedAttentionMap(
|
||||
"hair",
|
||||
torch.tensor([0.0, 0.0, 0.0, 1.0]),
|
||||
0.5,
|
||||
"layer",
|
||||
1,
|
||||
),
|
||||
)
|
||||
)
|
||||
ATTENTION_REGION_CAPTURE_STORE.publish(("search",), session)
|
||||
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_image(
|
||||
request_node_id="search",
|
||||
image=image,
|
||||
controls=_controls(),
|
||||
)
|
||||
|
||||
assert result.image.data_ptr() == image.data_ptr()
|
||||
assert torch.equal(result.image, image)
|
||||
assert len(result.segs) == 2
|
||||
assert result.mask.shape == (2, 4, 4)
|
||||
assert result.mask[0].argmax().item() != result.mask[1].argmax().item()
|
||||
|
||||
|
||||
def test_latent_service_supports_anima_singleton_frame_and_empty_no_op() -> None:
|
||||
"""Return exact BCTHW provenance and correctly sized empty masks."""
|
||||
|
||||
latent: dict[str, object] = {"samples": torch.zeros(2, 16, 1, 8, 6)}
|
||||
ATTENTION_REGION_CAPTURE_STORE.publish(
|
||||
("mask",),
|
||||
EmptyAttentionCaptureSession(("mask",), "unsupported model"),
|
||||
)
|
||||
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_latent(
|
||||
request_node_id="mask",
|
||||
latent=latent,
|
||||
controls=_controls(),
|
||||
)
|
||||
|
||||
assert result.latent is latent
|
||||
assert result.mask.shape == (2, 8, 6)
|
||||
assert result.mask.count_nonzero().item() == 0
|
||||
|
||||
|
||||
def test_mask_conditioning_uses_comfy_generic_mask_metadata(monkeypatch: Any) -> None:
|
||||
"""Attach a later-pass mask through Comfy's conditioning helper contract."""
|
||||
|
||||
calls: list[tuple[object, dict[str, object]]] = []
|
||||
|
||||
def setter(conditioning: object, values: dict[str, object]) -> object:
|
||||
"""Record exact conditioning metadata and return a sentinel."""
|
||||
|
||||
calls.append((conditioning, values))
|
||||
return "masked"
|
||||
|
||||
fake_module = type(
|
||||
"NodeHelpers",
|
||||
(),
|
||||
{"conditioning_set_values": staticmethod(setter)},
|
||||
)()
|
||||
import simple_syrup.services.attention_region_node_service as service_module
|
||||
|
||||
monkeypatch.setattr(service_module, "import_module", lambda _name: fake_module)
|
||||
mask = torch.ones(1, 4, 4)
|
||||
|
||||
result = ATTENTION_REGION_NODE_SERVICE.mask_conditioning("cond", mask, 1.25)
|
||||
|
||||
assert result == "masked"
|
||||
assert calls == [
|
||||
(
|
||||
"cond",
|
||||
{
|
||||
"mask": mask,
|
||||
"set_area_to_bounds": False,
|
||||
"mask_strength": 1.25,
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class _Session:
|
||||
"""Expose deterministic maps through the capture-store protocol."""
|
||||
|
||||
status_message = "captured"
|
||||
|
||||
def __init__(self, maps: tuple[CapturedAttentionMap, ...]) -> None:
|
||||
"""Retain immutable maps for one request."""
|
||||
|
||||
self._maps = maps
|
||||
|
||||
def maps_for(self, request_node_id: str) -> tuple[CapturedAttentionMap, ...]:
|
||||
"""Return maps for the test request."""
|
||||
|
||||
assert request_node_id == "search"
|
||||
return self._maps
|
||||
|
||||
|
||||
def _controls() -> AttentionRegionControls:
|
||||
"""Return permissive deterministic rendering controls."""
|
||||
|
||||
return AttentionRegionControls(
|
||||
0.0,
|
||||
1.0,
|
||||
0.1,
|
||||
0.0,
|
||||
0.0,
|
||||
1,
|
||||
AttentionCaptureProfile.EXHAUSTIVE,
|
||||
)
|
||||
@@ -0,0 +1,53 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Test graph-visible attention mask coordinate projection."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.attention_spatial_transform import (
|
||||
AttentionSpatialTransform,
|
||||
AttentionSpatialTransformKind,
|
||||
)
|
||||
from simple_syrup.services.attention_spatial_projection import (
|
||||
ATTENTION_SPATIAL_PROJECTION_SERVICE,
|
||||
)
|
||||
|
||||
|
||||
def test_cover_crop_projects_attention_with_requested_anchor() -> None:
|
||||
"""Crop a covered map using the same anchor geometry as image resizing."""
|
||||
|
||||
source = torch.tensor([[1.0, 0.75, 0.25, 0.0], [1.0, 0.75, 0.25, 0.0]])
|
||||
transform = AttentionSpatialTransform(
|
||||
AttentionSpatialTransformKind.COVER_CROP,
|
||||
width=2,
|
||||
height=2,
|
||||
anchor="center",
|
||||
)
|
||||
|
||||
result = ATTENTION_SPATIAL_PROJECTION_SERVICE.project(source, (transform,))
|
||||
|
||||
assert result.shape == (2, 2)
|
||||
assert torch.allclose(result, source[:, 1:3])
|
||||
|
||||
|
||||
def test_fit_pad_projects_attention_without_filling_the_padded_canvas() -> None:
|
||||
"""Keep padding outside the transformed source mask at zero alpha."""
|
||||
|
||||
source = torch.ones(2, 4)
|
||||
transform = AttentionSpatialTransform(
|
||||
AttentionSpatialTransformKind.FIT_PAD,
|
||||
width=4,
|
||||
height=4,
|
||||
anchor="center",
|
||||
)
|
||||
|
||||
result = ATTENTION_SPATIAL_PROJECTION_SERVICE.project(source, (transform,))
|
||||
|
||||
assert result.shape == (4, 4)
|
||||
assert result[0].count_nonzero().item() == 0
|
||||
assert result[-1].count_nonzero().item() == 0
|
||||
assert result[1:3].eq(1.0).all().item()
|
||||
@@ -0,0 +1,40 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Characterize shared connected-component extraction."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.masking.mask_components import connected_mask_components
|
||||
|
||||
|
||||
def test_connected_components_are_eight_connected_and_spatially_ordered() -> None:
|
||||
"""Join diagonal pixels and return components from top-left to bottom-right."""
|
||||
|
||||
active = torch.tensor(
|
||||
[
|
||||
[True, False, False, False],
|
||||
[False, True, False, False],
|
||||
[False, False, False, True],
|
||||
]
|
||||
)
|
||||
|
||||
components = connected_mask_components(active)
|
||||
|
||||
assert tuple(component.bbox for component in components) == (
|
||||
(0, 0, 2, 2),
|
||||
(3, 2, 4, 3),
|
||||
)
|
||||
assert tuple(int(component.mask.count_nonzero()) for component in components) == (
|
||||
2,
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def test_connected_components_return_empty_for_an_empty_mask() -> None:
|
||||
"""Return no components without inventing geometry."""
|
||||
|
||||
assert connected_mask_components(torch.zeros(4, 5, dtype=torch.bool)) == ()
|
||||
@@ -17,6 +17,10 @@ from typing import Any, Protocol, cast
|
||||
import pytest
|
||||
|
||||
BASE_NODE_IDS = [
|
||||
"SimpleSyrup.AllPromptAttentionSEGS",
|
||||
"SimpleSyrup.AttentionCaptureModel",
|
||||
"SimpleSyrup.AttentionMaskedConditioning",
|
||||
"SimpleSyrup.AttentionRegionMask",
|
||||
"SimpleSyrup.BatchRegionConditioning",
|
||||
"SimpleSyrup.BatchSEGS",
|
||||
"SimpleSyrup.ConditioningBatchAppend",
|
||||
@@ -51,6 +55,7 @@ BASE_NODE_IDS = [
|
||||
"SimpleSyrup.SAMModelLoader",
|
||||
"SimpleSyrup.SEGSFromSAMOutput",
|
||||
"SimpleSyrup.ScaleFactor",
|
||||
"SimpleSyrup.ConceptAttentionSEGS",
|
||||
"SimpleSyrup.Seed",
|
||||
"SimpleSyrup.SimpleLoadAnima",
|
||||
"SimpleSyrup.SimplePreviewSEGS",
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Run focused real-model acceptance for cohesive attention silhouettes."""
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Compose labeled visual evidence for attention-cohesion acceptance runs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
|
||||
|
||||
def compose_sheet(
|
||||
*,
|
||||
title: str,
|
||||
concept: str,
|
||||
image_path: Path,
|
||||
raw_path: Path,
|
||||
default_path: Path,
|
||||
solid_path: Path,
|
||||
destination: Path,
|
||||
) -> None:
|
||||
"""Write a four-panel proof sheet with coverage labels."""
|
||||
|
||||
panels = (
|
||||
("GENERATED IMAGE", image_path, False),
|
||||
("RAW ATTENTION ALPHA", raw_path, True),
|
||||
("DEFAULT COHESIVE", default_path, True),
|
||||
("FULL MATTE SOLIDITY", solid_path, True),
|
||||
)
|
||||
panel_width = 384
|
||||
panel_height = 512
|
||||
header_height = 112
|
||||
sheet = Image.new(
|
||||
"RGB", (panel_width * len(panels), header_height + panel_height), "#101216"
|
||||
)
|
||||
draw = ImageDraw.Draw(sheet)
|
||||
title_font = _font(30)
|
||||
label_font = _font(18)
|
||||
detail_font = _font(15)
|
||||
draw.text((18, 12), title, fill="white", font=title_font)
|
||||
draw.text(
|
||||
(18, 54),
|
||||
f'Concept: "{concept}" | Same seed and capture profile across masks',
|
||||
fill="#b9c1cc",
|
||||
font=detail_font,
|
||||
)
|
||||
for index, (label, path, is_mask) in enumerate(panels):
|
||||
panel = (
|
||||
Image.open(path)
|
||||
.convert("RGB")
|
||||
.resize((panel_width, panel_height), Image.Resampling.LANCZOS)
|
||||
)
|
||||
left = index * panel_width
|
||||
sheet.paste(panel, (left, header_height))
|
||||
draw.rectangle((left, 80, left + panel_width, 112), fill="#1c2028")
|
||||
detail = f" support {_coverage(path):.1f}%" if is_mask else ""
|
||||
draw.text((left + 10, 86), label + detail, fill="white", font=label_font)
|
||||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||||
sheet.save(destination)
|
||||
|
||||
|
||||
def _coverage(path: Path) -> float:
|
||||
"""Return the percentage of pixels with visible mask support."""
|
||||
|
||||
values = np.asarray(Image.open(path).convert("L"), dtype=np.uint8)
|
||||
return float((values >= 8).mean() * 100.0)
|
||||
|
||||
|
||||
def _font(size: int) -> ImageFont.FreeTypeFont | ImageFont.ImageFont:
|
||||
"""Load a stable Windows UI font with a portable fallback."""
|
||||
|
||||
path = Path("C:/Windows/Fonts/segoeui.ttf")
|
||||
return ImageFont.truetype(path, size) if path.exists() else ImageFont.load_default()
|
||||
@@ -0,0 +1,54 @@
|
||||
"""Submit ComfyUI acceptance graphs and resolve their saved outputs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .workflow import Graph
|
||||
|
||||
|
||||
def execute(server: str, graph: Graph, output_root: Path) -> dict[str, Path]:
|
||||
"""Execute one graph and return saved image paths by output node id."""
|
||||
|
||||
payload = json.dumps({"prompt": graph}).encode("utf-8")
|
||||
request = urllib.request.Request(
|
||||
f"{server}/prompt",
|
||||
data=payload,
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
with urllib.request.urlopen(request, timeout=30) as response:
|
||||
prompt_id = str(json.loads(response.read())["prompt_id"])
|
||||
deadline = time.monotonic() + 300
|
||||
while time.monotonic() < deadline:
|
||||
with urllib.request.urlopen(
|
||||
f"{server}/history/{prompt_id}", timeout=30
|
||||
) as response:
|
||||
history: dict[str, Any] = json.loads(response.read())
|
||||
if prompt_id in history:
|
||||
record = history[prompt_id]
|
||||
status = record.get("status", {})
|
||||
if status.get("status_str") == "error":
|
||||
raise RuntimeError(json.dumps(status, indent=2))
|
||||
outputs = _saved_outputs(record, output_root)
|
||||
if outputs:
|
||||
return outputs
|
||||
time.sleep(0.5)
|
||||
raise TimeoutError(f"ComfyUI prompt {prompt_id} did not finish within 300 seconds.")
|
||||
|
||||
|
||||
def _saved_outputs(record: dict[str, Any], output_root: Path) -> dict[str, Path]:
|
||||
"""Resolve the first saved image for each completed output node."""
|
||||
|
||||
outputs: dict[str, Path] = {}
|
||||
for node_id, output in record.get("outputs", {}).items():
|
||||
images = output.get("images", [])
|
||||
if images:
|
||||
image = images[0]
|
||||
outputs[str(node_id)] = (
|
||||
output_root / image.get("subfolder", "") / image["filename"]
|
||||
)
|
||||
return outputs
|
||||
@@ -0,0 +1,92 @@
|
||||
"""Run SDXL and Anima attention-cohesion visual acceptance cases."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from .artifacts import compose_sheet
|
||||
from .client import execute
|
||||
from .workflow import build_workflow
|
||||
|
||||
OUTPUT_ROOT = Path("E:/ComfyUI/output")
|
||||
PROOF_ROOT = OUTPUT_ROOT / "simple_syrup_attention_cohesion_proof"
|
||||
SERVER = "http://127.0.0.1:8207"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AcceptanceCase:
|
||||
"""Describe one real-model cohesion acceptance case."""
|
||||
|
||||
source: Path
|
||||
concept: str
|
||||
strength: float
|
||||
consensus: float
|
||||
split: float
|
||||
title: str
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Execute both models and record labeled artifact paths."""
|
||||
|
||||
cases = _cases()
|
||||
case_manifest: dict[str, object] = {}
|
||||
for key, case in cases.items():
|
||||
prefix = f"simple_syrup_attention_cohesion_proof/{key}"
|
||||
graph = build_workflow(
|
||||
case.source,
|
||||
output_prefix=prefix,
|
||||
concept=case.concept,
|
||||
strength=case.strength,
|
||||
consensus=case.consensus,
|
||||
split=case.split,
|
||||
)
|
||||
outputs = execute(SERVER, graph, OUTPUT_ROOT)
|
||||
sheet = PROOF_ROOT / f"{key}_cohesive_silhouette_proof.png"
|
||||
compose_sheet(
|
||||
title=case.title,
|
||||
concept=case.concept,
|
||||
image_path=outputs["900"],
|
||||
raw_path=outputs["101"],
|
||||
default_path=outputs["1101"],
|
||||
solid_path=outputs["1131"],
|
||||
destination=sheet,
|
||||
)
|
||||
case_manifest[key] = {
|
||||
"concept": case.concept,
|
||||
"outputs": {node_id: str(path) for node_id, path in outputs.items()},
|
||||
"proof_sheet": str(sheet),
|
||||
}
|
||||
manifest = {"server": SERVER, "cases": case_manifest}
|
||||
destination = PROOF_ROOT / "cohesion_manifest.json"
|
||||
destination.write_text(json.dumps(manifest, indent=2), encoding="utf-8")
|
||||
print(destination)
|
||||
|
||||
|
||||
def _cases() -> dict[str, AcceptanceCase]:
|
||||
"""Return the requested fixed-model, fixed-prompt acceptance matrix."""
|
||||
|
||||
prior_root = OUTPUT_ROOT / "simple_syrup_attention_acceptance_v2"
|
||||
return {
|
||||
"sdxl": AcceptanceCase(
|
||||
prior_root / "sdxl/final_many_image_00002_.png",
|
||||
"pink hair",
|
||||
0.12,
|
||||
0.25,
|
||||
0.30,
|
||||
"SDXL / Illustrious Amanatsu - Cohesive Attention Silhouette",
|
||||
),
|
||||
"anima": AcceptanceCase(
|
||||
prior_root / "anima/final_many_image_00002_.png",
|
||||
"holding cat",
|
||||
0.14,
|
||||
0.30,
|
||||
0.35,
|
||||
"Hassaku Anima - Cohesive Attention Silhouette",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,131 @@
|
||||
"""Build focused attention-cohesion acceptance workflows from saved prompts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from PIL import Image
|
||||
|
||||
Graph = dict[str, dict[str, Any]]
|
||||
|
||||
|
||||
def build_workflow(
|
||||
source_png: Path,
|
||||
*,
|
||||
output_prefix: str,
|
||||
concept: str,
|
||||
strength: float,
|
||||
consensus: float,
|
||||
split: float,
|
||||
) -> Graph:
|
||||
"""Return a minimal same-seed graph with raw, default, and solid masks."""
|
||||
|
||||
prompt = json.loads(Image.open(source_png).info["prompt"])
|
||||
graph: Graph = deepcopy(prompt)
|
||||
capture_id = next(
|
||||
node_id
|
||||
for node_id, node in graph.items()
|
||||
if node["class_type"] == "SimpleSyrup.AttentionCaptureModel"
|
||||
)
|
||||
sampler_id = next(
|
||||
node_id for node_id, node in graph.items() if node["class_type"] == "KSampler"
|
||||
)
|
||||
graph[sampler_id]["inputs"]["model"] = graph[capture_id]["inputs"]["model"]
|
||||
|
||||
retained = {
|
||||
"1",
|
||||
"2",
|
||||
"3",
|
||||
"4",
|
||||
"5",
|
||||
"6",
|
||||
"900",
|
||||
"10",
|
||||
"100",
|
||||
"101",
|
||||
"11",
|
||||
"1100",
|
||||
"1101",
|
||||
"14",
|
||||
"1130",
|
||||
"1131",
|
||||
}
|
||||
graph = {node_id: node for node_id, node in graph.items() if node_id in retained}
|
||||
graph["900"]["inputs"]["filename_prefix"] = f"{output_prefix}/image"
|
||||
_configure_mask(
|
||||
graph,
|
||||
request_id="10",
|
||||
save_id="101",
|
||||
concept=concept,
|
||||
strength=strength,
|
||||
consensus=consensus,
|
||||
split=split,
|
||||
minimum_size=1,
|
||||
keep_only=0,
|
||||
solidity=0.0,
|
||||
prefix=f"{output_prefix}/raw_alpha",
|
||||
)
|
||||
_configure_mask(
|
||||
graph,
|
||||
request_id="11",
|
||||
save_id="1101",
|
||||
concept=concept,
|
||||
strength=0.15,
|
||||
consensus=0.25,
|
||||
split=0.35,
|
||||
minimum_size=512,
|
||||
keep_only=1,
|
||||
solidity=0.75,
|
||||
prefix=f"{output_prefix}/default_cohesive",
|
||||
)
|
||||
_configure_mask(
|
||||
graph,
|
||||
request_id="14",
|
||||
save_id="1131",
|
||||
concept=concept,
|
||||
strength=0.15,
|
||||
consensus=0.25,
|
||||
split=0.35,
|
||||
minimum_size=512,
|
||||
keep_only=1,
|
||||
solidity=1.0,
|
||||
prefix=f"{output_prefix}/full_solid",
|
||||
)
|
||||
return graph
|
||||
|
||||
|
||||
def _configure_mask(
|
||||
graph: Graph,
|
||||
*,
|
||||
request_id: str,
|
||||
save_id: str,
|
||||
concept: str,
|
||||
strength: float,
|
||||
consensus: float,
|
||||
split: float,
|
||||
minimum_size: int,
|
||||
keep_only: int,
|
||||
solidity: float,
|
||||
prefix: str,
|
||||
) -> None:
|
||||
"""Configure one focused concept-mask output."""
|
||||
|
||||
inputs = graph[request_id]["inputs"]
|
||||
inputs.update(
|
||||
{
|
||||
"concepts": concept,
|
||||
"minimum_strength": strength,
|
||||
"minimum_consensus": consensus,
|
||||
"split_sensitivity": split,
|
||||
"minimum_region_size": minimum_size,
|
||||
"keep_only": keep_only,
|
||||
"combine_segs": False,
|
||||
"matte_solidity": solidity,
|
||||
"edge_feather": 8,
|
||||
"capture_profile": "fast",
|
||||
}
|
||||
)
|
||||
graph[save_id]["inputs"]["filename_prefix"] = prefix
|
||||
Reference in New Issue
Block a user