diff --git a/simple_syrup/domain/processed_regional_attention.py b/simple_syrup/domain/processed_regional_attention.py index a335018..574ffed 100644 --- a/simple_syrup/domain/processed_regional_attention.py +++ b/simple_syrup/domain/processed_regional_attention.py @@ -30,6 +30,7 @@ class ProcessedRegionalAttentionEntry: schedule: ConditioningScheduleRange cross_attention: torch.Tensor strength: float + cross_attention_value_multiplier: torch.Tensor | None = None def __post_init__(self) -> None: """Validate entry order, model context, and finite scalar strength.""" @@ -62,6 +63,21 @@ class ProcessedRegionalAttentionEntry: if not math.isfinite(float(self.strength)): raise ValueError("Processed conditioning strength must be finite.") object.__setattr__(self, "strength", float(self.strength)) + multiplier = self.cross_attention_value_multiplier + if multiplier is None: + return + if ( + not isinstance(multiplier, torch.Tensor) + or multiplier.shape != (*self.cross_attention.shape[:2], 1) + or not multiplier.is_floating_point() + or multiplier.device != self.cross_attention.device + or multiplier.dtype != self.cross_attention.dtype + or not bool(torch.isfinite(multiplier).all().item()) + ): + raise ValueError( + "Processed attention value multiplier must be a finite floating " + "BxSx1 tensor aligned with cross_attention." + ) @dataclass(frozen=True, slots=True) diff --git a/simple_syrup/domain/regional_attention_batch.py b/simple_syrup/domain/regional_attention_batch.py index d2a998f..1183d3f 100644 --- a/simple_syrup/domain/regional_attention_batch.py +++ b/simple_syrup/domain/regional_attention_batch.py @@ -48,6 +48,7 @@ class BatchedRegionalAttentionEntry: entry_index: int context: torch.Tensor strengths: tuple[float, ...] + cross_attention_value_multiplier: torch.Tensor | None = None def __post_init__(self) -> None: """Validate entry order, aligned context, and finite sample strengths.""" @@ -70,6 +71,11 @@ class BatchedRegionalAttentionEntry: ) if not math.isfinite(float(strength)): raise ValueError("Regional attention entry strength must be finite.") + _validate_value_multiplier( + self.cross_attention_value_multiplier, + self.context, + name="entry", + ) @dataclass(frozen=True, slots=True) @@ -109,6 +115,7 @@ class BatchedRegionalAttentionContexts: chunks: tuple[RegionalAttentionChunkBatch, ...] base_context: torch.Tensor regions: tuple[BatchedRegionalAttentionRegion, ...] + base_value_multiplier: torch.Tensor | None = None def __post_init__(self) -> None: """Validate complete chunk and tensor alignment.""" @@ -141,6 +148,11 @@ class BatchedRegionalAttentionContexts: expected_batch=expected_start, name="base", ) + _validate_value_multiplier( + self.base_value_multiplier, + self.base_context, + name="base", + ) if not isinstance(self.regions, tuple): raise TypeError("Regional attention regions must be a tuple.") if tuple(region.region_index for region in self.regions) != tuple( @@ -189,3 +201,27 @@ def _validate_aligned_context( raise ValueError( f"Regional attention {name} context must contain finite floating values." ) + + +def _validate_value_multiplier( + multiplier: object, + context: torch.Tensor, + *, + name: str, +) -> None: + """Validate one optional value multiplier against its aligned context.""" + + if multiplier is None: + return + if ( + not isinstance(multiplier, torch.Tensor) + or multiplier.shape != (*context.shape[:2], 1) + or not multiplier.is_floating_point() + or multiplier.device != context.device + or multiplier.dtype != context.dtype + or not bool(torch.isfinite(multiplier).all().item()) + ): + raise ValueError( + f"Regional attention {name} value multiplier must be a finite " + "floating BxSx1 tensor aligned with its context." + ) diff --git a/simple_syrup/runtime/attention_coupling/unet.py b/simple_syrup/runtime/attention_coupling/unet.py index e6f360f..a518fa9 100644 --- a/simple_syrup/runtime/attention_coupling/unet.py +++ b/simple_syrup/runtime/attention_coupling/unet.py @@ -10,6 +10,7 @@ from dataclasses import dataclass from ..model_attention_patch_mutations import ModelAttn2PatchesMutation from ..patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation +from ..ppm_negpip_interop import PpmNegpipInterop from ..regional_lora.standard_unet_native_admission import ( StandardUnetNativeLoraAdmission, ) @@ -43,6 +44,7 @@ class StandardUnetAttentionBackend: model: object, state: StandardUnetAttentionState, admission: StandardUnetNativeLoraAdmission, + negpip: PpmNegpipInterop | None = None, ) -> StandardUnetAttentionModel: """Return a direct MODEL child containing only the paired UNet patches.""" @@ -55,6 +57,8 @@ class StandardUnetAttentionBackend: "Standard UNet admission and processed conditioning must share " "the same regional LoRA plan." ) + if negpip is not None and not isinstance(negpip, PpmNegpipInterop): + raise TypeError("Standard UNet backend NegPiP state has an invalid type.") attention_phase = StandardUnetAttentionPhaseSession() template = ( STANDARD_UNET_VARIANT_TEMPLATE_CACHE.resolve(model, admission) @@ -68,6 +72,7 @@ class StandardUnetAttentionBackend: admission, attention_phase, template, + negpip, ), ) if template is not None @@ -85,6 +90,7 @@ class StandardUnetAttentionBackend: ModelAttn2PatchesMutation( patches.input_patch, patches.output_patch, + (() if negpip is None else (negpip.attention_patch,)), ), ) derived = PATCHER_LIFECYCLE.derive_model( diff --git a/simple_syrup/runtime/comfy_conditioning_processing.py b/simple_syrup/runtime/comfy_conditioning_processing.py index f5822f2..fe4bfd7 100644 --- a/simple_syrup/runtime/comfy_conditioning_processing.py +++ b/simple_syrup/runtime/comfy_conditioning_processing.py @@ -27,6 +27,7 @@ from ..services.attention_coupling_preparation_service import ( AttentionCouplingPreparation, ) from .attention_coupling.context_validation import RegionalContextValidator +from .ppm_negpip_interop import PpmNegpipInterop class ComfyRegionalConditioningProcessor: @@ -40,6 +41,7 @@ class ComfyRegionalConditioningProcessor: noise: torch.Tensor, device: torch.device, context_validator: RegionalContextValidator, + negpip: PpmNegpipInterop | None = None, ) -> ProcessedRegionalAttentionPlan: """Return model-ready positive and negative context banks.""" @@ -57,6 +59,8 @@ class ComfyRegionalConditioningProcessor: raise TypeError("Regional context processing device must be torch.device.") if not isinstance(context_validator, RegionalContextValidator): raise TypeError("Regional context validator has an invalid type.") + if negpip is not None and not isinstance(negpip, PpmNegpipInterop): + raise TypeError("Regional conditioning NegPiP state has an invalid type.") base_model = getattr(model, "model", None) extra_conds = getattr(base_model, "extra_conds", None) if not callable(extra_conds): @@ -75,6 +79,7 @@ class ComfyRegionalConditioningProcessor: noise=noise, device=device, context_validator=context_validator, + negpip=negpip, ) negative = self._process_branch( preparation.plan.negative, @@ -84,6 +89,7 @@ class ComfyRegionalConditioningProcessor: noise=noise, device=device, context_validator=context_validator, + negpip=negpip, ) return ProcessedRegionalAttentionPlan( positive=positive, @@ -102,6 +108,7 @@ class ComfyRegionalConditioningProcessor: noise: torch.Tensor, device: torch.device, context_validator: RegionalContextValidator, + negpip: PpmNegpipInterop | None, ) -> ProcessedRegionalAttentionBranch: """Process one base plus its ordered regional context bank.""" @@ -115,6 +122,7 @@ class ComfyRegionalConditioningProcessor: noise=noise, device=device, context_validator=context_validator, + negpip=negpip, ) regional = tuple( self._process_context( @@ -127,6 +135,7 @@ class ComfyRegionalConditioningProcessor: noise=noise, device=device, context_validator=context_validator, + negpip=negpip, ) for context in branch.regional_contexts ) @@ -144,6 +153,7 @@ class ComfyRegionalConditioningProcessor: noise: torch.Tensor, device: torch.device, context_validator: RegionalContextValidator, + negpip: PpmNegpipInterop | None, ) -> ProcessedRegionalAttentionContext: """Convert and extract one exact post-adapter Anima context tensor.""" @@ -168,6 +178,7 @@ class ComfyRegionalConditioningProcessor: conditioning_index=conditioning_index, prompt_type=prompt_type, context_validator=context_validator, + negpip=negpip, ) for entry_index, encoded_item in enumerate(encoded) ) @@ -185,6 +196,7 @@ class ComfyRegionalConditioningProcessor: conditioning_index: int, prompt_type: str, context_validator: RegionalContextValidator, + negpip: PpmNegpipInterop | None, ) -> ProcessedRegionalAttentionEntry: """Extract one exact post-adapter Anima context and Comfy strength.""" @@ -233,6 +245,11 @@ class ComfyRegionalConditioningProcessor: ), cross_attention=context, strength=float(strength), + cross_attention_value_multiplier=( + None + if negpip is None + else negpip.extract_value_multiplier(model_conds, context) + ), ) @staticmethod diff --git a/simple_syrup/runtime/model_attention_patch_mutations.py b/simple_syrup/runtime/model_attention_patch_mutations.py index 9a074a8..6f875d4 100644 --- a/simple_syrup/runtime/model_attention_patch_mutations.py +++ b/simple_syrup/runtime/model_attention_patch_mutations.py @@ -22,6 +22,7 @@ def _apply_paired_attention_patches( attention_name: str, input_patch: Callable[..., object], output_patch: Callable[..., object], + trailing_input_patches: tuple[Callable[..., object], ...], ) -> None: """Validate and atomically install one paired attention callback surface.""" @@ -29,6 +30,12 @@ def _apply_paired_attention_patches( raise TypeError(f"MODEL {attention_name} input patch must be callable.") if not callable(output_patch): raise TypeError(f"MODEL {attention_name} output patch must be callable.") + if not isinstance(trailing_input_patches, tuple) or any( + not callable(patch) for patch in trailing_input_patches + ): + raise TypeError( + f"MODEL {attention_name} preserved input patches must be callables." + ) input_setter = _require_bound_method( model, f"set_model_{attention_name}_patch", @@ -54,11 +61,21 @@ def _apply_paired_attention_patches( output_name = f"{attention_name}_output_patch" input_exists = _require_callable_patch_list(patches, input_name) output_exists = _require_callable_patch_list(patches, output_name) - if input_exists: + existing_input = patches.get(input_name, []) + if input_exists and ( + not trailing_input_patches or existing_input != list(trailing_input_patches) + ): raise ValueError(f"MODEL {attention_name} input patch is already installed.") + if not input_exists and trailing_input_patches: + raise ValueError( + f"MODEL {attention_name} preserved input patch is not installed." + ) if output_exists: raise ValueError(f"MODEL {attention_name} output patch is already installed.") - input_setter(input_patch) + if trailing_input_patches: + patches[input_name] = [input_patch, *trailing_input_patches] + else: + input_setter(input_patch) output_setter(output_patch) @@ -68,6 +85,7 @@ class ModelAttn2PatchesMutation: input_patch: Callable[..., object] output_patch: Callable[..., object] + trailing_input_patches: tuple[Callable[..., object], ...] = () def apply(self, model: object) -> None: """Validate both attn2 surfaces before either mutation.""" @@ -77,4 +95,5 @@ class ModelAttn2PatchesMutation: attention_name="attn2", input_patch=self.input_patch, output_patch=self.output_patch, + trailing_input_patches=self.trailing_input_patches, ) diff --git a/simple_syrup/runtime/patcher_lifecycle.py b/simple_syrup/runtime/patcher_lifecycle.py index c344985..4036857 100644 --- a/simple_syrup/runtime/patcher_lifecycle.py +++ b/simple_syrup/runtime/patcher_lifecycle.py @@ -7,7 +7,8 @@ from __future__ import annotations from collections.abc import Iterable -from typing import Protocol, TypeVar, cast +from dataclasses import dataclass +from typing import Any, Protocol, TypeVar, cast from .clip_patcher_model_alignment import align_clip_text_encoder_with_patcher @@ -27,6 +28,14 @@ class ClipMutation(Protocol): PatcherValue = TypeVar("PatcherValue") +_MODEL_FALLBACK_BOUNDARY_ATTACHMENT = "simple_syrup.model_fallback_boundary" + + +@dataclass(frozen=True, slots=True) +class _ModelFallbackBoundary: + """Retain the durable Comfy patcher beneath consecutive Syrup derivations.""" + + patcher: object class ComfyPatcherLifecycle: @@ -48,6 +57,7 @@ class ComfyPatcherLifecycle: disable_dynamic=disable_dynamic, ) self._require_direct_parent(source, derived, operation=operation) + self._stabilize_model_fallback(source, derived) for mutation in mutations: mutation.apply(derived) return derived @@ -66,13 +76,19 @@ class ComfyPatcherLifecycle: getter = getattr(model_override_source, "get_clone_model_override", None) if not callable(getter): raise TypeError(f"{operation} requires a model-override source.") + model_override = getter() derived = self._clone( source, operation=operation, disable_dynamic=disable_dynamic, - model_override=getter(), + model_override=model_override, ) self._require_direct_parent(source, derived, operation=operation) + self._stabilize_model_fallback( + source, + derived, + explicit_boundary=model_override_source, + ) for mutation in mutations: mutation.apply(derived) return derived @@ -107,6 +123,7 @@ class ComfyPatcherLifecycle: derived_patcher, operation=operation, ) + self._stabilize_model_fallback(source_patcher, derived_patcher) align_clip_text_encoder_with_patcher(derived) for mutation in mutations: mutation.apply(derived) @@ -181,6 +198,41 @@ class ComfyPatcherLifecycle: f"{operation} produced a derived patcher without its source as parent." ) + @staticmethod + def _stabilize_model_fallback( + source: object, + derived: object, + *, + explicit_boundary: object | None = None, + ) -> None: + """Collapse Syrup-only lineage onto the durable same-model boundary.""" + + derived_attachments = getattr(derived, "attachments", None) + if not isinstance(derived_attachments, dict): + return + boundary = explicit_boundary + if boundary is None: + source_attachments = getattr(source, "attachments", None) + inherited = ( + source_attachments.get(_MODEL_FALLBACK_BOUNDARY_ATTACHMENT) + if isinstance(source_attachments, dict) + else None + ) + if inherited is not None and not isinstance( + inherited, + _ModelFallbackBoundary, + ): + raise TypeError("SimpleSyrup MODEL fallback boundary is invalid.") + boundary = inherited.patcher if inherited is not None else source + + if getattr(boundary, "model", None) is not getattr(derived, "model", None): + derived_attachments.pop(_MODEL_FALLBACK_BOUNDARY_ATTACHMENT, None) + return + cast(Any, derived).parent = boundary + derived_attachments[_MODEL_FALLBACK_BOUNDARY_ATTACHMENT] = ( + _ModelFallbackBoundary(boundary) + ) + @staticmethod def _required_clip_patcher(value: object, *, value_name: str) -> object: """Return the CLIP patcher required for lineage validation.""" diff --git a/simple_syrup/runtime/ppm_negpip_interop.py b/simple_syrup/runtime/ppm_negpip_interop.py new file mode 100644 index 0000000..b162436 --- /dev/null +++ b/simple_syrup/runtime/ppm_negpip_interop.py @@ -0,0 +1,233 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Adapt the complete installed PPM NegPiP patch family to regional execution.""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from enum import StrEnum +from typing import cast + +import torch +from comfy.patcher_extension import WrappersMP + +from ..domain.regional_model_capabilities import RegionalModelFamily + +_MODEL_MARKER = "ppm_negpip" +_ANIMA_WRAPPER_KEY = "ppm_negpip_anima" +_ANIMA_CONDITION_KEY = "c_ppm_negpip_mask" +_ANIMA_TRANSFORMER_KEY = "ppm_negpip_mask" +_EXTRA_CONDS_PATH = "extra_conds" +_ATTN2_PATCH_NAME = "attn2_patch" +_UNET_CALLBACK = ( + "src.negpip.unet_negpip", + "sdxl_attn2_negpip", +) +_ANIMA_CALLBACK = ( + "src.negpip.anima_negpip", + "cosmos_attn2_negpip", +) +_ANIMA_WRAPPER = ( + "src.negpip.anima_negpip", + "cosmos_diffusion_negpip_wrapper", +) +_ANIMA_EXTRA_CONDS = ( + "src.negpip.anima_negpip", + "anima_extra_conds_negpip_wrapper.._anima_extra_conds_negpip_wrapper", +) + + +class PpmNegpipSemantics(StrEnum): + """Identify the family-specific NegPiP conditioning representation.""" + + STANDARD_UNET_SPLIT_KEY_VALUE = "standard-unet-split-key-value" + ANIMA_VALUE_MASK = "anima-value-mask" + + +@dataclass(frozen=True, slots=True) +class PpmNegpipInterop: + """Retain identity-validated PPM objects needed by regional execution.""" + + semantics: PpmNegpipSemantics + attention_patch: Callable[..., object] + + def __post_init__(self) -> None: + """Require one typed semantic mode and callable preserved callback.""" + + if not isinstance(self.semantics, PpmNegpipSemantics): + raise TypeError("NegPiP semantics have an invalid type.") + if not callable(self.attention_patch): + raise TypeError("NegPiP attention patch must be callable.") + + def extract_value_multiplier( + self, + model_conditions: dict[object, object], + context: torch.Tensor, + ) -> torch.Tensor | None: + """Return one validated Anima value multiplier or no UNet multiplier.""" + + if self.semantics is PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE: + return None + condition = model_conditions.get(_ANIMA_CONDITION_KEY) + if condition is None: + return context.new_ones((*context.shape[:2], 1)) + multiplier = getattr(condition, "cond", None) + if not isinstance(multiplier, torch.Tensor): + raise TypeError("Anima NegPiP value mask condition must contain a tensor.") + if ( + multiplier.ndim != 3 + or int(multiplier.shape[0]) != int(context.shape[0]) + or int(multiplier.shape[1]) != int(context.shape[1]) + or int(multiplier.shape[2]) != 1 + ): + raise ValueError( + "Anima NegPiP value mask must match the conditioning batch and " + "sequence with one multiplier channel." + ) + multiplier = multiplier.to(device=context.device, dtype=context.dtype) + if not bool(((multiplier == 1) | (multiplier == -1)).all().item()): + raise ValueError("Anima NegPiP value mask must contain only -1 and 1.") + return multiplier + + def prepare_anima_transformer_options( + self, + source: dict[str, object], + packed_multiplier: torch.Tensor, + ) -> dict[str, object]: + """Publish a packed mask on an isolated Anima cross-attention call.""" + + if self.semantics is not PpmNegpipSemantics.ANIMA_VALUE_MASK: + raise ValueError("Only Anima NegPiP semantics can publish a value mask.") + if not isinstance(packed_multiplier, torch.Tensor): + raise TypeError("Packed Anima NegPiP multiplier must be a tensor.") + prepared = source.copy() + prepared[_ANIMA_TRANSFORMER_KEY] = packed_multiplier + return prepared + + +class PpmNegpipInteropValidator: + """Admit only one complete identity-validated PPM NegPiP family.""" + + def validate( + self, + family: RegionalModelFamily, + *, + model_options: dict[object, object], + wrappers: dict[str, dict[object, list[object]]], + object_patches: dict[object, object], + transformer_patches: dict[str, list[object]], + ) -> PpmNegpipInterop | None: + """Return preserved NegPiP state or reject every partial/conflicting form.""" + + if not isinstance(family, RegionalModelFamily): + raise TypeError("NegPiP interop requires a model family.") + marker = model_options.get(_MODEL_MARKER, False) + if not isinstance(marker, bool): + raise TypeError("MODEL ppm_negpip marker must be boolean.") + attention = transformer_patches.get(_ATTN2_PATCH_NAME, []) + anima_wrappers = wrappers.get(WrappersMP.DIFFUSION_MODEL, {}).get( + _ANIMA_WRAPPER_KEY, + [], + ) + extra_conds = object_patches.get(_EXTRA_CONDS_PATH) + recognized_surface = any( + ( + any(_is_identity(item, *_UNET_CALLBACK) for item in attention), + any(_is_identity(item, *_ANIMA_CALLBACK) for item in attention), + bool(anima_wrappers), + _is_identity(extra_conds, *_ANIMA_EXTRA_CONDS), + ) + ) + if not marker: + if recognized_surface: + raise ValueError( + "MODEL contains an incomplete NegPiP patch family without its " + "marker. Reapply CLIP NegPip to a clean MODEL." + ) + return None + if family is RegionalModelFamily.STANDARD_UNET: + return self._validate_standard_unet( + attention, + anima_wrappers=anima_wrappers, + extra_conds=extra_conds, + ) + if family is RegionalModelFamily.ANIMA: + return self._validate_anima( + attention, + anima_wrappers=anima_wrappers, + extra_conds=extra_conds, + ) + raise ValueError(f"NegPiP does not support model family {family.value!r}.") + + @staticmethod + def _validate_standard_unet( + attention: list[object], + *, + anima_wrappers: list[object], + extra_conds: object, + ) -> PpmNegpipInterop: + """Require exactly PPM's single UNet split-K/V callback surface.""" + + if ( + len(attention) != 1 + or not _is_identity(attention[0], *_UNET_CALLBACK) + or anima_wrappers + or _is_identity(extra_conds, *_ANIMA_EXTRA_CONDS) + ): + raise ValueError( + "Standard UNet NegPiP requires exactly its PPM split-K/V attention " + "patch and no Anima NegPiP surfaces." + ) + return PpmNegpipInterop( + PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE, + cast(Callable[..., object], attention[0]), + ) + + @staticmethod + def _validate_anima( + attention: list[object], + *, + anima_wrappers: list[object], + extra_conds: object, + ) -> PpmNegpipInterop: + """Require PPM's exact callback, wrapper, and extra-condition surfaces.""" + + if ( + len(attention) != 1 + or not _is_identity(attention[0], *_ANIMA_CALLBACK) + or len(anima_wrappers) != 1 + or not _is_identity(anima_wrappers[0], *_ANIMA_WRAPPER) + or not _is_identity(extra_conds, *_ANIMA_EXTRA_CONDS) + ): + raise ValueError( + "Anima NegPiP requires exactly its PPM attention patch, keyed " + "diffusion wrapper, and extra_conds object patch." + ) + return PpmNegpipInterop( + PpmNegpipSemantics.ANIMA_VALUE_MASK, + cast(Callable[..., object], attention[0]), + ) + + +def _is_identity( + value: object, + module_suffix: str, + qualified_name: str, +) -> bool: + """Match one callable by its stable defining module suffix and qualified name.""" + + if not callable(value): + return False + module = getattr(value, "__module__", None) + qualname = getattr(value, "__qualname__", None) + return ( + isinstance(module, str) + and (module == module_suffix or module.endswith(f".{module_suffix}")) + and qualname == qualified_name + ) + + +PPM_NEGPIP_INTEROP_VALIDATOR = PpmNegpipInteropValidator() diff --git a/simple_syrup/runtime/regional_attention_batching.py b/simple_syrup/runtime/regional_attention_batching.py index 8cef93a..73437b0 100644 --- a/simple_syrup/runtime/regional_attention_batching.py +++ b/simple_syrup/runtime/regional_attention_batching.py @@ -7,6 +7,7 @@ from __future__ import annotations import torch +import torch.nn.functional as functional from comfy.utils import repeat_to_batch_size from ..domain.processed_regional_attention import ( @@ -143,6 +144,13 @@ class RegionalAttentionBatchingService: ) for region_index in range(plan.mask_bank.region_count) ), + base_value_multiplier=self._align_value_multiplier( + tuple(chunk.base_entry for chunk in selected), + latent_batch_size=latent_batch_size, + device=aligned_base_context.device, + dtype=aligned_base_context.dtype, + target_sequence_length=target_sequence_length, + ), ) def _align_region( @@ -163,6 +171,7 @@ class RegionalAttentionBatchingService: for entry_index in range(entry_count): context_parts: list[torch.Tensor] = [] strengths: list[float] = [] + multiplier_entries: list[ProcessedRegionalAttentionEntry] = [] for chunk, source in zip(chunks, sources, strict=True): if source is None: entry = chunk.base_entry @@ -186,12 +195,20 @@ class RegionalAttentionBatchingService: target_length=target_sequence_length, ) context_parts.append(repeated) + multiplier_entries.append(entry) strengths.extend((strength,) * latent_batch_size) entries.append( BatchedRegionalAttentionEntry( entry_index, torch.cat(context_parts, dim=0), tuple(strengths), + self._align_value_multiplier( + tuple(multiplier_entries), + latent_batch_size=latent_batch_size, + device=device, + dtype=dtype, + target_sequence_length=target_sequence_length, + ), ) ) return BatchedRegionalAttentionRegion(region_index, tuple(entries)) @@ -201,6 +218,7 @@ class RegionalAttentionBatchingService: entries = self._entries(plan) authority = entries[0].cross_attention + has_value_multiplier = entries[0].cross_attention_value_multiplier is not None for entry in entries[1:]: tensor = entry.cross_attention shape_mismatch = ( @@ -216,6 +234,48 @@ class RegionalAttentionBatchingService: raise ValueError("Regional attention context devices must match.") if tensor.dtype != authority.dtype: raise ValueError("Regional attention context dtypes must match.") + if ( + entry.cross_attention_value_multiplier is not None + ) is not has_value_multiplier: + raise ValueError( + "Regional attention value multiplier presence must be uniform." + ) + + @staticmethod + def _align_value_multiplier( + entries: tuple[ProcessedRegionalAttentionEntry, ...], + *, + latent_batch_size: int, + device: torch.device, + dtype: torch.dtype, + target_sequence_length: int, + ) -> torch.Tensor | None: + """Repeat and sequence-align one chunk-major value-multiplier bank.""" + + if not entries or entries[0].cross_attention_value_multiplier is None: + return None + parts: list[torch.Tensor] = [] + for entry in entries: + multiplier = entry.cross_attention_value_multiplier + if multiplier is None: + raise ValueError( + "Regional attention value multiplier presence must be uniform." + ) + repeated = repeat_to_batch_size(multiplier, latent_batch_size).to( + device=device, + dtype=dtype, + ) + sequence_length = int(repeated.shape[1]) + if sequence_length < target_sequence_length: + repeated = functional.pad( + repeated, + (0, 0, 0, target_sequence_length - sequence_length), + value=1.0, + ) + elif sequence_length > target_sequence_length: + repeated = repeated[:, :target_sequence_length] + parts.append(repeated) + return torch.cat(parts, dim=0) @staticmethod def _entries( diff --git a/simple_syrup/runtime/regional_lora/anima_attention_coupling.py b/simple_syrup/runtime/regional_lora/anima_attention_coupling.py index bcaefdf..db8af91 100644 --- a/simple_syrup/runtime/regional_lora/anima_attention_coupling.py +++ b/simple_syrup/runtime/regional_lora/anima_attention_coupling.py @@ -7,6 +7,7 @@ from __future__ import annotations from ..patcher_lifecycle import ModelMutation +from ..ppm_negpip_interop import PpmNegpipInterop from .anima_activation_context import ( ANIMA_ACTIVATION_CONTEXT, AnimaActivationContext, @@ -75,6 +76,7 @@ def anima_attention_coupling_mutations( ), phase_context: AnimaCompositionPhaseContext = (ANIMA_COMPOSITION_PHASE_CONTEXT), query_mask_context: AnimaQueryMaskContext = ANIMA_QUERY_MASK_CONTEXT, + negpip: PpmNegpipInterop | None = None, ) -> tuple[ModelMutation, ...]: """Return attention-only or complete regional-LoRA mutation composition.""" @@ -106,6 +108,7 @@ def anima_attention_coupling_mutations( invocation_context=cross_attention_context, phase_context=phase_context, query_activity=query_activity, + negpip=negpip, ) phase_wrapper = anima_composition_phase_wrapper_mutation( surface, diff --git a/simple_syrup/runtime/regional_lora/anima_cross_attention.py b/simple_syrup/runtime/regional_lora/anima_cross_attention.py index 5a7ea3f..70f6c5f 100644 --- a/simple_syrup/runtime/regional_lora/anima_cross_attention.py +++ b/simple_syrup/runtime/regional_lora/anima_cross_attention.py @@ -18,6 +18,7 @@ from ...domain.regional_conditioning_output import ( RegionalConditioningOutputCombiner, ) from ..model_patcher_mutations import ModelExactObjectPatchMutation +from ..ppm_negpip_interop import PpmNegpipInterop from .anima_activation_context import ( ANIMA_ACTIVATION_CONTEXT, AnimaActivationContext, @@ -26,6 +27,7 @@ from .anima_activation_context import ( from .anima_attention_execution import AnimaRegionalAttentionExecution from .anima_branch_batch import ( ANIMA_BASE_BRANCH_KEY, + AnimaRegionalBranchBatch, AnimaRegionalBranchKey, ) from .anima_composition_phase_context import AnimaCompositionPhaseContext @@ -71,6 +73,7 @@ class AnimaRegionalCrossAttentionPatch(nn.Module): ), weighting: RegionalAttentionWeightingPolicy | None = None, entry_combiner: RegionalConditioningOutputCombiner | None = None, + negpip: PpmNegpipInterop | None = None, ) -> None: """Retain the exact original attention owner and focused collaborators.""" @@ -92,6 +95,9 @@ class AnimaRegionalCrossAttentionPatch(nn.Module): self._query_activity = query_activity self._weighting = weighting or ANIMA_CROSS_ATTENTION_WEIGHTING_POLICY self._entry_combiner = entry_combiner or REGIONAL_CONDITIONING_OUTPUT_COMBINER + if negpip is not None and not isinstance(negpip, PpmNegpipInterop): + raise TypeError("Anima cross-attention NegPiP state has an invalid type.") + self._negpip = negpip def __setattr__(self, name: str, value: Any) -> None: """Keep later exact child patches synchronized with installed attention.""" @@ -143,14 +149,18 @@ class AnimaRegionalCrossAttentionPatch(nn.Module): branch_batch = activity.attention_branches branch_x = branch_batch.pack_source(x) branch_context = branch_batch.pack_branch_values(branch_values) + original_options = {} if transformer_options is None else transformer_options + forwarded_options = self._prepare_transformer_options( + original_options, + branch_batch, + execution_contexts, + ) with self._invocation_context.activate(branch_batch.invocation): branch_output = self._backing.module( branch_x, branch_context, rope_emb=rope_emb, - transformer_options=( - {} if transformer_options is None else transformer_options - ), + transformer_options=forwarded_options, ) if not isinstance(branch_output, torch.Tensor): raise TypeError("Original Anima cross-attention must return a tensor.") @@ -186,6 +196,41 @@ class AnimaRegionalCrossAttentionPatch(nn.Module): regional_outputs=torch.stack(regional_outputs), ) + def _prepare_transformer_options( + self, + source: dict[str, Any], + branch_batch: AnimaRegionalBranchBatch, + contexts: BatchedRegionalAttentionContexts, + ) -> dict[str, object]: + """Pack NegPiP multipliers or preserve ordinary option identity.""" + + if self._negpip is None: + if contexts.base_value_multiplier is not None: + raise ValueError( + "Anima value multipliers require admitted NegPiP semantics." + ) + return source + if not isinstance(branch_batch, AnimaRegionalBranchBatch): + raise TypeError("Anima NegPiP branch batch has an invalid type.") + base_multiplier = contexts.base_value_multiplier + if base_multiplier is None: + raise ValueError("Anima NegPiP requires an aligned base value multiplier.") + values = {ANIMA_BASE_BRANCH_KEY: base_multiplier} + for region in contexts.regions: + for entry in region.entries: + multiplier = entry.cross_attention_value_multiplier + if multiplier is None: + raise ValueError( + "Anima NegPiP requires every regional value multiplier." + ) + values[ + AnimaRegionalBranchKey(region.region_index, entry.entry_index) + ] = multiplier + return self._negpip.prepare_anima_transformer_options( + source, + branch_batch.pack_branch_values(values), + ) + def _validate_inputs( self, *, @@ -246,6 +291,7 @@ def anima_cross_attention_mutations( query_activity: AnimaRegionalQueryActivityContext = ( ANIMA_REGIONAL_QUERY_ACTIVITY_CONTEXT ), + negpip: PpmNegpipInterop | None = None, ) -> tuple[ModelExactObjectPatchMutation, ...]: """Build one exact clone-local cross-attention replacement per Anima block.""" @@ -262,6 +308,7 @@ def anima_cross_attention_mutations( invocation_context=invocation_context, phase_context=phase_context, query_activity=query_activity, + negpip=negpip, ), ) for block in surface.blocks diff --git a/simple_syrup/runtime/regional_lora/anima_full_context_backend.py b/simple_syrup/runtime/regional_lora/anima_full_context_backend.py index fd2c4e0..f7d5969 100644 --- a/simple_syrup/runtime/regional_lora/anima_full_context_backend.py +++ b/simple_syrup/runtime/regional_lora/anima_full_context_backend.py @@ -10,6 +10,7 @@ from dataclasses import dataclass from ...domain.processed_regional_attention import ProcessedRegionalAttentionPlan from ..patcher_lifecycle import PATCHER_LIFECYCLE +from ..ppm_negpip_interop import PpmNegpipInterop from ..regional_attention_template import build_regional_attention_template from ..regional_lora_plan_adapter import RegionalLoraPlanAdaptation from .anima_attention_context_wrapper import ( @@ -44,6 +45,7 @@ class FullContextAnimaAttentionBackend: adaptation: RegionalLoraPlanAdaptation, region_strengths: tuple[float, ...], latent_batch_size: int, + negpip: PpmNegpipInterop | None = None, ) -> FullContextAnimaAttentionModel: """Return one collision-safe clone prepared for dynamic sampler calls.""" @@ -51,6 +53,8 @@ class FullContextAnimaAttentionBackend: raise TypeError("Anima backend requires a processed attention plan.") if not isinstance(adaptation, RegionalLoraPlanAdaptation): raise TypeError("Anima backend requires a regional LoRA adaptation.") + if negpip is not None and not isinstance(negpip, PpmNegpipInterop): + raise TypeError("Anima backend NegPiP state has an invalid type.") admitted = ANIMA_REGIONAL_LORA_PLAN_ADMISSION_SERVICE.admit(adaptation) ANIMA_GLOBAL_REGIONAL_LORA_OVERLAP_VALIDATOR.validate(model, admitted) template = build_regional_attention_template( @@ -89,6 +93,7 @@ class FullContextAnimaAttentionBackend: surface, attention, composition=composition, + negpip=negpip, ), ) derived = PATCHER_LIFECYCLE.derive_model( diff --git a/simple_syrup/runtime/regional_lora/fused_active_accumulation.py b/simple_syrup/runtime/regional_lora/fused_active_accumulation.py index b5cee4c..7a11d13 100644 --- a/simple_syrup/runtime/regional_lora/fused_active_accumulation.py +++ b/simple_syrup/runtime/regional_lora/fused_active_accumulation.py @@ -5,22 +5,62 @@ from __future__ import annotations +from types import ModuleType +from typing import Protocol, cast, runtime_checkable + import torch -from .fused_active_accumulation_kernel import ( - MAX_FUSED_ADAPTERS_PER_LAUNCH, - REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL, +from .fused_active_accumulation_contract import MAX_FUSED_ADAPTERS_PER_LAUNCH +from .triton_runtime import TRITON_RUNTIME_RESOLVER, TritonRuntimeResolver + +_TRITON_BACKEND_MODULE = ( + "simple_syrup.runtime.regional_lora.fused_active_accumulation_kernel" ) +@runtime_checkable +class _FusedAccumulationKernel(Protocol): + """Describe the lazily resolved fused CUDA launch surface.""" + + def launch( + self, + output: torch.Tensor, + *, + rank_values: torch.Tensor, + up: torch.Tensor, + multipliers: tuple[torch.Tensor, ...], + indices: torch.Tensor | None, + adapter_start: int, + target_indices: torch.Tensor | None = None, + ) -> None: + """Launch one validated fused accumulation chunk.""" + + ... + + +class _FusedAccumulationBackend(Protocol): + """Describe the exported lazy backend module surface.""" + + REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL: _FusedAccumulationKernel + + class RegionalLoraFusedActiveAccumulator: """Own verified CUDA fusion for B projection and ordered output updates.""" - @staticmethod - def supports(output: torch.Tensor, up: torch.Tensor) -> bool: + def __init__( + self, + resolver: TritonRuntimeResolver = TRITON_RUNTIME_RESOLVER, + ) -> None: + """Retain the process-level optional acceleration authority.""" + + if not isinstance(resolver, TritonRuntimeResolver): + raise TypeError("Fused accumulation requires a Triton resolver.") + self._resolver = resolver + + def supports(self, output: torch.Tensor, up: torch.Tensor) -> bool: """Admit only verified contiguous CUDA bf16/fp16 projection shapes.""" - return ( + eligible = ( isinstance(output, torch.Tensor) and isinstance(up, torch.Tensor) and output.device.type == "cuda" @@ -32,6 +72,7 @@ class RegionalLoraFusedActiveAccumulator: and output.is_contiguous() and up.is_contiguous() ) + return eligible and self._backend() is not None def add( self, @@ -59,7 +100,7 @@ class RegionalLoraFusedActiveAccumulator: return output for start in range(0, adapter_count, MAX_FUSED_ADAPTERS_PER_LAUNCH): stop = min(start + MAX_FUSED_ADAPTERS_PER_LAUNCH, adapter_count) - REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL.launch( + self._require_backend().launch( output, rank_values=rank_values, up=up, @@ -97,7 +138,7 @@ class RegionalLoraFusedActiveAccumulator: ) if active_row_count == 0: return output - REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL.launch( + self._require_backend().launch( output, rank_values=rank_values, up=up, @@ -137,7 +178,7 @@ class RegionalLoraFusedActiveAccumulator: return output for start in range(0, len(multipliers), MAX_FUSED_ADAPTERS_PER_LAUNCH): stop = min(start + MAX_FUSED_ADAPTERS_PER_LAUNCH, len(multipliers)) - REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL.launch( + self._require_backend().launch( output, rank_values=rank_values, up=up, @@ -148,9 +189,28 @@ class RegionalLoraFusedActiveAccumulator: ) return output - @classmethod + def _backend(self) -> _FusedAccumulationKernel | None: + """Return the cached optional fused kernel without hiding failures.""" + + backend = self._resolver.resolve(_TRITON_BACKEND_MODULE) + if backend is None: + return None + module = cast(_FusedAccumulationBackend, cast(ModuleType, backend)) + kernel = module.REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL + if not isinstance(kernel, _FusedAccumulationKernel): + raise TypeError("Triton fused backend has an invalid kernel surface.") + return kernel + + def _require_backend(self) -> _FusedAccumulationKernel: + """Return the admitted kernel or reject an invalid direct fused call.""" + + backend = self._backend() + if backend is None: + raise RuntimeError("Triton fused accumulation is unavailable.") + return backend + def _validate_mapped_projection( - cls, + self, output: torch.Tensor, rank_values: torch.Tensor, up: torch.Tensor, @@ -159,7 +219,7 @@ class RegionalLoraFusedActiveAccumulator: ) -> None: """Require one unique-target batch and valid declared group mapping.""" - if not cls.supports(output, up): + if not self.supports(output, up): raise ValueError("Mapped fused accumulation received an unsupported path.") if ( rank_values.device != output.device @@ -191,9 +251,8 @@ class RegionalLoraFusedActiveAccumulator: ): raise ValueError("Mapped fused output indices are misaligned.") - @classmethod def _validate_projection( - cls, + self, output: torch.Tensor, rank_values: torch.Tensor, up: torch.Tensor, @@ -201,7 +260,7 @@ class RegionalLoraFusedActiveAccumulator: ) -> None: """Require one complete aligned CUDA projection contract.""" - if not cls.supports(output, up): + if not self.supports(output, up): raise ValueError("Fused active accumulation received an unsupported path.") if ( rank_values.device != output.device diff --git a/simple_syrup/runtime/regional_lora/fused_active_accumulation_contract.py b/simple_syrup/runtime/regional_lora/fused_active_accumulation_contract.py new file mode 100644 index 0000000..232cb37 --- /dev/null +++ b/simple_syrup/runtime/regional_lora/fused_active_accumulation_contract.py @@ -0,0 +1,7 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define the shared launch-width contract for fused LoRA accumulation.""" + +MAX_FUSED_ADAPTERS_PER_LAUNCH = 8 diff --git a/simple_syrup/runtime/regional_lora/fused_active_accumulation_kernel.py b/simple_syrup/runtime/regional_lora/fused_active_accumulation_kernel.py index dba343a..7e4d152 100644 --- a/simple_syrup/runtime/regional_lora/fused_active_accumulation_kernel.py +++ b/simple_syrup/runtime/regional_lora/fused_active_accumulation_kernel.py @@ -14,7 +14,8 @@ import torch import triton # type: ignore[import-untyped] import triton.language as tl # type: ignore[import-untyped] -MAX_FUSED_ADAPTERS_PER_LAUNCH = 8 +from .fused_active_accumulation_contract import MAX_FUSED_ADAPTERS_PER_LAUNCH + _BLOCK_ROWS = 16 _BLOCK_OUTPUT_FEATURES = 64 diff --git a/simple_syrup/runtime/regional_lora/ordered_accumulation.py b/simple_syrup/runtime/regional_lora/ordered_accumulation.py index 60ad362..fea3ef2 100644 --- a/simple_syrup/runtime/regional_lora/ordered_accumulation.py +++ b/simple_syrup/runtime/regional_lora/ordered_accumulation.py @@ -1,67 +1,48 @@ # SimpleSyrup - workflow-focused ComfyUI extensions for image generation # Copyright (C) 2026 Artificial Sweetener and contributors # SPDX-License-Identifier: AGPL-3.0-or-later -# mypy: disable-error-code="no-untyped-def" -# ruff: noqa: ANN001, ANN202 - """Accumulate ordered adapter outputs with exact execution-dtype rounding.""" from __future__ import annotations -from typing import Any, cast +from types import ModuleType +from typing import Protocol, cast import torch -import triton # type: ignore[import-untyped] -import triton.language as tl # type: ignore[import-untyped] -_MAX_DELTAS_PER_LAUNCH = 8 -_BLOCK_SIZE = 256 +from .triton_runtime import TRITON_RUNTIME_RESOLVER, TritonRuntimeResolver + +_TRITON_BACKEND_MODULE = ( + "simple_syrup.runtime.regional_lora.ordered_accumulation_triton" +) -@triton.jit # type: ignore[untyped-decorator] -def _ordered_accumulation_kernel( - output, - base, - delta_0, - delta_1, - delta_2, - delta_3, - delta_4, - delta_5, - delta_6, - delta_7, - element_count, - delta_count: tl.constexpr, - execution_dtype: tl.constexpr, - block_size: tl.constexpr, -): - """Add one ordered chunk and round after every declared adapter.""" +class _OrderedAccumulationBackend(Protocol): + """Describe the lazy CUDA backend surface consumed by this owner.""" - offsets = tl.program_id(0) * block_size + tl.arange(0, block_size) - active = offsets < element_count - value = tl.load(base + offsets, mask=active) - if delta_count > 0: - value = (value + tl.load(delta_0 + offsets, mask=active)).to(execution_dtype) - if delta_count > 1: - value = (value + tl.load(delta_1 + offsets, mask=active)).to(execution_dtype) - if delta_count > 2: - value = (value + tl.load(delta_2 + offsets, mask=active)).to(execution_dtype) - if delta_count > 3: - value = (value + tl.load(delta_3 + offsets, mask=active)).to(execution_dtype) - if delta_count > 4: - value = (value + tl.load(delta_4 + offsets, mask=active)).to(execution_dtype) - if delta_count > 5: - value = (value + tl.load(delta_5 + offsets, mask=active)).to(execution_dtype) - if delta_count > 6: - value = (value + tl.load(delta_6 + offsets, mask=active)).to(execution_dtype) - if delta_count > 7: - value = (value + tl.load(delta_7 + offsets, mask=active)).to(execution_dtype) - tl.store(output + offsets, value, mask=active) + def accumulate( + self, + base: torch.Tensor, + deltas: tuple[torch.Tensor, ...], + ) -> torch.Tensor: + """Accumulate validated CUDA tensors in declared order.""" + + ... class OrderedTensorAccumulator: """Own exact ordered accumulation and its CUDA launch policy.""" + def __init__( + self, + resolver: TritonRuntimeResolver = TRITON_RUNTIME_RESOLVER, + ) -> None: + """Retain the process-level optional acceleration authority.""" + + if not isinstance(resolver, TritonRuntimeResolver): + raise TypeError("Ordered accumulation requires a Triton resolver.") + self._resolver = resolver + def accumulate( self, base: torch.Tensor, @@ -74,25 +55,13 @@ class OrderedTensorAccumulator: return base if base.device.type != "cuda": return self._torch_accumulate(base, deltas) - execution_dtype = self._triton_dtype(base.dtype) - remaining = deltas - result = base - while remaining: - chunk = remaining[:_MAX_DELTAS_PER_LAUNCH] - remaining = remaining[_MAX_DELTAS_PER_LAUNCH:] - padded = (*chunk, *((result,) * (_MAX_DELTAS_PER_LAUNCH - len(chunk)))) - grid = (triton.cdiv(result.numel(), _BLOCK_SIZE),) - kernel = cast(Any, _ordered_accumulation_kernel) - kernel[grid]( - result, - result, - *padded, - result.numel(), - delta_count=len(chunk), - execution_dtype=execution_dtype, - block_size=_BLOCK_SIZE, - ) - return result + backend = self._resolver.resolve(_TRITON_BACKEND_MODULE) + if backend is None: + return self._torch_accumulate(base, deltas) + return cast(_OrderedAccumulationBackend, cast(ModuleType, backend)).accumulate( + base, + deltas, + ) @staticmethod def _torch_accumulate( @@ -106,18 +75,6 @@ class OrderedTensorAccumulator: result.add_(delta) return result - @staticmethod - def _triton_dtype(dtype: torch.dtype) -> Any: - """Map the admitted floating execution dtype to a Triton scalar dtype.""" - - if dtype is torch.bfloat16: - return tl.bfloat16 - if dtype is torch.float16: - return tl.float16 - if dtype is torch.float32: - return tl.float32 - raise TypeError(f"Ordered CUDA accumulation does not support {dtype}.") - @staticmethod def _validate( base: torch.Tensor, diff --git a/simple_syrup/runtime/regional_lora/ordered_accumulation_triton.py b/simple_syrup/runtime/regional_lora/ordered_accumulation_triton.py new file mode 100644 index 0000000..bc44b62 --- /dev/null +++ b/simple_syrup/runtime/regional_lora/ordered_accumulation_triton.py @@ -0,0 +1,95 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later +# mypy: disable-error-code="no-untyped-def" +# ruff: noqa: ANN001, ANN202 + +"""Provide the lazily imported Triton ordered-accumulation backend.""" + +from __future__ import annotations + +from typing import Any, cast + +import torch +import triton # type: ignore[import-untyped] +import triton.language as tl # type: ignore[import-untyped] + +_MAX_DELTAS_PER_LAUNCH = 8 +_BLOCK_SIZE = 256 + + +@triton.jit # type: ignore[untyped-decorator] +def _ordered_accumulation_kernel( + output, + base, + delta_0, + delta_1, + delta_2, + delta_3, + delta_4, + delta_5, + delta_6, + delta_7, + element_count, + delta_count: tl.constexpr, + execution_dtype: tl.constexpr, + block_size: tl.constexpr, +): + """Add one ordered chunk and round after every declared adapter.""" + + offsets = tl.program_id(0) * block_size + tl.arange(0, block_size) + active = offsets < element_count + value = tl.load(base + offsets, mask=active) + if delta_count > 0: + value = (value + tl.load(delta_0 + offsets, mask=active)).to(execution_dtype) + if delta_count > 1: + value = (value + tl.load(delta_1 + offsets, mask=active)).to(execution_dtype) + if delta_count > 2: + value = (value + tl.load(delta_2 + offsets, mask=active)).to(execution_dtype) + if delta_count > 3: + value = (value + tl.load(delta_3 + offsets, mask=active)).to(execution_dtype) + if delta_count > 4: + value = (value + tl.load(delta_4 + offsets, mask=active)).to(execution_dtype) + if delta_count > 5: + value = (value + tl.load(delta_5 + offsets, mask=active)).to(execution_dtype) + if delta_count > 6: + value = (value + tl.load(delta_6 + offsets, mask=active)).to(execution_dtype) + if delta_count > 7: + value = (value + tl.load(delta_7 + offsets, mask=active)).to(execution_dtype) + tl.store(output + offsets, value, mask=active) + + +def accumulate(base: torch.Tensor, deltas: tuple[torch.Tensor, ...]) -> torch.Tensor: + """Accumulate CUDA deltas with the established launch and rounding policy.""" + + execution_dtype = _triton_dtype(base.dtype) + remaining = deltas + result = base + while remaining: + chunk = remaining[:_MAX_DELTAS_PER_LAUNCH] + remaining = remaining[_MAX_DELTAS_PER_LAUNCH:] + padded = (*chunk, *((result,) * (_MAX_DELTAS_PER_LAUNCH - len(chunk)))) + grid = (triton.cdiv(result.numel(), _BLOCK_SIZE),) + kernel = cast(Any, _ordered_accumulation_kernel) + kernel[grid]( + result, + result, + *padded, + result.numel(), + delta_count=len(chunk), + execution_dtype=execution_dtype, + block_size=_BLOCK_SIZE, + ) + return result + + +def _triton_dtype(dtype: torch.dtype) -> Any: + """Map the admitted floating execution dtype to a Triton scalar dtype.""" + + if dtype is torch.bfloat16: + return tl.bfloat16 + if dtype is torch.float16: + return tl.float16 + if dtype is torch.float32: + return tl.float32 + raise TypeError(f"Ordered CUDA accumulation does not support {dtype}.") diff --git a/simple_syrup/runtime/regional_lora/standard_unet_variant_base_attention.py b/simple_syrup/runtime/regional_lora/standard_unet_variant_base_attention.py index 8d59e02..c68e104 100644 --- a/simple_syrup/runtime/regional_lora/standard_unet_variant_base_attention.py +++ b/simple_syrup/runtime/regional_lora/standard_unet_variant_base_attention.py @@ -10,6 +10,7 @@ from ..attention_coupling.unet_attn2_execution_resolver import ( UnetAttn2ExecutionResolver, ) from ..attention_coupling.unet_attn2_patch import UnetAttn2PatchPair +from ..ppm_negpip_interop import PpmNegpipInterop _INPUT_PATCH_KEY = "attn2_patch" _OUTPUT_PATCH_KEY = "attn2_output_patch" @@ -18,12 +19,20 @@ _OUTPUT_PATCH_KEY = "attn2_output_patch" class StandardUnetVariantBaseAttention: """Install Attention Couple only on graph-local unpatched base execution.""" - def __init__(self, resolver: UnetAttn2ExecutionResolver) -> None: + def __init__( + self, + resolver: UnetAttn2ExecutionResolver, + *, + negpip: PpmNegpipInterop | None = None, + ) -> None: """Retain one paired callback authority for the active request state.""" if not isinstance(resolver, UnetAttn2ExecutionResolver): raise TypeError("Standard UNet base attention requires a resolver.") self._patches = UnetAttn2PatchPair(resolver) + if negpip is not None and not isinstance(negpip, PpmNegpipInterop): + raise TypeError("Standard UNet base attention NegPiP state is invalid.") + self._negpip = negpip def prepare(self, source: dict[str, object]) -> dict[str, object]: """Return isolated transformer options with one collision-free pair.""" @@ -40,10 +49,22 @@ class StandardUnetVariantBaseAttention: patches = source_patches.copy() else: raise TypeError("Standard UNet transformer patches must be a dictionary.") - for key in (_INPUT_PATCH_KEY, _OUTPUT_PATCH_KEY): - if key in patches: - raise ValueError(f"Standard UNet base graph already contains {key!r}.") - patches[_INPUT_PATCH_KEY] = [self._patches.input_patch] + expected_input = [] if self._negpip is None else [self._negpip.attention_patch] + if ( + self._negpip is None + and _INPUT_PATCH_KEY in patches + or self._negpip is not None + and patches.get(_INPUT_PATCH_KEY) != expected_input + ): + raise ValueError( + "Standard UNet base graph already contains an unadmitted attn2 " + "input patch." + ) + if _OUTPUT_PATCH_KEY in patches: + raise ValueError( + "Standard UNet base graph already contains 'attn2_output_patch'." + ) + patches[_INPUT_PATCH_KEY] = [self._patches.input_patch, *expected_input] patches[_OUTPUT_PATCH_KEY] = [self._patches.output_patch] prepared["patches"] = patches return prepared diff --git a/simple_syrup/runtime/regional_lora/standard_unet_variant_runtime.py b/simple_syrup/runtime/regional_lora/standard_unet_variant_runtime.py index 53149a9..1eed337 100644 --- a/simple_syrup/runtime/regional_lora/standard_unet_variant_runtime.py +++ b/simple_syrup/runtime/regional_lora/standard_unet_variant_runtime.py @@ -22,6 +22,7 @@ from ..model_patcher_mutations import ( ModelKeyedCallbackMutation, ModelKeyedWrapperMutation, ) +from ..ppm_negpip_interop import PpmNegpipInterop from .standard_unet_cold_sampling import ( StandardUnetColdSamplingDiagnosticsMutation, ) @@ -48,6 +49,7 @@ class StandardUnetVariantRuntimeMutation: admission: StandardUnetNativeLoraAdmission attention_phase: StandardUnetAttentionPhaseSession template: StandardUnetVariantTemplate + negpip: PpmNegpipInterop | None = None def apply(self, model: object) -> None: """Build persistent variants before installing the private root clone.""" @@ -69,7 +71,8 @@ class StandardUnetVariantRuntimeMutation: ), attention_phase=self.attention_phase, base_attention=StandardUnetVariantBaseAttention( - StandardUnetAttn2ExecutionResolver(self.state) + StandardUnetAttn2ExecutionResolver(self.state), + negpip=self.negpip, ), ) execution.prime() diff --git a/simple_syrup/runtime/regional_lora/standard_unet_variant_template.py b/simple_syrup/runtime/regional_lora/standard_unet_variant_template.py index 64203c7..4094a62 100644 --- a/simple_syrup/runtime/regional_lora/standard_unet_variant_template.py +++ b/simple_syrup/runtime/regional_lora/standard_unet_variant_template.py @@ -148,8 +148,10 @@ class StandardUnetVariantTemplate: ) if not isinstance(request, ModelPatcher) or request.is_dynamic(): raise TypeError("Comfy did not bind a static standard-UNet request.") - if request.parent is not source: - raise RuntimeError("Static standard-UNet request lost source lineage.") + if request.parent is not self.model: + raise RuntimeError( + "Static standard-UNet request lost its template fallback boundary." + ) if "diffusion_model" in request.object_patches_backup: ModelSharedObjectPatchMutation( "diffusion_model", diff --git a/simple_syrup/runtime/regional_lora/triton_runtime.py b/simple_syrup/runtime/regional_lora/triton_runtime.py new file mode 100644 index 0000000..7bdea54 --- /dev/null +++ b/simple_syrup/runtime/regional_lora/triton_runtime.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 + +"""Resolve optional Triton backends without importing them on Torch paths.""" + +from __future__ import annotations + +import logging +from collections.abc import Callable +from importlib import import_module +from threading import Lock + +LOGGER = logging.getLogger(__name__) + + +class TritonRuntimeResolver: + """Own thread-safe lazy backend imports and optional-package fallback.""" + + def __init__( + self, + *, + import_module: Callable[[str], object] = import_module, + ) -> None: + """Retain an injectable importer and empty process-lifetime cache.""" + + if not callable(import_module): + raise TypeError("Triton runtime importer must be callable.") + self._import_module = import_module + self._lock = Lock() + self._backends: dict[str, object | None] = {} + self._missing_warning_emitted = False + + def resolve(self, backend_module: str) -> object | None: + """Return one cached backend or None only when Triton is absent.""" + + if not isinstance(backend_module, str) or not backend_module: + raise ValueError("Triton backend module must be a non-empty string.") + with self._lock: + if backend_module in self._backends: + return self._backends[backend_module] + try: + backend = self._import_module(backend_module) + except ModuleNotFoundError as error: + if error.name != "triton" and not ( + isinstance(error.name, str) and error.name.startswith("triton.") + ): + raise RuntimeError( + f"Triton backend {backend_module!r} failed to initialize." + ) from error + backend = None + if not self._missing_warning_emitted: + LOGGER.warning( + "Triton acceleration is unavailable; using the Torch " + "execution path", + extra={ + "backend_module": backend_module, + "missing_dependency": error.name, + }, + ) + self._missing_warning_emitted = True + except Exception as error: + raise RuntimeError( + f"Triton backend {backend_module!r} failed to initialize." + ) from error + self._backends[backend_module] = backend + return backend + + +TRITON_RUNTIME_RESOLVER = TritonRuntimeResolver() diff --git a/simple_syrup/runtime/regional_model_patch_interop.py b/simple_syrup/runtime/regional_model_patch_interop.py index 0b54a39..3d7d4e6 100644 --- a/simple_syrup/runtime/regional_model_patch_interop.py +++ b/simple_syrup/runtime/regional_model_patch_interop.py @@ -20,11 +20,13 @@ from ..domain.regional_model_capabilities import ( RegionalModelFamily, RegionalPatchConflict, ) +from .ppm_negpip_interop import ( + PPM_NEGPIP_INTEROP_VALIDATOR, + PpmNegpipInterop, +) LOGGER = logging.getLogger(__name__) -_NEGPIP_MODEL_OPTION = "ppm_negpip" -_NEGPIP_ANIMA_WRAPPER_KEY = "ppm_negpip_anima" _EASYCACHE_OPTION = "easycache" _ATTN2_PATCH_CONFLICTS = { RegionalPatchConflict.ATTN2_INPUT_PATCH: "attn2_patch", @@ -49,6 +51,7 @@ class RegionalModelPatchInteropReport: model_family: RegionalModelFamily preserved_modifiers: tuple[RegionalPreservedModelModifier, ...] + negpip: PpmNegpipInterop | None = None def __post_init__(self) -> None: """Require a typed family and canonical unique modifier order.""" @@ -62,6 +65,8 @@ class RegionalModelPatchInteropReport: raise TypeError("Regional interop report modifiers have invalid types.") if len(set(self.preserved_modifiers)) != len(self.preserved_modifiers): raise ValueError("Regional interop report modifiers must be unique.") + if self.negpip is not None and not isinstance(self.negpip, PpmNegpipInterop): + raise TypeError("Regional interop report NegPiP state has an invalid type.") @property def cache_modifier(self) -> RegionalPreservedModelModifier | None: @@ -99,12 +104,22 @@ class RegionalModelPatchInteropValidator: model_weight_patches = _require_dictionary_attribute(model, "patches") patches = _require_optional_patch_state(transformer_options) - self._reject_negpip(model_options, wrappers) + negpip = PPM_NEGPIP_INTEROP_VALIDATOR.validate( + capabilities.model_family, + model_options=model_options, + wrappers=wrappers, + object_patches=object_patches, + transformer_patches=patches, + ) cache_modifier = self._validate_cache_state( transformer_options, wrappers, ) - self._reject_attention_collisions(patches, capabilities) + self._reject_attention_collisions( + patches, + capabilities, + admitted_negpip=negpip, + ) modifiers: list[RegionalPreservedModelModifier] = [] model_wrapper = model_options.get("model_function_wrapper") @@ -135,6 +150,7 @@ class RegionalModelPatchInteropValidator: report = RegionalModelPatchInteropReport( capabilities.model_family, tuple(modifiers), + negpip, ) LOGGER.info( "Regional MODEL patch interoperability admitted", @@ -188,30 +204,6 @@ class RegionalModelPatchInteropValidator: }, ) - @staticmethod - def _reject_negpip( - model_options: dict[object, object], - wrappers: dict[str, dict[object, list[object]]], - ) -> None: - """Reject installed NegPiP before its mask can enter branch packing.""" - - marker = model_options.get(_NEGPIP_MODEL_OPTION, False) - if not isinstance(marker, bool): - raise TypeError("MODEL ppm_negpip marker must be boolean.") - negpip_wrapper = bool( - wrappers.get(WrappersMP.DIFFUSION_MODEL, {}).get( - _NEGPIP_ANIMA_WRAPPER_KEY, - (), - ) - ) - if marker or negpip_wrapper: - raise ValueError( - "Attention Coupling does not support NegPiP because its attention " - "mask is aligned to the ordinary conditioning batch rather than " - "SimpleSyrup's regional branch batch. Remove CLIP NegPip before " - "the Attention Coupling sampler." - ) - @staticmethod def _validate_cache_state( transformer_options: dict[object, object], @@ -260,6 +252,8 @@ class RegionalModelPatchInteropValidator: def _reject_attention_collisions( patches: dict[str, list[object]], capabilities: RegionalModelCapabilities, + *, + admitted_negpip: PpmNegpipInterop | None, ) -> None: """Reject every populated attention surface owned by the backend.""" @@ -268,6 +262,11 @@ class RegionalModelPatchInteropValidator: for conflict, patch_name in _ATTN2_PATCH_CONFLICTS.items() if conflict in capabilities.known_patch_conflicts and patches.get(patch_name) + and not ( + patch_name == "attn2_patch" + and admitted_negpip is not None + and patches[patch_name] == [admitted_negpip.attention_patch] + ) ) if conflicts: raise ValueError( diff --git a/simple_syrup/services/anima_attention_coupling_model_family.py b/simple_syrup/services/anima_attention_coupling_model_family.py index bb5656f..07bd6f7 100644 --- a/simple_syrup/services/anima_attention_coupling_model_family.py +++ b/simple_syrup/services/anima_attention_coupling_model_family.py @@ -107,5 +107,6 @@ class AnimaAttentionCouplingModelFamily: adaptation=admission.adaptation, region_strengths=region_strengths, latent_batch_size=latent_batch_size, + negpip=interop_report.negpip, ) return built.model diff --git a/simple_syrup/services/attention_coupling_model_preparation_service.py b/simple_syrup/services/attention_coupling_model_preparation_service.py index be7d5a4..1c1450e 100644 --- a/simple_syrup/services/attention_coupling_model_preparation_service.py +++ b/simple_syrup/services/attention_coupling_model_preparation_service.py @@ -218,6 +218,7 @@ class AttentionCouplingModelPreparationService: noise=samples.to(device), device=device, context_validator=model_family.context_validator, + negpip=interop_report.negpip, ) interop_validator.validate_execution( interop_report, diff --git a/simple_syrup/services/unet_attention_coupling_model_family.py b/simple_syrup/services/unet_attention_coupling_model_family.py index 831e2bb..87e367b 100644 --- a/simple_syrup/services/unet_attention_coupling_model_family.py +++ b/simple_syrup/services/unet_attention_coupling_model_family.py @@ -142,6 +142,7 @@ class StandardUnetAttentionCouplingModelFamily: model=model, state=state, admission=admission, + negpip=interop_report.negpip, ) .model ) diff --git a/tests/test_anima_attention_coupling_model_family.py b/tests/test_anima_attention_coupling_model_family.py index fb21fe1..81d49e5 100644 --- a/tests/test_anima_attention_coupling_model_family.py +++ b/tests/test_anima_attention_coupling_model_family.py @@ -76,6 +76,7 @@ def test_anima_family_retains_single_frame_context_and_backend_policy() -> None: "adaptation": adaptation, "region_strengths": (0.75,), "latent_batch_size": 2, + "negpip": None, } ] diff --git a/tests/test_anima_cross_attention.py b/tests/test_anima_cross_attention.py index a3f6f2b..42eb691 100644 --- a/tests/test_anima_cross_attention.py +++ b/tests/test_anima_cross_attention.py @@ -31,6 +31,10 @@ from simple_syrup.domain.spatial_views import ( SpatialViewKind, ) from simple_syrup.runtime.patcher_lifecycle import PATCHER_LIFECYCLE +from simple_syrup.runtime.ppm_negpip_interop import ( + PpmNegpipInterop, + PpmNegpipSemantics, +) from simple_syrup.runtime.regional_lora.anima_activation_context import ( AnimaActivationContext, AnimaActivationGeometry, @@ -258,6 +262,79 @@ def test_patch_batches_complete_branches_and_blends_expected_outputs( assert invocation_context.current_or_none() is None +def test_patch_packs_anima_negpip_masks_with_the_same_branch_segments() -> None: + """Align each compact regional context with its own NegPiP value mask.""" + + activation_context = AnimaActivationContext() + invocation_context = AnimaCrossAttentionInvocationContext() + original = _DeterministicCrossAttention(invocation_context) + base = torch.zeros((1, 2, 1)) + region = torch.ones_like(base) + base_multiplier = torch.tensor([[[1.0], [-1.0]]]) + region_multiplier = torch.tensor([[[-1.0], [1.0]]]) + contexts = BatchedRegionalAttentionContexts( + latent_batch_size=1, + chunks=( + RegionalAttentionChunkBatch( + 0, + RegionalAttentionBranch.POSITIVE, + 0, + 1, + ), + ), + base_context=base, + regions=( + BatchedRegionalAttentionRegion( + 0, + ( + BatchedRegionalAttentionEntry( + 0, + region, + (1.0,), + region_multiplier, + ), + ), + ), + ), + base_value_multiplier=base_multiplier, + ) + execution = AnimaRegionalAttentionExecution( + contexts, + _bank(torch.full((1, 1, 1), 0.5)), + (1.0,), + ) + negpip = PpmNegpipInterop( + PpmNegpipSemantics.ANIMA_VALUE_MASK, + lambda *args, **kwargs: (args, kwargs), + ) + patch = AnimaRegionalCrossAttentionPatch( + original, + execution, + activation_context=activation_context, + invocation_context=invocation_context, + phase_context=_FullRegionalPhaseContext(), + negpip=negpip, + ) + old_mask = torch.ones_like(base_multiplier) + options: dict[str, object] = {"ppm_negpip_mask": old_mask} + + with activation_context.activate(_geometry(batch=1, height=1, width=1)): + patch( + torch.zeros((1, 1, 1)), + base, + transformer_options=options, + ) + + observed_options = original.calls[0][3] + assert isinstance(observed_options, dict) + assert observed_options is not options + assert torch.equal( + observed_options["ppm_negpip_mask"], + torch.cat((base_multiplier, region_multiplier)), + ) + assert options["ppm_negpip_mask"] is old_mask + + def test_cross_attention_backing_module_does_not_leak_a_host_weight_namespace() -> None: """Keep the retained installed attention outside PyTorch child discovery.""" diff --git a/tests/test_attention_coupling_model_preparation_service.py b/tests/test_attention_coupling_model_preparation_service.py index 3c3cbbc..cc619d0 100644 --- a/tests/test_attention_coupling_model_preparation_service.py +++ b/tests/test_attention_coupling_model_preparation_service.py @@ -25,10 +25,14 @@ from simple_syrup.domain.regional_attention_execution import ( RegionalAttentionExecutionMode, ) from simple_syrup.domain.regional_lora_plan import EMPTY_REGIONAL_LORA_PLAN +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily from simple_syrup.runtime.attention_coupling.family_admission import ( AttentionCouplingFamilyAdmission, ) from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation +from simple_syrup.runtime.regional_model_patch_interop import ( + RegionalModelPatchInteropReport, +) from simple_syrup.services.attention_coupling_model_family import ( AttentionCouplingPreparedModelReuse, AttentionCouplingSamplerConditioning, @@ -59,10 +63,16 @@ class _CapabilityService: class _InteropValidator: """Record centralized modifier admission without requiring a real patcher.""" - report: ClassVar[object] = object() + report: ClassVar[RegionalModelPatchInteropReport] = RegionalModelPatchInteropReport( + RegionalModelFamily.ANIMA, () + ) calls: ClassVar[list[tuple[object, ...]]] = [] - def validate(self, model: object, capabilities: object) -> object: + def validate( + self, + model: object, + capabilities: object, + ) -> RegionalModelPatchInteropReport: """Record exact orchestration inputs without changing them.""" type(self).calls.append((model, capabilities)) diff --git a/tests/test_comfy_patcher_lifecycle.py b/tests/test_comfy_patcher_lifecycle.py index 8e39b52..11abf6b 100644 --- a/tests/test_comfy_patcher_lifecycle.py +++ b/tests/test_comfy_patcher_lifecycle.py @@ -14,7 +14,10 @@ from typing import Any, cast import pytest import torch -from simple_syrup.runtime.patcher_lifecycle import ComfyPatcherLifecycle +from simple_syrup.runtime.patcher_lifecycle import ( + PATCHER_LIFECYCLE, + ComfyPatcherLifecycle, +) from simple_syrup.runtime.regional_lora.execution_cache import ModelCloneLineage @@ -64,6 +67,46 @@ def test_real_comfy_anima_lifecycle_regression( _assert_supported_model_mutations_share_one_clone() +@pytest.mark.parametrize("derivation_count", (2, 3, 5)) +def test_stacked_model_derivations_survive_simultaneous_cyclic_release( + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, + derivation_count: int, +) -> None: + """Keep Comfy on a foreign boundary when any Syrup stack dies together.""" + + import comfy.model_management + from comfy.model_management import LoadedModel + + encoder = AnimaTEModel_() + loader = _patcher(encoder) + foreign_boundary = loader.clone() + derived_models: list[object] = [] + current = foreign_boundary + for stage in range(derivation_count): + current = PATCHER_LIFECYCLE.derive_model( + current, + (), + operation=f"stacked lifecycle regression stage {stage}", + ) + derived_models.append(current) + loaded = LoadedModel(current) + loaded.real_model = weakref.ref(encoder) + monkeypatch.setattr(comfy.model_management, "current_loaded_models", [loaded]) + + execution_cycle: list[object] = [*derived_models] + execution_cycle.append(execution_cycle) + del current, derived_models, execution_cycle + gc.collect() + with caplog.at_level(logging.INFO): + comfy.model_management.cleanup_models_gc() + + assert loaded.model is foreign_boundary + assert loaded.is_dead() is False + assert "Potential memory leak detected" not in caplog.text + assert "WARNING, memory leak" not in caplog.text + + def test_clip_alignment_precedes_mutations_after_dynamic_to_static_clone() -> None: """Mutate the same independently reloaded encoder the returned CLIP executes.""" diff --git a/tests/test_comfy_regional_conditioning_processing.py b/tests/test_comfy_regional_conditioning_processing.py index 1fdd790..b2fba77 100644 --- a/tests/test_comfy_regional_conditioning_processing.py +++ b/tests/test_comfy_regional_conditioning_processing.py @@ -8,7 +8,7 @@ from __future__ import annotations from pathlib import Path from types import SimpleNamespace -from typing import Any +from typing import Any, cast from uuid import UUID import comfy.conds @@ -31,6 +31,10 @@ from simple_syrup.runtime.attention_coupling.unet_context import ( from simple_syrup.runtime.comfy_conditioning_processing import ( COMFY_REGIONAL_CONDITIONING_PROCESSOR, ) +from simple_syrup.runtime.ppm_negpip_interop import ( + PpmNegpipInterop, + PpmNegpipSemantics, +) from simple_syrup.services.attention_coupling_preparation_service import ( ATTENTION_COUPLING_PREPARATION_SERVICE, AttentionCouplingPreparation, @@ -89,6 +93,26 @@ class _LinearModelSampling: return 100.0 * (1.0 - float(percent)) +class _NegpipRecordingAnimaModel(_RecordingAnimaModel): + """Return a distinct PPM-style value mask for every processed context.""" + + def extra_conds(self, **kwargs: Any) -> dict[str, object]: + """Add a binary value mask beside the ordinary cross-attention output.""" + + result = super().extra_conds(**kwargs) + output = self.outputs[-1] + value = int(kwargs["negpip_value"]) + result["c_ppm_negpip_mask"] = comfy.conds.CONDRegular( + torch.full( + (*output.shape[:2], 1), + value, + dtype=torch.int32, + device=output.device, + ) + ) + return result + + def test_processor_uses_comfy_conversion_and_model_post_adapter_contexts() -> None: """Retain weighted padded model outputs in exact positive/negative order.""" @@ -146,6 +170,47 @@ def test_processor_uses_comfy_conversion_and_model_post_adapter_contexts() -> No assert source_positive_context.shape == (1, 3, 1024) +def test_processor_retains_each_anima_negpip_mask_with_its_scheduled_entry() -> None: + """Keep value semantics attached through Comfy conversion and UUID ownership.""" + + model = _NegpipRecordingAnimaModel() + preparation = _preparation( + positive=( + _conditioning(1.0, negpip_value=1), + _conditioning(2.0, negpip_value=-1), + ), + negative=( + _conditioning(-1.0, negpip_value=-1), + _conditioning(-2.0, negpip_value=1), + ), + ) + negpip = PpmNegpipInterop( + PpmNegpipSemantics.ANIMA_VALUE_MASK, + lambda *args, **kwargs: (args, kwargs), + ) + + processed = COMFY_REGIONAL_CONDITIONING_PROCESSOR.process( + preparation, + model=SimpleNamespace(model=model), + noise=torch.zeros((1, 16, 8, 8)), + device=torch.device("cpu"), + context_validator=ANIMA_REGIONAL_CONTEXT_VALIDATOR, + negpip=negpip, + ) + + entries = ( + processed.positive.base_context.entries[0], + processed.positive.regional_contexts[0].entries[0], + processed.negative.base_context.entries[0], + processed.negative.regional_contexts[0].entries[0], + ) + assert all(entry.cross_attention_value_multiplier is not None for entry in entries) + assert [ + int(cast(torch.Tensor, entry.cross_attention_value_multiplier)[0, 0, 0].item()) + for entry in entries + ] == [1, -1, -1, 1] + + @pytest.mark.parametrize( ("sequence_length", "feature_width", "message"), [ @@ -347,6 +412,7 @@ def _conditioning( strength: float | None = None, start_percent: float | None = None, end_percent: float | None = None, + negpip_value: int | None = None, ) -> list[list[object]]: """Build one small standard conditioning with model-consumed metadata.""" @@ -357,6 +423,8 @@ def _conditioning( metadata["start_percent"] = start_percent if end_percent is not None: metadata["end_percent"] = end_percent + if negpip_value is not None: + metadata["negpip_value"] = negpip_value return [ [ torch.full((1, 3, ANIMA_CONTEXT_FEATURE_WIDTH), value), diff --git a/tests/test_model_patcher_mutations.py b/tests/test_model_patcher_mutations.py index 5b621dd..7d9df8d 100644 --- a/tests/test_model_patcher_mutations.py +++ b/tests/test_model_patcher_mutations.py @@ -235,6 +235,35 @@ def test_collision_safe_mutations_integrate_through_one_real_comfy_clone() -> No assert ModelPatcher.set_model_attn2_patch is original_attn2_setter +def test_attn2_mutation_prepends_before_exact_preserved_input_patch() -> None: + """Compose regional packing before an identity-admitted input transformer.""" + + model = _patcher(torch.nn.Linear(1, 1)) + + def preserved(*args: object) -> tuple[object, ...]: + """Return preserved callback arguments.""" + + return args + + def regional(*args: object) -> tuple[object, ...]: + """Return regional callback arguments.""" + + return args + + def output(*args: object) -> tuple[object, ...]: + """Return output callback arguments.""" + + return args + + model.set_model_attn2_patch(preserved) + + ModelAttn2PatchesMutation(regional, output, (preserved,)).apply(model) + + patches = model.model_options["transformer_options"]["patches"] + assert patches["attn2_patch"] == [regional, preserved] + assert patches["attn2_output_patch"] == [output] + + @pytest.mark.parametrize( ("wrapper_type", "key", "wrapper", "message"), [ diff --git a/tests/test_optional_triton_runtime.py b/tests/test_optional_triton_runtime.py new file mode 100644 index 0000000..fd51d72 --- /dev/null +++ b/tests/test_optional_triton_runtime.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 + +"""Prove Triton remains an optional CUDA acceleration dependency.""" + +from __future__ import annotations + +import subprocess +import sys +from pathlib import Path + +import pytest + +from simple_syrup.runtime.regional_lora.triton_runtime import ( + TritonRuntimeResolver, +) + + +def test_node_registration_and_cpu_accumulation_do_not_import_triton() -> None: + """Load public nodes and execute CPU accumulation with Triton blocked.""" + + script = """ +import importlib.abc +import sys +from pathlib import Path + +sys.path.insert(0, str(Path.cwd().parents[1])) +sys.argv = [sys.argv[0], "--cpu"] +import comfy.options +comfy.options.enable_args_parsing() + +class BlockTriton(importlib.abc.MetaPathFinder): + def find_spec(self, fullname, path, target=None): + if fullname == "triton" or fullname.startswith("triton."): + raise ModuleNotFoundError("blocked optional Triton", name=fullname) + return None + +sys.meta_path.insert(0, BlockTriton()) +import torch +from simple_syrup.nodes_v3 import get_nodes +from simple_syrup.runtime.regional_lora import ordered_accumulation + +assert get_nodes() +base = torch.tensor([1.0, 2.0]) +result = ordered_accumulation.OrderedTensorAccumulator().accumulate( + base, + (torch.tensor([3.0, 4.0]), torch.tensor([5.0, 6.0])), +) +assert result.tolist() == [9.0, 12.0] +assert not any(name == "triton" or name.startswith("triton.") for name in sys.modules) +""" + completed = subprocess.run( + [sys.executable, "-c", script], + cwd=Path(__file__).resolve().parents[1], + capture_output=True, + text=True, + timeout=60, + check=False, + ) + + assert completed.returncode == 0, completed.stderr + + +def test_resolver_caches_one_missing_result_and_warns_once( + caplog: pytest.LogCaptureFixture, +) -> None: + """Treat an absent Triton package as one observable optional miss.""" + + calls: list[str] = [] + + def missing_import(name: str) -> object: + calls.append(name) + raise ModuleNotFoundError("missing", name="triton") + + resolver = TritonRuntimeResolver(import_module=missing_import) + + with caplog.at_level("WARNING"): + assert resolver.resolve("fake.backend") is None + assert resolver.resolve("fake.backend") is None + + assert calls == ["fake.backend"] + assert [record.message for record in caplog.records] == [ + "Triton acceleration is unavailable; using the Torch execution path" + ] + + +def test_resolver_exposes_broken_backend_import_with_original_cause() -> None: + """Fail visibly when a present backend cannot initialize correctly.""" + + failure = RuntimeError("JIT initialization failed") + + def broken_import(_name: str) -> object: + raise failure + + resolver = TritonRuntimeResolver(import_module=broken_import) + + with pytest.raises(RuntimeError, match="failed to initialize") as raised: + resolver.resolve("fake.backend") + + assert raised.value.__cause__ is failure + + +def test_resolver_is_thread_safe_and_returns_one_cached_backend() -> None: + """Publish exactly one imported backend across concurrent callers.""" + + from concurrent.futures import ThreadPoolExecutor + + backend = object() + calls: list[str] = [] + + def import_backend(name: str) -> object: + calls.append(name) + return backend + + resolver = TritonRuntimeResolver(import_module=import_backend) + + with ThreadPoolExecutor(max_workers=8) as executor: + results = tuple(executor.map(resolver.resolve, ("fake.backend",) * 32)) + + assert all(result is backend for result in results) + assert calls == ["fake.backend"] + + +@pytest.mark.parametrize( + "missing_name", + ["fake.backend", "unrelated_dependency"], +) +def test_resolver_does_not_hide_non_triton_module_failures(missing_name: str) -> None: + """Reserve optional fallback exclusively for the Triton package family.""" + + def missing_import(_name: str) -> object: + raise ModuleNotFoundError("missing", name=missing_name) + + resolver = TritonRuntimeResolver( + import_module=missing_import, + ) + + with pytest.raises(RuntimeError, match="failed to initialize"): + resolver.resolve("fake.backend") diff --git a/tests/test_ppm_negpip_interop.py b/tests/test_ppm_negpip_interop.py new file mode 100644 index 0000000..698c693 --- /dev/null +++ b/tests/test_ppm_negpip_interop.py @@ -0,0 +1,94 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Verify typed PPM NegPiP conditioning and call-local option adaptation.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import comfy.conds +import pytest +import torch + +from simple_syrup.runtime.ppm_negpip_interop import ( + PpmNegpipInterop, + PpmNegpipSemantics, +) + + +def test_anima_adapter_extracts_exact_typed_value_multiplier() -> None: + """Convert PPM's model condition to context-aligned execution state.""" + + interop = _anima_interop() + context = torch.zeros((1, 3, 4), dtype=torch.float16) + source = torch.tensor([[[1], [-1], [1]]], dtype=torch.int32) + + multiplier = interop.extract_value_multiplier( + {"c_ppm_negpip_mask": comfy.conds.CONDRegular(source)}, + context, + ) + + assert multiplier is not None + assert multiplier.dtype is context.dtype + assert multiplier.device == context.device + assert multiplier.tolist() == [[[1.0], [-1.0], [1.0]]] + + +def test_anima_adapter_uses_neutral_multiplier_when_condition_is_absent() -> None: + """Represent an all-positive prompt without leaving branch state partial.""" + + context = torch.zeros((2, 3, 4)) + + multiplier = _anima_interop().extract_value_multiplier({}, context) + + assert multiplier is not None + assert torch.equal(multiplier, torch.ones((2, 3, 1))) + + +@pytest.mark.parametrize( + "source", + [ + torch.ones((1, 2, 1)), + torch.zeros((1, 3, 1)), + torch.ones((1, 3, 2)), + ], +) +def test_anima_adapter_rejects_misaligned_or_nonbinary_masks( + source: torch.Tensor, +) -> None: + """Fail closed before malformed PPM state reaches regional packing.""" + + with pytest.raises(ValueError): + _anima_interop().extract_value_multiplier( + {"c_ppm_negpip_mask": SimpleNamespace(cond=source)}, + torch.zeros((1, 3, 4)), + ) + + +def test_anima_adapter_publishes_mask_on_an_isolated_option_copy() -> None: + """Keep the ordinary call options unchanged outside original cross-attention.""" + + old_mask = torch.ones((1, 2, 1)) + packed = torch.tensor([[[1.0], [-1.0]], [[-1.0], [1.0]]]) + source: dict[str, object] = { + "ppm_negpip_mask": old_mask, + "preserved": object(), + } + + prepared = _anima_interop().prepare_anima_transformer_options(source, packed) + + assert prepared is not source + assert prepared["preserved"] is source["preserved"] + assert prepared["ppm_negpip_mask"] is packed + assert source["ppm_negpip_mask"] is old_mask + + +def _anima_interop() -> PpmNegpipInterop: + """Return one focused admitted Anima semantic adapter.""" + + return PpmNegpipInterop( + PpmNegpipSemantics.ANIMA_VALUE_MASK, + lambda *args, **kwargs: (args, kwargs), + ) diff --git a/tests/test_preparation_collaborator_profile.py b/tests/test_preparation_collaborator_profile.py index 9f9ad3f..936f918 100644 --- a/tests/test_preparation_collaborator_profile.py +++ b/tests/test_preparation_collaborator_profile.py @@ -26,6 +26,7 @@ from simple_syrup.runtime.comfy_conditioning_model_loader import ( from simple_syrup.runtime.comfy_conditioning_processing import ( ComfyRegionalConditioningProcessor, ) +from simple_syrup.runtime.ppm_negpip_interop import PpmNegpipInterop from simple_syrup.runtime.regional_lora_conditioning_adapter import ( RegionalLoraConditioningAdapter, ) @@ -135,6 +136,7 @@ def test_profiled_collaborators_preserve_arguments_results_and_stage_order( noise: torch.Tensor, device: torch.device, context_validator: RegionalContextValidator, + negpip: PpmNegpipInterop | None = None, ) -> ProcessedRegionalAttentionPlan: calls.append( ( @@ -145,6 +147,7 @@ def test_profiled_collaborators_preserve_arguments_results_and_stage_order( "noise": noise, "device": device, "context_validator": context_validator, + "negpip": negpip, }, ) ) @@ -190,6 +193,7 @@ def test_profiled_collaborators_preserve_arguments_results_and_stage_order( "noise": noise, "device": device, "context_validator": validator, + "negpip": None, } assert _stages(caplog) == [ "source_model_load", diff --git a/tests/test_regional_attention_batching.py b/tests/test_regional_attention_batching.py index 240ced2..535c749 100644 --- a/tests/test_regional_attention_batching.py +++ b/tests/test_regional_attention_batching.py @@ -108,6 +108,56 @@ def test_batching_builds_canonical_regions_with_base_fallback() -> None: ] +def test_batching_aligns_value_multipliers_with_cfg_regions_and_fallbacks() -> None: + """Keep each scheduled branch's value semantics in identical chunk order.""" + + source = _plan() + plan = ProcessedRegionalAttentionPlan( + ProcessedRegionalAttentionBranch( + _context_with_multiplier(source.positive.base_context, 1.0), + ( + _context_with_multiplier( + source.positive.regional_contexts[0], + -1.0, + ), + _context_with_multiplier( + source.positive.regional_contexts[1], + 1.0, + ), + ), + ), + ProcessedRegionalAttentionBranch( + _context_with_multiplier(source.negative.base_context, -1.0), + ( + _context_with_multiplier( + source.negative.regional_contexts[0], + 1.0, + ), + ), + ), + source.mask_bank, + source.lora_plan, + ) + + aligned = REGIONAL_ATTENTION_BATCHING_SERVICE.align( + plan, + base_context=_runtime_base(plan, [1, 0], 2), + cond_or_uncond=[1, 0], + conditioning_uuids=_uuids_for_selectors(plan, [1, 0]), + sigma=0.5, + latent_batch_size=2, + ) + + assert aligned.base_value_multiplier is not None + assert aligned.base_value_multiplier[:, 0, 0].tolist() == [-1.0, -1.0, 1.0, 1.0] + region_zero = aligned.regions[0].entries[0].cross_attention_value_multiplier + region_one = aligned.regions[1].entries[0].cross_attention_value_multiplier + assert region_zero is not None + assert region_one is not None + assert region_zero[:, 0, 0].tolist() == [1.0, 1.0, -1.0, -1.0] + assert region_one[:, 0, 0].tolist() == [-1.0, -1.0, 1.0, 1.0] + + def test_batching_preserves_all_regional_entries_and_per_sample_strengths() -> None: """Align simultaneous regional entries without collapsing their order.""" @@ -409,6 +459,34 @@ def _multi_entry_context( ) +def _context_with_multiplier( + context: ProcessedRegionalAttentionContext, + value: float, +) -> ProcessedRegionalAttentionContext: + """Copy one context with a uniform sequence-aligned value multiplier.""" + + return ProcessedRegionalAttentionContext( + context.conditioning_index, + context.region_index, + tuple( + ProcessedRegionalAttentionEntry( + entry.entry_index, + entry.uuid, + entry.schedule, + entry.cross_attention, + entry.strength, + torch.full( + (*entry.cross_attention.shape[:2], 1), + value, + dtype=entry.cross_attention.dtype, + device=entry.cross_attention.device, + ), + ) + for entry in context.entries + ), + ) + + def _entry( entry_index: int, tensor: torch.Tensor, diff --git a/tests/test_regional_model_patch_interop.py b/tests/test_regional_model_patch_interop.py index ae7ad63..a20c34b 100644 --- a/tests/test_regional_model_patch_interop.py +++ b/tests/test_regional_model_patch_interop.py @@ -30,6 +30,7 @@ from simple_syrup.domain.regional_model_capabilities import ( RegionalReferenceLatentPolicy, RegionalSpatialPatchSupport, ) +from simple_syrup.runtime.ppm_negpip_interop import PpmNegpipSemantics from simple_syrup.runtime.regional_model_patch_interop import ( REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR, RegionalPreservedModelModifier, @@ -46,6 +47,11 @@ class _FixtureModel(torch.nn.Module): self.projection = torch.nn.Linear(1, 1) self.latent_format = SimpleNamespace(latent_channels=4) + def extra_conds(self, **_kwargs: object) -> dict[str, object]: + """Expose the object path patched by Anima NegPiP.""" + + return {} + def test_validator_preserves_easycache_and_unrelated_model_state() -> None: """Accept EasyCache while retaining every collaborator-owned surface.""" @@ -186,28 +192,103 @@ def test_validator_rejects_both_core_caches_without_mutating_them() -> None: assert combined.wrappers == before_wrappers -@pytest.mark.parametrize( - "family", - [RegionalModelFamily.ANIMA, RegionalModelFamily.STANDARD_UNET], -) -def test_validator_rejects_named_negpip_before_generic_attn2_collision( - family: RegionalModelFamily, -) -> None: - """Report the installed modifier and regional mask misalignment by name.""" +def test_validator_admits_exact_standard_unet_negpip_without_mutation() -> None: + """Retain PPM's exact split-K/V callback as typed interop evidence.""" model = _patcher() + callback = _identity_callback( + "custom_nodes.ComfyUI-ppm.src.negpip.unet_negpip", + "sdxl_attn2_negpip", + ) model.model_options["ppm_negpip"] = True + model.set_model_attn2_patch(callback) + + report = REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR.validate( + model, + _capabilities(RegionalModelFamily.STANDARD_UNET), + ) + + assert report.negpip is not None + assert report.negpip.semantics is PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE + assert report.negpip.attention_patch is callback + assert model.model_options["transformer_options"]["patches"]["attn2_patch"] == [ + callback + ] + + +def test_validator_admits_exact_anima_negpip_without_mutation() -> None: + """Retain PPM's complete Anima callback, wrapper, and object-patch family.""" + + model = _patcher() + callback = _identity_callback( + "custom_nodes.ComfyUI-ppm.src.negpip.anima_negpip", + "cosmos_attn2_negpip", + ) + wrapper = _identity_callback( + "custom_nodes.ComfyUI-ppm.src.negpip.anima_negpip", + "cosmos_diffusion_negpip_wrapper", + ) + extra_conds = _identity_callback( + "custom_nodes.ComfyUI-ppm.src.negpip.anima_negpip", + ("anima_extra_conds_negpip_wrapper.._anima_extra_conds_negpip_wrapper"), + ) + model.model_options["ppm_negpip"] = True + model.set_model_attn2_patch(callback) model.add_wrapper_with_key( WrappersMP.DIFFUSION_MODEL, "ppm_negpip_anima", - lambda executor, *args, **kwargs: executor(*args, **kwargs), + wrapper, ) - model.set_model_attn2_patch(lambda q, k, v, **kwargs: {"q": q, "k": k, "v": v}) + model.add_object_patch("extra_conds", extra_conds) - with pytest.raises( - ValueError, - match="NegPiP.*ordinary conditioning batch.*regional branch batch", - ): + report = REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR.validate( + model, + _capabilities(RegionalModelFamily.ANIMA), + ) + + assert report.negpip is not None + assert report.negpip.semantics is PpmNegpipSemantics.ANIMA_VALUE_MASK + assert report.negpip.attention_patch is callback + assert model.wrappers[WrappersMP.DIFFUSION_MODEL]["ppm_negpip_anima"] == [wrapper] + assert model.object_patches["extra_conds"] is extra_conds + + +@pytest.mark.parametrize( + ("family", "configure", "message"), + [ + ( + RegionalModelFamily.STANDARD_UNET, + lambda model: model.model_options.__setitem__("ppm_negpip", True), + "requires exactly its PPM split-K/V", + ), + ( + RegionalModelFamily.STANDARD_UNET, + lambda model: model.set_model_attn2_patch( + _identity_callback( + "custom_nodes.ComfyUI-ppm.src.negpip.unet_negpip", + "sdxl_attn2_negpip", + ) + ), + "incomplete NegPiP patch family", + ), + ( + RegionalModelFamily.ANIMA, + lambda model: model.model_options.__setitem__("ppm_negpip", True), + "requires exactly its PPM attention patch", + ), + ], +) +def test_validator_rejects_partial_or_foreign_negpip_families( + family: RegionalModelFamily, + configure: Callable[[ModelPatcher], object], + message: str, +) -> None: + """Fail closed before partial or identity-foreign NegPiP state is composed.""" + + model = _patcher() + configure(model) + + with pytest.raises(ValueError, match=message): REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR.validate( model, _capabilities(family), @@ -313,6 +394,19 @@ def _patcher() -> ModelPatcher: ) +def _identity_callback(module: str, qualname: str) -> Callable[..., object]: + """Build one executable callback carrying a stable PPM definition identity.""" + + def callback(*args: object, **_kwargs: object) -> object: + """Return callback inputs for model-state-only admission tests.""" + + return args + + callback.__module__ = module + callback.__qualname__ = qualname + return callback + + def _capabilities(family: RegionalModelFamily) -> RegionalModelCapabilities: """Build the exact family contract consumed by modifier admission.""" diff --git a/tests/test_regional_model_patch_stack.py b/tests/test_regional_model_patch_stack.py index 2619c47..37e426b 100644 --- a/tests/test_regional_model_patch_stack.py +++ b/tests/test_regional_model_patch_stack.py @@ -100,7 +100,7 @@ def test_regional_patch_stack_preserves_state_lineage_and_runtime_nesting() -> N assert stack.user_model is source assert stack.attention_model.parent is source - assert stack.sampling_model.parent is stack.attention_model + assert stack.sampling_model.parent is source assert stack.sampling_model.patches["weight"] == [global_lora_patch] assert ( source.get_wrappers("diffusion_model", "simple_syrup.attention_coupling") == [] diff --git a/tests/test_regional_patch_interop_evidence.py b/tests/test_regional_patch_interop_evidence.py index 0f5a47f..280f2cf 100644 --- a/tests/test_regional_patch_interop_evidence.py +++ b/tests/test_regional_patch_interop_evidence.py @@ -41,6 +41,9 @@ def test_every_matrix_case_requires_exact_modifier_and_terminal_evidence( observed: RegionalPatchInteropHistory if case.expect_success: model_call_count = STEPS + diagnostic_record_count = model_call_count * ( + 2 if case.model_family is PatchInteropModelFamily.SDXL else 1 + ) observed = RegionalPatchInteropSuccess( snapshot, { @@ -49,7 +52,7 @@ def test_every_matrix_case_requires_exact_modifier_and_terminal_evidence( "runtime_ms": 10.0, "peak_vram_bytes": 1, }, - _diagnostics(case, workflow, record_count=model_call_count), + _diagnostics(case, workflow, record_count=diagnostic_record_count), ImageReference("result.png", "", "output"), None, ) @@ -69,7 +72,7 @@ def test_every_matrix_case_requires_exact_modifier_and_terminal_evidence( assert validated.model_call_count == (STEPS if case.expect_success else 0) -def test_success_requires_one_diagnostic_record_per_actual_model_call() -> None: +def test_success_requires_family_specific_diagnostic_records_per_model_call() -> None: """Reject cache evidence whose diagnostics omit an executed model call.""" case = next( @@ -89,7 +92,7 @@ def test_success_requires_one_diagnostic_record_per_actual_model_call() -> None: None, ) - with pytest.raises(ValueError, match="one diagnostic record per model call"): + with pytest.raises(ValueError, match="model-family execution shape"): validate_case(case, workflow, observed) @@ -216,7 +219,8 @@ def _diagnostics( PatchInteropSpatialMode.CONTEXTUAL: ("tile", "contextual_global"), }[case.spatial_mode] snapshots = [ - _diagnostic_snapshot(modes[index % len(modes)]) for index in range(record_count) + _diagnostic_snapshot(case, modes[index % len(modes)]) + for index in range(record_count) ] return { "run_id": workflow.diagnostics_run_id, @@ -225,8 +229,25 @@ def _diagnostics( } -def _diagnostic_snapshot(mode: str) -> JsonObject: - """Return one exact static-PRIMARY_ADAPTER diagnostic record.""" +def _diagnostic_snapshot( + case: RegionalPatchInteropCase, + mode: str, +) -> JsonObject: + """Return one exact family-specific regional diagnostic record.""" + + if case.model_family is PatchInteropModelFamily.SDXL: + return { + "strategy": "attention_coupling", + "backend": "comfy.ldm.modules.diffusionmodules.openaimodel.UNetModel", + "spatial_mode": mode, + "region_count": 2, + "active_region_indices": [0, 1], + "estimated_work": { + "cross_attention_branch_multiplier": 3.0, + "cross_attention_formula": "base_plus_region_count", + "denoiser_call_multiplier": 1.0, + }, + } return { "strategy": "attention_coupling", diff --git a/tests/test_regional_patch_interop_workflow.py b/tests/test_regional_patch_interop_workflow.py index 65d0676..575f677 100644 --- a/tests/test_regional_patch_interop_workflow.py +++ b/tests/test_regional_patch_interop_workflow.py @@ -23,13 +23,13 @@ from tools.regional_patch_interop_integration.workflow import ( ) -def test_matrix_contains_five_acceptances_and_six_exact_rejections() -> None: +def test_matrix_contains_seven_acceptances_and_four_exact_rejections() -> None: """Keep every required modifier, spatial, scheduled, and family case.""" definitions = cases() assert len(definitions) == 11 - assert sum(case.expect_success for case in definitions) == 5 + assert sum(case.expect_success for case in definitions) == 7 assert {case.modifier for case in definitions} == set(PatchInteropModifier) assert {case.spatial_mode for case in definitions} == set(PatchInteropSpatialMode) assert {case.model_family for case in definitions} == set(PatchInteropModelFamily) @@ -154,11 +154,9 @@ def test_scheduled_cache_graph_authors_the_exact_regional_adapter_interval() -> def test_sdxl_negpip_graph_uses_public_modifier_snapshot_and_sampler() -> None: - """Submit SDXL NegPiP through the same evidence and rejection boundary.""" + """Submit SDXL NegPiP through the same evidence and execution boundary.""" - definition = next( - case for case in cases() if case.case_id == "sdxl-negpip-rejected" - ) + definition = next(case for case in cases() if case.case_id == "sdxl-negpip") workflow = RegionalPatchInteropWorkflowBuilder().build( definition, run_id="run", diff --git a/tests/test_standard_unet_variant_base_attention.py b/tests/test_standard_unet_variant_base_attention.py index 731cb13..b042be7 100644 --- a/tests/test_standard_unet_variant_base_attention.py +++ b/tests/test_standard_unet_variant_base_attention.py @@ -14,6 +14,10 @@ import torch from simple_syrup.runtime.attention_coupling.unet_attn2_execution import ( UnetAttn2Execution, ) +from simple_syrup.runtime.ppm_negpip_interop import ( + PpmNegpipInterop, + PpmNegpipSemantics, +) from simple_syrup.runtime.regional_lora.standard_unet_variant_base_attention import ( StandardUnetVariantBaseAttention, ) @@ -67,3 +71,32 @@ def test_prepare_rejects_any_preexisting_attn2_callback_surface(key: str) -> Non with pytest.raises(ValueError, match="already contains"): StandardUnetVariantBaseAttention(_Resolver()).prepare({"patches": {key: []}}) + + +def test_prepare_places_coupling_before_exact_preserved_negpip_callback() -> None: + """Pack regional alternating tokens before PPM selects K and V views.""" + + def negpip(*args: object, **_kwargs: object) -> tuple[object, ...]: + """Represent the identity-validated PPM split callback.""" + + return args + + interop = PpmNegpipInterop( + PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE, + negpip, + ) + source: dict[str, object] = {"patches": {"attn2_patch": [negpip]}} + + prepared = StandardUnetVariantBaseAttention( + _Resolver(), + negpip=interop, + ).prepare(source) + + patches = prepared["patches"] + assert isinstance(patches, dict) + installed = patches["attn2_patch"] + assert isinstance(installed, list) + assert len(installed) == 2 + assert installed[1] is negpip + assert callable(installed[0]) + assert source == {"patches": {"attn2_patch": [negpip]}} diff --git a/tests/test_unet_attn2_patch.py b/tests/test_unet_attn2_patch.py index c9c9ac4..5126f0a 100644 --- a/tests/test_unet_attn2_patch.py +++ b/tests/test_unet_attn2_patch.py @@ -29,6 +29,13 @@ from simple_syrup.runtime.attention_coupling.unet_attn2_execution_resolver impor from simple_syrup.runtime.attention_coupling.unet_attn2_patch import ( UnetAttn2PatchPair, ) +from simple_syrup.runtime.ppm_negpip_interop import ( + PpmNegpipInterop, + PpmNegpipSemantics, +) +from simple_syrup.runtime.regional_lora.standard_unet_variant_base_attention import ( + StandardUnetVariantBaseAttention, +) class _ZeroAttention(nn.Module): @@ -81,6 +88,32 @@ class _RegionalAttention(nn.Module): return query + context.mean(dim=1, keepdim=True) +class _RecordingKeyValueAttention(nn.Module): + """Record exact post-patch key/value inputs and return packed zeros.""" + + def __init__(self) -> None: + """Initialize an empty invocation record.""" + + super().__init__() + self.calls: list[tuple[torch.Tensor, torch.Tensor]] = [] + + def forward( + self, + query: torch.Tensor, + *, + context: torch.Tensor | None, + value: torch.Tensor | None, + transformer_options: dict[str, Any], + ) -> torch.Tensor: + """Retain post-patch K/V sources without performing attention.""" + + del transformer_options + if context is None or value is None: + raise AssertionError("NegPiP test requires explicit K and V tensors.") + self.calls.append((context, value)) + return torch.zeros_like(query) + + class _CountingZeroFeedForward(nn.Module): """Return zero while counting the retained feed-forward trajectory.""" @@ -165,6 +198,149 @@ def test_unet_patch_clears_callback_state_after_output_failure() -> None: patches.input_patch(query, context, context, options) +@pytest.mark.parametrize("persistent_base_graph", [False, True]) +def test_negpip_splits_exact_packed_regional_key_and_value_views( + persistent_base_graph: bool, +) -> None: + """Select even K and odd V tokens after packing in both UNet base routes.""" + + contexts = _negpip_contexts() + execution = UnetAttn2Execution( + contexts, + torch.tensor([[[1.0, 0.0]], [[0.0, 1.0]]]), + (1.0, 1.0), + 1, + 2, + ) + pair = UnetAttn2PatchPair(StaticUnetAttn2ExecutionResolver(execution)) + + def split_negpip( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + _extra_options: dict[str, Any], + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Apply PPM's public alternating-token UNet semantics exactly.""" + + return query, key[:, 0::2], value[:, 1::2] + + if persistent_base_graph: + interop = PpmNegpipInterop( + PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE, + split_negpip, + ) + options = StandardUnetVariantBaseAttention( + StaticUnetAttn2ExecutionResolver(execution), + negpip=interop, + ).prepare({"patches": {"attn2_patch": [split_negpip]}}) + else: + options = { + "patches": { + "attn2_patch": [pair.input_patch, split_negpip], + "attn2_output_patch": [pair.output_patch], + } + } + recording = _RecordingKeyValueAttention() + block = _block(recording) + + block( + torch.zeros((1, 2, 1)), + context=contexts.base_context, + transformer_options=options, + ) + + assert len(recording.calls) == 1 + key, value = recording.calls[0] + assert key[:, :, 0].tolist() == [[30.0, 40.0], [70.0, 80.0]] + assert value[:, :, 0].tolist() == [[31.0, 41.0], [71.0, 81.0]] + + +def test_persistent_regional_graph_keeps_native_negpip_split_semantics() -> None: + """Split one manually selected regional graph without adding branch packing.""" + + recording = _RecordingKeyValueAttention() + block = _block(recording) + context = torch.tensor([[[30.0], [31.0], [40.0], [41.0]]]) + + def split_negpip( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + _extra_options: dict[str, Any], + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Apply PPM's public alternating-token UNet semantics exactly.""" + + return query, key[:, 0::2], value[:, 1::2] + + block( + torch.zeros((1, 2, 1)), + context=context, + transformer_options={"patches": {"attn2_patch": [split_negpip]}}, + ) + + key, value = recording.calls[0] + assert key[:, :, 0].tolist() == [[30.0, 40.0]] + assert value[:, :, 0].tolist() == [[31.0, 41.0]] + + +def _block(cross_attention: nn.Module) -> BasicTransformerBlock: + """Build one deterministic block around a supplied cross-attention owner.""" + + block = BasicTransformerBlock( + dim=1, + n_heads=1, + d_head=1, + context_dim=1, + checkpoint=False, + ) + block.norm1 = nn.Identity() + block.attn1 = _ZeroAttention() + block.norm2 = nn.Identity() + block.attn2 = cross_attention + block.norm3 = nn.Identity() + block.ff = _CountingZeroFeedForward() + return block + + +def _negpip_contexts() -> BatchedRegionalAttentionContexts: + """Return alternating K/V token pairs for base and two regions.""" + + return BatchedRegionalAttentionContexts( + latent_batch_size=1, + chunks=( + RegionalAttentionChunkBatch( + 0, + RegionalAttentionBranch.POSITIVE, + 0, + 1, + ), + ), + base_context=torch.tensor([[[10.0], [11.0], [20.0], [21.0]]]), + regions=( + BatchedRegionalAttentionRegion( + 0, + ( + BatchedRegionalAttentionEntry( + 0, + torch.tensor([[[30.0], [31.0], [40.0], [41.0]]]), + (1.0,), + ), + ), + ), + BatchedRegionalAttentionRegion( + 1, + ( + BatchedRegionalAttentionEntry( + 0, + torch.tensor([[[70.0], [71.0], [80.0], [81.0]]]), + (1.0,), + ), + ), + ), + ), + ) + + def _contexts() -> BatchedRegionalAttentionContexts: """Return one base and two single-entry regional contexts.""" diff --git a/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py b/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py index 6d859f7..f8d73d9 100644 --- a/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py +++ b/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py @@ -23,6 +23,7 @@ from simple_syrup.runtime.comfy_conditioning_model_loader import ( from simple_syrup.runtime.comfy_conditioning_processing import ( ComfyRegionalConditioningProcessor, ) +from simple_syrup.runtime.ppm_negpip_interop import PpmNegpipInterop from simple_syrup.runtime.regional_lora_conditioning_adapter import ( RegionalLoraConditioningAdapter, ) @@ -101,6 +102,7 @@ class ProfiledComfyRegionalConditioningProcessor(ComfyRegionalConditioningProces noise: torch.Tensor, device: torch.device, context_validator: RegionalContextValidator, + negpip: PpmNegpipInterop | None = None, ) -> ProcessedRegionalAttentionPlan: """Delegate conditioning processing with synchronized device timing.""" @@ -114,6 +116,7 @@ class ProfiledComfyRegionalConditioningProcessor(ComfyRegionalConditioningProces noise=noise, device=device, context_validator=context_validator, + negpip=negpip, ) diff --git a/tools/regional_patch_interop_integration/matrix.py b/tools/regional_patch_interop_integration/matrix.py index dd569a4..3b8b918 100644 --- a/tools/regional_patch_interop_integration/matrix.py +++ b/tools/regional_patch_interop_integration/matrix.py @@ -109,11 +109,6 @@ class RegionalPatchInteropCase: def cases() -> tuple[RegionalPatchInteropCase, ...]: """Return accepted and rejected cases in authoritative evidence order.""" - negpip_error = ( - "does not support NegPiP", - "ordinary conditioning batch", - "regional branch batch", - ) return ( _accepted("anima-full-baseline", "Anima full static PRIMARY_ADAPTER baseline"), _accepted( @@ -183,24 +178,22 @@ def cases() -> tuple[RegionalPatchInteropCase, ...]: ("easycache", "Contextual spatial views", "view coordinates"), ), RegionalPatchInteropCase( - "anima-negpip-rejected", - "Reject Anima NegPiP before regional branch packing", + "anima-negpip", + "Anima NegPiP with aligned regional value masks", PatchInteropModelFamily.ANIMA, PatchInteropSpatialMode.FULL, PatchInteropModifier.NEGPIP, - PatchInteropOutcome.REJECTED, + PatchInteropOutcome.ACCEPTED, False, - negpip_error, ), RegionalPatchInteropCase( - "sdxl-negpip-rejected", - "Reject SDXL NegPiP before paired regional attention patches", + "sdxl-negpip", + "SDXL NegPiP with packed regional split-K/V conditioning", PatchInteropModelFamily.SDXL, PatchInteropSpatialMode.FULL, PatchInteropModifier.NEGPIP, - PatchInteropOutcome.REJECTED, + PatchInteropOutcome.ACCEPTED, False, - negpip_error, ), ) diff --git a/tools/regional_patch_interop_integration/validation.py b/tools/regional_patch_interop_integration/validation.py index 937303b..8a4cbfb 100644 --- a/tools/regional_patch_interop_integration/validation.py +++ b/tools/regional_patch_interop_integration/validation.py @@ -133,7 +133,7 @@ def _validate_success( workflow: BuiltRegionalPatchInteropWorkflow, observed: RegionalPatchInteropSuccess, ) -> ValidatedRegionalPatchInterop: - """Require one exact single-trajectory regional PRIMARY_ADAPTER execution.""" + """Require one exact single-trajectory regional execution.""" metrics = observed.metrics if metrics.get("run_id") != workflow.metrics_run_id: @@ -150,62 +150,26 @@ def _validate_success( raise ValueError("P9.7 diagnostics identity changed.") snapshots = _array(diagnostics.get("snapshots"), "diagnostic snapshots") record_count = _integer(diagnostics.get("record_count"), "record count") - if record_count != model_calls or len(snapshots) != record_count: - raise ValueError("P9.7 requires one diagnostic record per model call.") + expected_records = model_calls * _diagnostic_records_per_model_call(case) + if record_count != expected_records or len(snapshots) != record_count: + raise ValueError( + "P9.7 diagnostic records do not match the model-family execution shape." + ) if record_count < 1: raise ValueError("P9.7 diagnostics must contain aligned model-call records.") spatial_modes: set[str] = set() adapter_tokens: set[str] = set() for item in snapshots: snapshot = _object(item, "diagnostic snapshot") - if ( - snapshot.get("strategy") != "attention_coupling" - or snapshot.get("backend") != "comfy.ldm.anima.model.Anima" - ): + if snapshot.get("strategy") != "attention_coupling" or snapshot.get( + "backend" + ) != _expected_backend(case): raise ValueError("P9.7 diagnostic strategy or backend changed.") spatial_modes.add(_string(snapshot.get("spatial_mode"), "spatial mode")) - uses = _array(snapshot.get("adapter_uses"), "adapter uses") - if len(uses) != 2: - raise ValueError( - "P9.7 accepted cases require paired positive/negative adapter uses." - ) - normalized_uses = tuple(_object(use, "adapter use") for use in uses) - if { - (_string(use.get("branch"), "adapter branch"), use.get("composition_index")) - for use in normalized_uses - } != {("positive", 0), ("negative", 1)}: - raise ValueError( - "P9.7 paired regional PRIMARY_ADAPTER branch ownership changed." - ) - for use in normalized_uses: - if ( - use.get("active") is not True - or use.get("region_index") != 0 - or use.get("target_count") != 448 - ): - raise ValueError( - "P9.7 exact regional PRIMARY_ADAPTER execution changed." - ) - if not math.isclose( - _number(use.get("effective_strength"), "effective strength"), - 0.75, - abs_tol=1e-8, - ): - raise ValueError("P9.7 regional PRIMARY_ADAPTER strength changed.") - adapter_tokens.add(_string(use.get("adapter_token"), "adapter token")) - work = _object(snapshot.get("estimated_work"), "estimated work") - if ( - work.get("active_adapter_uses") != 2 - or work.get("active_target_count") != 448 - or work.get("target_use_count") != 896 - ): - raise ValueError("P9.7 paired LoRA target-use accounting changed.") - if not math.isclose( - _number(work.get("denoiser_call_multiplier"), "denoiser multiplier"), - 1.0, - abs_tol=1e-8, - ): - raise ValueError("P9.7 denoiser trajectory multiplier changed.") + if case.model_family is PatchInteropModelFamily.ANIMA: + _validate_anima_snapshot(snapshot, adapter_tokens) + else: + _validate_sdxl_snapshot(snapshot) expected_modes = { PatchInteropSpatialMode.FULL: {"full"}, PatchInteropSpatialMode.TILED: {"tile"}, @@ -213,7 +177,7 @@ def _validate_success( }[case.spatial_mode] if not expected_modes <= spatial_modes: raise ValueError("P9.7 spatial diagnostics are incomplete.") - if len(adapter_tokens) != 1: + if case.model_family is PatchInteropModelFamily.ANIMA and len(adapter_tokens) != 1: raise ValueError( "P9.7 regional PRIMARY_ADAPTER identity changed during sampling." ) @@ -225,6 +189,93 @@ def _validate_success( ) +def _validate_anima_snapshot( + snapshot: JsonObject, + adapter_tokens: set[str], +) -> None: + """Require exact Anima regional-LoRA execution evidence.""" + + uses = _array(snapshot.get("adapter_uses"), "adapter uses") + if len(uses) != 2: + raise ValueError( + "P9.7 accepted cases require paired positive/negative adapter uses." + ) + normalized_uses = tuple(_object(use, "adapter use") for use in uses) + if { + (_string(use.get("branch"), "adapter branch"), use.get("composition_index")) + for use in normalized_uses + } != {("positive", 0), ("negative", 1)}: + raise ValueError( + "P9.7 paired regional PRIMARY_ADAPTER branch ownership changed." + ) + for use in normalized_uses: + if ( + use.get("active") is not True + or use.get("region_index") != 0 + or use.get("target_count") != 448 + ): + raise ValueError("P9.7 exact regional PRIMARY_ADAPTER execution changed.") + if not math.isclose( + _number(use.get("effective_strength"), "effective strength"), + 0.75, + abs_tol=1e-8, + ): + raise ValueError("P9.7 regional PRIMARY_ADAPTER strength changed.") + adapter_tokens.add(_string(use.get("adapter_token"), "adapter token")) + work = _object(snapshot.get("estimated_work"), "estimated work") + if ( + work.get("active_adapter_uses") != 2 + or work.get("active_target_count") != 448 + or work.get("target_use_count") != 896 + ): + raise ValueError("P9.7 paired LoRA target-use accounting changed.") + _validate_single_denoiser_trajectory(work) + + +def _validate_sdxl_snapshot(snapshot: JsonObject) -> None: + """Require exact SDXL regional cross-attention execution evidence.""" + + if snapshot.get("region_count") != 2 or snapshot.get("active_region_indices") != [ + 0, + 1, + ]: + raise ValueError("P9.7 SDXL regional branch execution changed.") + work = _object(snapshot.get("estimated_work"), "estimated work") + if ( + work.get("cross_attention_branch_multiplier") != 3.0 + or work.get("cross_attention_formula") != "base_plus_region_count" + ): + raise ValueError("P9.7 SDXL regional attention accounting changed.") + _validate_single_denoiser_trajectory(work) + + +def _validate_single_denoiser_trajectory(work: JsonObject) -> None: + """Require regional work to retain one denoiser trajectory.""" + + if not math.isclose( + _number(work.get("denoiser_call_multiplier"), "denoiser multiplier"), + 1.0, + abs_tol=1e-8, + ): + raise ValueError("P9.7 denoiser trajectory multiplier changed.") + + +def _diagnostic_records_per_model_call(case: RegionalPatchInteropCase) -> int: + """Return the baseline diagnostic cardinality for one model family.""" + + if case.model_family is PatchInteropModelFamily.SDXL: + return 2 + return 1 + + +def _expected_backend(case: RegionalPatchInteropCase) -> str: + """Return the exact Comfy denoiser backend identity for one family.""" + + if case.model_family is PatchInteropModelFamily.SDXL: + return "comfy.ldm.modules.diffusionmodules.openaimodel.UNetModel" + return "comfy.ldm.anima.model.Anima" + + def _validate_model_call_count( case: RegionalPatchInteropCase, model_calls: int, diff --git a/tools/regional_patch_interop_integration/workflow.py b/tools/regional_patch_interop_integration/workflow.py index b710bf1..d78d313 100644 --- a/tools/regional_patch_interop_integration/workflow.py +++ b/tools/regional_patch_interop_integration/workflow.py @@ -171,7 +171,7 @@ class RegionalPatchInteropWorkflowBuilder: mask_names: tuple[str, ...], checkpoint_name: str, ) -> BuiltRegionalPatchInteropWorkflow: - """Build the focused SDXL NegPiP rejection graph.""" + """Build the focused SDXL NegPiP execution graph.""" graph = AnimaWorkflowGraph() loader = graph.add("CheckpointLoaderSimple", ckpt_name=checkpoint_name)