From 02710d254489ced0492da2cf1c342c0c49019388 Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Wed, 26 Aug 2026 18:43:33 -0400 Subject: [PATCH] feat(attention): add sampler-derived concept regions --- __init__.py | 4 + simple_syrup/domain/attention_concepts.py | 15 + simple_syrup/domain/attention_geometry.py | 30 ++ .../domain/attention_region_capture.py | 157 ++++++ simple_syrup/domain/attention_region_maps.py | 164 ++++++ .../domain/attention_sampler_lineage.py | 78 +++ .../domain/attention_spatial_transform.py | 70 +++ simple_syrup/masking/mask_components.py | 83 +-- simple_syrup/nodes_v3/__init__.py | 10 + .../nodes_v3/all_prompt_attention_segs.py | 119 ++++ .../nodes_v3/attention_capture_model.py | 89 +++ .../nodes_v3/attention_masked_conditioning.py | 151 ++++++ .../nodes_v3/attention_region_inputs.py | 177 ++++++ .../nodes_v3/attention_region_mask.py | 124 +++++ .../nodes_v3/concept_attention_segs.py | 129 +++++ .../runtime/attention_region_affinity.py | 262 +++++++++ .../runtime/attention_region_capture.py | 403 ++++++++++++++ .../attention_region_capture_backend.py | 116 ++++ .../runtime/attention_region_graph.py | 306 +++++++++++ .../attention_region_open_vocabulary.py | 119 ++++ .../runtime/attention_region_plan_codec.py | 208 +++++++ .../attention_region_prompt_handler.py | 73 +++ .../runtime/attention_region_status.py | 38 ++ .../runtime/attention_region_store.py | 76 +++ .../runtime/attention_region_tokens.py | 254 +++++++++ .../runtime/attention_sampler_lineage.py | 510 ++++++++++++++++++ .../attention_capture_model_service.py | 167 ++++++ .../services/attention_region_components.py | 99 ++++ .../services/attention_region_evidence.py | 114 ++++ .../attention_region_instance_splitting.py | 66 +++ .../services/attention_region_matte.py | 86 +++ .../services/attention_region_node_service.py | 194 +++++++ .../services/attention_region_rendering.py | 140 +++++ .../services/attention_spatial_projection.py | 92 ++++ tests/test_attention_region_capture.py | 440 +++++++++++++++ tests/test_attention_region_graph.py | 368 +++++++++++++ tests/test_attention_region_prompt_handler.py | 40 ++ tests/test_attention_region_rendering.py | 379 +++++++++++++ tests/test_attention_region_store.py | 53 ++ tests/test_attention_region_tokens.py | 106 ++++ tests/test_attention_region_v3_nodes.py | 200 +++++++ tests/test_attention_spatial_projection.py | 53 ++ tests/test_mask_components.py | 40 ++ tests/test_registration.py | 5 + .../__init__.py | 1 + .../artifacts.py | 72 +++ .../client.py | 54 ++ .../run.py | 92 ++++ .../workflow.py | 131 +++++ 49 files changed, 6700 insertions(+), 57 deletions(-) create mode 100644 simple_syrup/domain/attention_concepts.py create mode 100644 simple_syrup/domain/attention_geometry.py create mode 100644 simple_syrup/domain/attention_region_capture.py create mode 100644 simple_syrup/domain/attention_region_maps.py create mode 100644 simple_syrup/domain/attention_sampler_lineage.py create mode 100644 simple_syrup/domain/attention_spatial_transform.py create mode 100644 simple_syrup/nodes_v3/all_prompt_attention_segs.py create mode 100644 simple_syrup/nodes_v3/attention_capture_model.py create mode 100644 simple_syrup/nodes_v3/attention_masked_conditioning.py create mode 100644 simple_syrup/nodes_v3/attention_region_inputs.py create mode 100644 simple_syrup/nodes_v3/attention_region_mask.py create mode 100644 simple_syrup/nodes_v3/concept_attention_segs.py create mode 100644 simple_syrup/runtime/attention_region_affinity.py create mode 100644 simple_syrup/runtime/attention_region_capture.py create mode 100644 simple_syrup/runtime/attention_region_capture_backend.py create mode 100644 simple_syrup/runtime/attention_region_graph.py create mode 100644 simple_syrup/runtime/attention_region_open_vocabulary.py create mode 100644 simple_syrup/runtime/attention_region_plan_codec.py create mode 100644 simple_syrup/runtime/attention_region_prompt_handler.py create mode 100644 simple_syrup/runtime/attention_region_status.py create mode 100644 simple_syrup/runtime/attention_region_store.py create mode 100644 simple_syrup/runtime/attention_region_tokens.py create mode 100644 simple_syrup/runtime/attention_sampler_lineage.py create mode 100644 simple_syrup/services/attention_capture_model_service.py create mode 100644 simple_syrup/services/attention_region_components.py create mode 100644 simple_syrup/services/attention_region_evidence.py create mode 100644 simple_syrup/services/attention_region_instance_splitting.py create mode 100644 simple_syrup/services/attention_region_matte.py create mode 100644 simple_syrup/services/attention_region_node_service.py create mode 100644 simple_syrup/services/attention_region_rendering.py create mode 100644 simple_syrup/services/attention_spatial_projection.py create mode 100644 tests/test_attention_region_capture.py create mode 100644 tests/test_attention_region_graph.py create mode 100644 tests/test_attention_region_prompt_handler.py create mode 100644 tests/test_attention_region_rendering.py create mode 100644 tests/test_attention_region_store.py create mode 100644 tests/test_attention_region_tokens.py create mode 100644 tests/test_attention_region_v3_nodes.py create mode 100644 tests/test_attention_spatial_projection.py create mode 100644 tests/test_mask_components.py create mode 100644 tools/attention_region_cohesion_acceptance/__init__.py create mode 100644 tools/attention_region_cohesion_acceptance/artifacts.py create mode 100644 tools/attention_region_cohesion_acceptance/client.py create mode 100644 tools/attention_region_cohesion_acceptance/run.py create mode 100644 tools/attention_region_cohesion_acceptance/workflow.py diff --git a/__init__.py b/__init__.py index 673cb00..69a946e 100644 --- a/__init__.py +++ b/__init__.py @@ -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", diff --git a/simple_syrup/domain/attention_concepts.py b/simple_syrup/domain/attention_concepts.py new file mode 100644 index 0000000..09b2d91 --- /dev/null +++ b/simple_syrup/domain/attention_concepts.py @@ -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()) diff --git a/simple_syrup/domain/attention_geometry.py b/simple_syrup/domain/attention_geometry.py new file mode 100644 index 0000000..a0c5aaa --- /dev/null +++ b/simple_syrup/domain/attention_geometry.py @@ -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 diff --git a/simple_syrup/domain/attention_region_capture.py b/simple_syrup/domain/attention_region_capture.py new file mode 100644 index 0000000..aeeaeab --- /dev/null +++ b/simple_syrup/domain/attention_region_capture.py @@ -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}" diff --git a/simple_syrup/domain/attention_region_maps.py b/simple_syrup/domain/attention_region_maps.py new file mode 100644 index 0000000..4299617 --- /dev/null +++ b/simple_syrup/domain/attention_region_maps.py @@ -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()) diff --git a/simple_syrup/domain/attention_sampler_lineage.py b/simple_syrup/domain/attention_sampler_lineage.py new file mode 100644 index 0000000..9901a40 --- /dev/null +++ b/simple_syrup/domain/attention_sampler_lineage.py @@ -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, + ) diff --git a/simple_syrup/domain/attention_spatial_transform.py b/simple_syrup/domain/attention_spatial_transform.py new file mode 100644 index 0000000..230aba3 --- /dev/null +++ b/simple_syrup/domain/attention_spatial_transform.py @@ -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.") diff --git a/simple_syrup/masking/mask_components.py b/simple_syrup/masking/mask_components.py index 0109e4c..f10a38b 100644 --- a/simple_syrup/masking/mask_components.py +++ b/simple_syrup/masking/mask_components.py @@ -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) diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index 6eda12a..2c2d53c 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -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, diff --git a/simple_syrup/nodes_v3/all_prompt_attention_segs.py b/simple_syrup/nodes_v3/all_prompt_attention_segs.py new file mode 100644 index 0000000..e07207c --- /dev/null +++ b/simple_syrup/nodes_v3/all_prompt_attention_segs.py @@ -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 diff --git a/simple_syrup/nodes_v3/attention_capture_model.py b/simple_syrup/nodes_v3/attention_capture_model.py new file mode 100644 index 0000000..1b597b0 --- /dev/null +++ b/simple_syrup/nodes_v3/attention_capture_model.py @@ -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, + ), + ) diff --git a/simple_syrup/nodes_v3/attention_masked_conditioning.py b/simple_syrup/nodes_v3/attention_masked_conditioning.py new file mode 100644 index 0000000..03576b5 --- /dev/null +++ b/simple_syrup/nodes_v3/attention_masked_conditioning.py @@ -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 diff --git a/simple_syrup/nodes_v3/attention_region_inputs.py b/simple_syrup/nodes_v3/attention_region_inputs.py new file mode 100644 index 0000000..8af8985 --- /dev/null +++ b/simple_syrup/nodes_v3/attention_region_inputs.py @@ -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, + ) diff --git a/simple_syrup/nodes_v3/attention_region_mask.py b/simple_syrup/nodes_v3/attention_region_mask.py new file mode 100644 index 0000000..36dce05 --- /dev/null +++ b/simple_syrup/nodes_v3/attention_region_mask.py @@ -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 diff --git a/simple_syrup/nodes_v3/concept_attention_segs.py b/simple_syrup/nodes_v3/concept_attention_segs.py new file mode 100644 index 0000000..e631c0f --- /dev/null +++ b/simple_syrup/nodes_v3/concept_attention_segs.py @@ -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 diff --git a/simple_syrup/runtime/attention_region_affinity.py b/simple_syrup/runtime/attention_region_affinity.py new file mode 100644 index 0000000..d514ccc --- /dev/null +++ b/simple_syrup/runtime/attention_region_affinity.py @@ -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() diff --git a/simple_syrup/runtime/attention_region_capture.py b/simple_syrup/runtime/attention_region_capture.py new file mode 100644 index 0000000..ea4aa0e --- /dev/null +++ b/simple_syrup/runtime/attention_region_capture.py @@ -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 diff --git a/simple_syrup/runtime/attention_region_capture_backend.py b/simple_syrup/runtime/attention_region_capture_backend.py new file mode 100644 index 0000000..008d8e5 --- /dev/null +++ b/simple_syrup/runtime/attention_region_capture_backend.py @@ -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() diff --git a/simple_syrup/runtime/attention_region_graph.py b/simple_syrup/runtime/attention_region_graph.py new file mode 100644 index 0000000..1f6260c --- /dev/null +++ b/simple_syrup/runtime/attention_region_graph.py @@ -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() diff --git a/simple_syrup/runtime/attention_region_open_vocabulary.py b/simple_syrup/runtime/attention_region_open_vocabulary.py new file mode 100644 index 0000000..c29ad5f --- /dev/null +++ b/simple_syrup/runtime/attention_region_open_vocabulary.py @@ -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) diff --git a/simple_syrup/runtime/attention_region_plan_codec.py b/simple_syrup/runtime/attention_region_plan_codec.py new file mode 100644 index 0000000..c646562 --- /dev/null +++ b/simple_syrup/runtime/attention_region_plan_codec.py @@ -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() diff --git a/simple_syrup/runtime/attention_region_prompt_handler.py b/simple_syrup/runtime/attention_region_prompt_handler.py new file mode 100644 index 0000000..7317fe9 --- /dev/null +++ b/simple_syrup/runtime/attention_region_prompt_handler.py @@ -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) diff --git a/simple_syrup/runtime/attention_region_status.py b/simple_syrup/runtime/attention_region_status.py new file mode 100644 index 0000000..dfc3dd3 --- /dev/null +++ b/simple_syrup/runtime/attention_region_status.py @@ -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() diff --git a/simple_syrup/runtime/attention_region_store.py b/simple_syrup/runtime/attention_region_store.py new file mode 100644 index 0000000..8409243 --- /dev/null +++ b/simple_syrup/runtime/attention_region_store.py @@ -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() diff --git a/simple_syrup/runtime/attention_region_tokens.py b/simple_syrup/runtime/attention_region_tokens.py new file mode 100644 index 0000000..ad94859 --- /dev/null +++ b/simple_syrup/runtime/attention_region_tokens.py @@ -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() diff --git a/simple_syrup/runtime/attention_sampler_lineage.py b/simple_syrup/runtime/attention_sampler_lineage.py new file mode 100644 index 0000000..8eb9c3d --- /dev/null +++ b/simple_syrup/runtime/attention_sampler_lineage.py @@ -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) diff --git a/simple_syrup/services/attention_capture_model_service.py b/simple_syrup/services/attention_capture_model_service.py new file mode 100644 index 0000000..a105623 --- /dev/null +++ b/simple_syrup/services/attention_capture_model_service.py @@ -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()) diff --git a/simple_syrup/services/attention_region_components.py b/simple_syrup/services/attention_region_components.py new file mode 100644 index 0000000..9a1c8aa --- /dev/null +++ b/simple_syrup/services/attention_region_components.py @@ -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() diff --git a/simple_syrup/services/attention_region_evidence.py b/simple_syrup/services/attention_region_evidence.py new file mode 100644 index 0000000..5431479 --- /dev/null +++ b/simple_syrup/services/attention_region_evidence.py @@ -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() diff --git a/simple_syrup/services/attention_region_instance_splitting.py b/simple_syrup/services/attention_region_instance_splitting.py new file mode 100644 index 0000000..3a9b7c4 --- /dev/null +++ b/simple_syrup/services/attention_region_instance_splitting.py @@ -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() diff --git a/simple_syrup/services/attention_region_matte.py b/simple_syrup/services/attention_region_matte.py new file mode 100644 index 0000000..b4babe5 --- /dev/null +++ b/simple_syrup/services/attention_region_matte.py @@ -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() diff --git a/simple_syrup/services/attention_region_node_service.py b/simple_syrup/services/attention_region_node_service.py new file mode 100644 index 0000000..42b2a81 --- /dev/null +++ b/simple_syrup/services/attention_region_node_service.py @@ -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() diff --git a/simple_syrup/services/attention_region_rendering.py b/simple_syrup/services/attention_region_rendering.py new file mode 100644 index 0000000..a048c88 --- /dev/null +++ b/simple_syrup/services/attention_region_rendering.py @@ -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() diff --git a/simple_syrup/services/attention_spatial_projection.py b/simple_syrup/services/attention_spatial_projection.py new file mode 100644 index 0000000..7c8f647 --- /dev/null +++ b/simple_syrup/services/attention_spatial_projection.py @@ -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() diff --git a/tests/test_attention_region_capture.py b/tests/test_attention_region_capture.py new file mode 100644 index 0000000..ea77d6b --- /dev/null +++ b/tests/test_attention_region_capture.py @@ -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) diff --git a/tests/test_attention_region_graph.py b/tests/test_attention_region_graph.py new file mode 100644 index 0000000..b4164ad --- /dev/null +++ b/tests/test_attention_region_graph.py @@ -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 diff --git a/tests/test_attention_region_prompt_handler.py b/tests/test_attention_region_prompt_handler.py new file mode 100644 index 0000000..a1cafbb --- /dev/null +++ b/tests/test_attention_region_prompt_handler.py @@ -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 diff --git a/tests/test_attention_region_rendering.py b/tests/test_attention_region_rendering.py new file mode 100644 index 0000000..81b427d --- /dev/null +++ b/tests/test_attention_region_rendering.py @@ -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, + ) diff --git a/tests/test_attention_region_store.py b/tests/test_attention_region_store.py new file mode 100644 index 0000000..85852b9 --- /dev/null +++ b/tests/test_attention_region_store.py @@ -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()) diff --git a/tests/test_attention_region_tokens.py b/tests/test_attention_region_tokens.py new file mode 100644 index 0000000..1c22884 --- /dev/null +++ b/tests/test_attention_region_tokens.py @@ -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, + ) + == () + ) diff --git a/tests/test_attention_region_v3_nodes.py b/tests/test_attention_region_v3_nodes.py new file mode 100644 index 0000000..8ca79f3 --- /dev/null +++ b/tests/test_attention_region_v3_nodes.py @@ -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, + ) diff --git a/tests/test_attention_spatial_projection.py b/tests/test_attention_spatial_projection.py new file mode 100644 index 0000000..65b6e50 --- /dev/null +++ b/tests/test_attention_spatial_projection.py @@ -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() diff --git a/tests/test_mask_components.py b/tests/test_mask_components.py new file mode 100644 index 0000000..867dc81 --- /dev/null +++ b/tests/test_mask_components.py @@ -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)) == () diff --git a/tests/test_registration.py b/tests/test_registration.py index 55f70d8..60d58f7 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -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", diff --git a/tools/attention_region_cohesion_acceptance/__init__.py b/tools/attention_region_cohesion_acceptance/__init__.py new file mode 100644 index 0000000..ded1b1d --- /dev/null +++ b/tools/attention_region_cohesion_acceptance/__init__.py @@ -0,0 +1 @@ +"""Run focused real-model acceptance for cohesive attention silhouettes.""" diff --git a/tools/attention_region_cohesion_acceptance/artifacts.py b/tools/attention_region_cohesion_acceptance/artifacts.py new file mode 100644 index 0000000..b4797fd --- /dev/null +++ b/tools/attention_region_cohesion_acceptance/artifacts.py @@ -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() diff --git a/tools/attention_region_cohesion_acceptance/client.py b/tools/attention_region_cohesion_acceptance/client.py new file mode 100644 index 0000000..f8cb02c --- /dev/null +++ b/tools/attention_region_cohesion_acceptance/client.py @@ -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 diff --git a/tools/attention_region_cohesion_acceptance/run.py b/tools/attention_region_cohesion_acceptance/run.py new file mode 100644 index 0000000..96f4e95 --- /dev/null +++ b/tools/attention_region_cohesion_acceptance/run.py @@ -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() diff --git a/tools/attention_region_cohesion_acceptance/workflow.py b/tools/attention_region_cohesion_acceptance/workflow.py new file mode 100644 index 0000000..6f74087 --- /dev/null +++ b/tools/attention_region_cohesion_acceptance/workflow.py @@ -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