feat(attention): add sampler-derived concept regions

This commit is contained in:
Artificial Sweetener
2026-08-26 18:43:33 -04:00
parent b80d2bc420
commit 02710d2544
49 changed files with 6700 additions and 57 deletions
+4
View File
@@ -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",
+15
View File
@@ -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())
+30
View File
@@ -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.")
+26 -57
View File
@@ -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)
+10
View File
@@ -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()
+440
View File
@@ -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)
+368
View File
@@ -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
+379
View File
@@ -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,
)
+53
View File
@@ -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())
+106
View File
@@ -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,
)
== ()
)
+200
View File
@@ -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()
+40
View File
@@ -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)) == ()
+5
View File
@@ -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