Compare commits

..
2 Commits
48 changed files with 2309 additions and 982 deletions
+7
View File
@@ -1,3 +1,10 @@
## [1.9.3](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.9.2...v1.9.3) (2026-09-20)
### Bug Fixes
* **attention-coupling:** restore regional LoRA sampling ([0255a0f](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/0255a0f5044278f14452b6b2582ec6646083f756))
## [1.9.2](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.9.1...v1.9.2) (2026-09-20)
+2 -2
View File
@@ -1,12 +1,12 @@
{
"name": "simple-syrup-comfyui",
"version": "1.9.2",
"version": "1.9.3",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "simple-syrup-comfyui",
"version": "1.9.2",
"version": "1.9.3",
"license": "AGPL-3.0-or-later",
"devDependencies": {
"@eslint/js": "^9.39.1",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "simple-syrup-comfyui",
"version": "1.9.2",
"version": "1.9.3",
"private": true,
"license": "AGPL-3.0-or-later",
"type": "module",
+1 -1
View File
@@ -5,7 +5,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "SimpleSyrup"
description = "Workflow-focused ComfyUI extensions for image generation."
version = "1.9.2"
version = "1.9.3"
license = "AGPL-3.0-or-later"
license-files = ["LICENSE"]
requires-python = ">=3.11"
+1 -1
View File
@@ -6,6 +6,6 @@
from __future__ import annotations
__version__ = "1.9.2"
__version__ = "1.9.3"
__all__: list[str] = ["__version__"]
@@ -53,12 +53,14 @@ class KSamplerAttentionCouplingV3(_ComfyNodeBase):
"With conditioning batches and masks, denoises supported Anima "
"and standard SD/SDXL models through one "
"shared trajectory while coupling global and masked regional "
"cross-attention. The input MODEL may carry a global LoRA. Anima "
"regions may also carry ordered, independently scheduled Prompt "
"Control model LoRAs whose overlapping deltas compose in declared "
"order. Runtime scales with active adapters, ranks, and targets. "
"Standard SD/SDXL regional model-side hooks and unsupported Anima "
"adapter targets fail before sampling."
"cross-attention. LoRAs on the input MODEL and Prompt Control model "
"LoRAs on global conditioning entry 0 apply across the image. Regions "
"may also carry ordered, independently scheduled model LoRAs "
"whose overlapping deltas compose in declared order. Runtime scales "
"with active adapters, ranks, and targets. "
"Global LoRA and regional LoRA retain independent schedules; "
"regional model-side hooks are supported on admitted model families. "
"Unsupported adapter targets fail before sampling."
),
search_aliases=[
"attention coupling",
@@ -50,12 +50,14 @@ class KSamplerContextualAttentionCouplingV3(_ComfyNodeBase):
description=(
"Preserves large-image composition through Contextual Diffusion "
"while coupling regional attention in every local and reduced-global "
"Anima or standard SD/SDXL view. Global LoRAs remain on the input "
"model. Anima regional LoRA stacks are prepared once, retain "
"Anima or standard SD/SDXL view. LoRAs on the input MODEL and Prompt "
"Control model LoRAs on global conditioning entry 0 apply in every "
"view. Regional LoRA stacks are prepared once, retain "
"independent schedules and full quality, and skip inactive work. "
"Optional SEGS guide the shared local tile plan. Standard SD/SDXL "
"regional model-side hooks and unsupported Anima targets fail before "
"sampling."
"Global LoRA and regional LoRA stacks remain independently scheduled; "
"regional model-side hooks are supported on admitted model families. "
"Optional SEGS guide the shared local tile plan. Unsupported adapter "
"targets fail before sampling."
),
search_aliases=[
"contextual attention coupling",
+10 -8
View File
@@ -248,7 +248,8 @@ def attention_coupling_ksampler_inputs(
"model",
tooltip=(
"Supported Anima or standard SD/SDXL model used for one shared "
"denoiser trajectory; apply global model LoRAs before connecting it."
"denoiser trajectory. LoRAs patched on this model and Prompt Control "
"model LoRAs on conditioning entry 0 apply globally."
),
),
*base[1:6],
@@ -258,8 +259,8 @@ def attention_coupling_ksampler_inputs(
tooltip=(
"Global-first positive conditioning: entry 0 is global and later "
"entries pair with masks. Regional Prompt Control WeightHooks may "
"contain ordered full-rank Anima LoRA stacks with independent "
"schedules; standard SD/SDXL rejects regional model-side hooks."
"contain ordered regional LoRA stacks with independent schedules. "
"Model LoRA hooks on entry 0 apply across the image."
),
),
comfy_io.MultiType.Input(
@@ -267,8 +268,9 @@ def attention_coupling_ksampler_inputs(
[comfy_io.Conditioning, conditioning_batch],
tooltip=(
"Global-first negative conditioning aligned to the same masks; "
"Anima regional LoRA hooks retain their negative-branch ownership "
"and independent schedules."
"its global model hooks must match the positive global entry. "
"Regional LoRA hooks retain their negative-branch ownership and "
"independent schedules."
),
),
comfy_io.Mask.Input(
@@ -278,7 +280,7 @@ def attention_coupling_ksampler_inputs(
"Optional ordered masks paired with conditioning entries 1 onward. "
"Leave disconnected with ordinary conditioning to bypass Attention "
"Coupling. In overlaps, prompt contributions are normalized while "
"Anima regional LoRA deltas add in declared adapter and region order."
"regional LoRA deltas add in declared adapter and region order."
),
),
comfy_io.Float.Input(
@@ -291,7 +293,7 @@ def attention_coupling_ksampler_inputs(
tooltip=(
"Balances regional cross-attention against the global prompt from "
"0 (global only) to 1 (regional only inside solid masks); regional "
"Anima LoRA strength remains controlled by each hook."
"LoRA strength remains controlled by each hook."
),
),
comfy_io.Int.Input(
@@ -301,7 +303,7 @@ def attention_coupling_ksampler_inputs(
max=512,
step=1,
tooltip=(
"Softens Attention Coupling and Anima regional LoRA boundaries by "
"Softens Attention Coupling and regional LoRA boundaries by "
"this many image pixels; 0 preserves authored mask values."
),
),
@@ -54,12 +54,15 @@ class KSamplerTiledAttentionCouplingV3(_ComfyNodeBase):
"batches and masks, denoises large Anima and standard SD/SDXL "
"latents in tiles through "
"one shared model trajectory per tile batch while coupling global "
"and masked regional cross-attention. The input MODEL may carry "
"global LoRAs. Anima regions may carry independently scheduled "
"and masked regional cross-attention. LoRAs on the input MODEL and "
"Prompt Control model LoRAs on global conditioning entry 0 apply "
"across every tile. Regions may carry independently scheduled "
"regional LoRA stacks; inactive attention and LoRA work is pruned "
"without changing quality. MultiDiffusion or Mixture of Diffusers "
"fuses restored tile predictions. Standard SD/SDXL regional "
"model-side hooks and unsupported Anima targets fail before sampling."
"fuses restored tile predictions. Global LoRA and regional LoRA "
"stacks retain independent schedules; regional model-side hooks are "
"supported on admitted model families. Unsupported adapter targets "
"fail before sampling."
),
search_aliases=[
"attention coupling tiled",
+53 -44
View File
@@ -8,17 +8,19 @@ from __future__ import annotations
from dataclasses import dataclass
import comfy.model_patcher
from comfy.patcher_extension import CallbacksMP
from ..model_attention_patch_mutations import ModelAttn2PatchesMutation
from ..model_patcher_mutations import ModelKeyedCallbackMutation
from ..patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
from ..ppm_negpip_interop import PpmNegpipInterop
from ..regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
from ..regional_lora.operation_assembly import REGIONAL_OPERATION_ASSEMBLER
from ..regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
)
from ..regional_lora.standard_unet_variant_runtime import (
StandardUnetVariantRuntimeMutation,
)
from ..regional_lora.standard_unet_variant_template import (
STANDARD_UNET_VARIANT_TEMPLATE_CACHE,
from ..regional_lora.standard_unet_operation_session import (
StandardUnetRegionalOperationSession,
)
from .unet_attention_context_wrapper import unet_attention_context_wrapper_mutation
from .unet_attention_phase_session import StandardUnetAttentionPhaseSession
@@ -43,15 +45,15 @@ class StandardUnetAttentionBackend:
*,
model: object,
state: StandardUnetAttentionState,
admission: StandardUnetNativeLoraAdmission,
admission: StandardUnetOperationAdmission,
negpip: PpmNegpipInterop | None = None,
) -> StandardUnetAttentionModel:
"""Return a direct MODEL child containing only the paired UNet patches."""
if not isinstance(state, StandardUnetAttentionState):
raise TypeError("Standard UNet backend requires attention state.")
if not isinstance(admission, StandardUnetNativeLoraAdmission):
raise TypeError("Standard UNet backend requires native admission.")
if not isinstance(admission, StandardUnetOperationAdmission):
raise TypeError("Standard UNet backend requires operation admission.")
if admission.adaptation.plan != state.plan.lora_plan:
raise ValueError(
"Standard UNet admission and processed conditioning must share "
@@ -60,48 +62,55 @@ class StandardUnetAttentionBackend:
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)
if admission.adaptation.plan.adapters
else None
)
variant_mutations = (
(
StandardUnetVariantRuntimeMutation(
state,
admission,
attention_phase,
template,
negpip,
operation_session: StandardUnetRegionalOperationSession | None = None
operation_mutations: tuple[ModelMutation, ...] = ()
if admission.adaptation.plan.adapters:
if (
not isinstance(model, comfy.model_patcher.ModelPatcher)
or admission.binding is None
or admission.cache is None
):
raise TypeError(
"Standard UNet regional operations require complete MODEL "
"admission."
)
assembly = REGIONAL_OPERATION_ASSEMBLER.assemble(
admission.binding,
model=model,
cache=admission.cache,
)
operation_session = StandardUnetRegionalOperationSession(
admission.adaptation.plan,
state.plan.mask_bank,
admission.module_roles,
assembly.call_scope,
)
operation_mutations = (
assembly.cache_lifecycle.mutation(),
ModelKeyedCallbackMutation(
CallbacksMP.ON_DETACH,
"simple_syrup.standard_unet_regional_operation_schedule",
operation_session.clear,
),
)
if template is not None
else ()
patches = UnetAttn2PatchPair(
StandardUnetAttn2ExecutionResolver(state),
operation_scope=operation_session,
)
derivation_source = (
template.bind_request(model) if template is not None else model
)
attention_mutations: tuple[ModelMutation, ...] = ()
if template is None:
patches = UnetAttn2PatchPair(
StandardUnetAttn2ExecutionResolver(state),
)
attention_mutations = (
derived = PATCHER_LIFECYCLE.derive_model(
model,
(
unet_attention_context_wrapper_mutation(
state,
attention_phase,
operation_session,
),
ModelAttn2PatchesMutation(
patches.input_patch,
patches.output_patch,
(() if negpip is None else (negpip.attention_patch,)),
),
)
derived = PATCHER_LIFECYCLE.derive_model(
derivation_source,
(
unet_attention_context_wrapper_mutation(
state,
attention_phase,
),
*attention_mutations,
*variant_mutations,
*operation_mutations,
),
operation="standard UNet Attention Coupling",
)
@@ -6,12 +6,17 @@
from __future__ import annotations
from contextlib import ExitStack
import torch
from ..diffusion_wrapper_executor import DiffusionWrapperExecutor
from ..diffusion_wrapper_invocation import DIFFUSION_WRAPPER_INVOCATION_VALIDATOR
from ..model_patcher_mutations import ModelDiffusionWrapperMutation
from ..regional_attention_model_call import RegionalAttentionModelCallResolver
from ..regional_lora.standard_unet_operation_session import (
StandardUnetRegionalOperationSession,
)
from .standard_unet_model_output_validation import (
STANDARD_UNET_MODEL_OUTPUT_VALIDATOR,
StandardUnetModelOutputValidator,
@@ -32,6 +37,7 @@ class StandardUnetAttentionContextDiffusionWrapper:
self,
state: StandardUnetAttentionState,
attention_phase: StandardUnetAttentionPhaseSession,
operation_session: StandardUnetRegionalOperationSession | None = None,
*,
model_call_resolver: RegionalAttentionModelCallResolver = (
STANDARD_UNET_MODEL_CALL_RESOLVER
@@ -46,12 +52,18 @@ class StandardUnetAttentionContextDiffusionWrapper:
raise TypeError("Standard UNet context wrapper requires attention state.")
if not isinstance(attention_phase, StandardUnetAttentionPhaseSession):
raise TypeError("Standard UNet context wrapper requires phase state.")
if operation_session is not None and not isinstance(
operation_session,
StandardUnetRegionalOperationSession,
):
raise TypeError("Standard UNet operation session has an invalid type.")
if not isinstance(model_call_resolver, RegionalAttentionModelCallResolver):
raise TypeError(
"Standard UNet context wrapper requires a model-call resolver."
)
self._state = state
self._attention_phase = attention_phase
self._operation_session = operation_session
self._model_call_resolver = model_call_resolver
if not isinstance(output_validator, StandardUnetModelOutputValidator):
raise TypeError("Standard UNet output validator has an invalid type.")
@@ -90,11 +102,14 @@ class StandardUnetAttentionContextDiffusionWrapper:
transformer_options=args[5],
)
forwarded_args = (*args[:2], contexts.base_context, *args[3:])
with (
self._attention_phase.activate(args[5]),
self._state.execution_context.activate(contexts),
self._state.resolution_cache.activate(),
):
with ExitStack() as scopes:
scopes.enter_context(self._attention_phase.activate(args[5]))
scopes.enter_context(self._state.execution_context.activate(contexts))
scopes.enter_context(self._state.resolution_cache.activate())
if self._operation_session is not None:
scopes.enter_context(
self._operation_session.activate(contexts, args[5])
)
output = executor(*forwarded_args, **kwargs)
return self._output_validator.validate(output, model_input=args[0])
@@ -102,6 +117,7 @@ class StandardUnetAttentionContextDiffusionWrapper:
def unet_attention_context_wrapper_mutation(
state: StandardUnetAttentionState,
attention_phase: StandardUnetAttentionPhaseSession,
operation_session: StandardUnetRegionalOperationSession | None = None,
) -> ModelDiffusionWrapperMutation:
"""Return the clone-local standard-UNet context wrapper mutation."""
@@ -110,5 +126,6 @@ def unet_attention_context_wrapper_mutation(
StandardUnetAttentionContextDiffusionWrapper(
state,
attention_phase,
operation_session,
),
)
@@ -0,0 +1,74 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Compose global-first conditioning hooks for conventional regional sampling."""
from __future__ import annotations
from typing import Any, TypeAlias
from comfy.hooks import HookGroup
from .regional_lora_conditioning_sources import conditioning_hook_groups
Conditioning: TypeAlias = list[list[Any]]
class GlobalFirstConditioningHookComposer:
"""Apply one global HookGroup to every regional conditioning model state."""
def global_hooks(
self,
conditioning: Conditioning,
*,
source_label: str,
) -> HookGroup | None:
"""Return the single uniform HookGroup carried by global conditioning."""
groups = conditioning_hook_groups(conditioning)
if not groups:
return None
authority = groups[0]
if any(group is not authority for group in groups[1:]):
raise ValueError(
f"{source_label} uses different HookGroups across conditioning "
"entries. Keep one shared Prompt Control hook schedule on the "
"global segment."
)
return authority
def compose(
self,
conditioning: Conditioning,
global_hooks: HookGroup | None,
*,
source_label: str,
cache: dict[tuple[HookGroup, HookGroup], HookGroup],
) -> Conditioning:
"""Prepend global hooks to every local HookGroup without mutating inputs."""
if global_hooks is None:
return [[item[0], dict(item[1])] for item in conditioning]
composed: Conditioning = []
for item_index, item in enumerate(conditioning):
metadata = dict(item[1])
local_hooks = metadata.get("hooks")
if local_hooks is None:
metadata["hooks"] = global_hooks
elif not isinstance(local_hooks, HookGroup):
raise TypeError(
f"{source_label} item {item_index} hooks must be a Comfy HookGroup."
)
else:
key = (global_hooks, local_hooks)
combined = cache.get(key)
if combined is None:
combined = global_hooks.clone_and_combine(local_hooks)
cache[key] = combined
metadata["hooks"] = combined
composed.append([item[0], metadata])
return composed
GLOBAL_FIRST_CONDITIONING_HOOK_COMPOSER = GlobalFirstConditioningHookComposer()
@@ -0,0 +1,63 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Resolve the effective Comfy MODEL before global conditioning hooks execute."""
from __future__ import annotations
from typing import cast
from ..domain.conditioning_batch import select_conditioning
from .regional_lora_conditioning_sources import conditioning_hook_groups
class GlobalHookModelResolver:
"""Mirror Comfy's dynamic-to-static handoff before regional derivation."""
def resolve(
self,
model: object,
*,
positive: object,
negative: object,
) -> object:
"""Return the model that Comfy will use for global hooked conditioning."""
if not self._has_global_hooks(positive) and not self._has_global_hooks(
negative
):
return model
is_dynamic = getattr(model, "is_dynamic", None)
if not callable(is_dynamic):
raise TypeError(
"Global conditioning hooks require MODEL dynamic-mode state."
)
dynamic = is_dynamic()
if not isinstance(dynamic, bool):
raise TypeError("MODEL is_dynamic() must return a bool.")
if not dynamic:
return model
delegate_factory = getattr(model, "get_non_dynamic_delegate", None)
if not callable(delegate_factory):
raise TypeError(
"Dynamic MODEL global conditioning hooks require Comfy's "
"get_non_dynamic_delegate()."
)
resolved = delegate_factory()
if resolved is model:
raise RuntimeError("Dynamic MODEL returned itself as its static delegate.")
resolved_is_dynamic = getattr(resolved, "is_dynamic", None)
if not callable(resolved_is_dynamic) or resolved_is_dynamic() is not False:
raise RuntimeError("Global conditioning hook delegate must be static.")
return cast(object, resolved)
@staticmethod
def _has_global_hooks(conditioning: object) -> bool:
"""Report hooks only on consumer-defined global entry zero."""
global_conditioning = select_conditioning(conditioning, 0)
return bool(conditioning_hook_groups(global_conditioning))
GLOBAL_HOOK_MODEL_RESOLVER = GlobalHookModelResolver()
@@ -20,7 +20,6 @@ from .anima_attention_coupling import anima_attention_coupling_mutations
from .anima_attention_execution import AnimaRegionalAttentionExecution
from .anima_composition import AnimaRegionalLoraComposition
from .anima_execution_scope import AnimaRegionalLoraAdapterExecution
from .anima_global_lora_overlap import ANIMA_GLOBAL_REGIONAL_LORA_OVERLAP_VALIDATOR
from .anima_model_patcher_surface import ANIMA_MODEL_PATCHER_SURFACE_RESOLVER
from .anima_plan_admission import ANIMA_REGIONAL_LORA_PLAN_ADMISSION_SERVICE
from .execution_cache import ModelCloneLineage, RegionalLoraExecutionCache
@@ -56,7 +55,6 @@ class FullContextAnimaAttentionBackend:
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(
processed_plan,
latent_batch_size=latent_batch_size,
@@ -1,134 +0,0 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Reject exact static-global and admitted-regional Anima LoRA overlap."""
from __future__ import annotations
import math
from collections.abc import Mapping
import torch
from comfy.weight_adapter.lora import LoRAAdapter
from .anima_plan_admission import (
AnimaRegionalLoraAdapterAdmission,
AnimaRegionalLoraPlanAdmission,
)
from .standard_adapter import StandardLoraTarget
class AnimaGlobalRegionalLoraOverlapError(ValueError):
"""Report regional adapters already present in the static global MODEL."""
class AnimaGlobalRegionalLoraOverlapValidator:
"""Compare exact admitted regional A/B tensors with static global patches."""
def validate(
self,
model: object,
admission: AnimaRegionalLoraPlanAdmission,
) -> None:
"""Reject every regional adapter whose complete content is global."""
if not isinstance(admission, AnimaRegionalLoraPlanAdmission):
raise TypeError("Anima global LoRA overlap requires an admitted plan.")
patches = getattr(model, "patches", None)
if not isinstance(patches, Mapping):
raise TypeError("Anima global LoRA overlap requires MODEL patches.")
duplicates = tuple(
adapter.adapter_plan.adapter_identity.value
for adapter in admission.adapters
if self._duplicates_global_content(patches, adapter)
)
unique_duplicates = tuple(dict.fromkeys(duplicates))
if unique_duplicates:
identities = ", ".join(repr(value) for value in unique_duplicates)
raise AnimaGlobalRegionalLoraOverlapError(
"Regional Anima LoRA content is already applied globally to the "
f"input MODEL: {identities}. Remove either the global or regional "
"application before sampling."
)
def _duplicates_global_content(
self,
patches: Mapping[object, object],
adapter: AnimaRegionalLoraAdapterAdmission,
) -> bool:
"""Return whether every admitted regional target has an exact global pair."""
targets = adapter.admission.targets
return bool(targets) and all(
self._target_matches(patches, target.adapter) for target in targets
)
def _target_matches(
self,
patches: Mapping[object, object],
regional: StandardLoraTarget,
) -> bool:
"""Match one regional target against nonzero comparable static patches."""
key = f"{regional.target}.weight"
entries = patches.get(key, ())
if entries == ():
return False
if not isinstance(entries, list):
raise TypeError(f"MODEL patches[{key!r}] must be a list.")
for index, entry in enumerate(entries):
if not isinstance(entry, tuple) or len(entry) < 3:
raise TypeError(
f"MODEL patches[{key!r}][{index}] must be a Comfy patch tuple."
)
if _nonzero_strength(entry[0], key=key, index=index) and _matches_pair(
entry[1],
regional,
):
return True
return False
def _nonzero_strength(value: object, *, key: str, index: int) -> bool:
"""Validate one installed static patch strength and report its activity."""
if isinstance(value, bool) or not isinstance(value, int | float):
raise TypeError(f"MODEL patches[{key!r}][{index}] strength must be numeric.")
strength = float(value)
if not math.isfinite(strength):
raise ValueError(f"MODEL patches[{key!r}][{index}] strength must be finite.")
return strength != 0.0
def _matches_pair(value: object, regional: StandardLoraTarget) -> bool:
"""Compare one installed standard LoRA patch without copies or transfers."""
if not isinstance(value, LoRAAdapter):
return False
weights = value.weights
if not isinstance(weights, tuple) or len(weights) != 6:
return False
up, down, alpha, mid, dora_scale, reshape = weights
if any(item is not None for item in (alpha, mid, dora_scale, reshape)):
return False
if not isinstance(down, torch.Tensor) or not isinstance(up, torch.Tensor):
return False
return _same_tensor(down, regional.down) and _same_tensor(up, regional.up)
def _same_tensor(left: torch.Tensor, right: torch.Tensor) -> bool:
"""Use an identity fast path before exact same-residency tensor equality."""
if left is right:
return True
if (
left.shape != right.shape
or left.dtype != right.dtype
or left.device != right.device
):
return False
return bool(torch.equal(left, right))
ANIMA_GLOBAL_REGIONAL_LORA_OVERLAP_VALIDATOR = AnimaGlobalRegionalLoraOverlapValidator()
@@ -0,0 +1,126 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Admit complete standard-UNet regional LoRA operation surfaces."""
from __future__ import annotations
from dataclasses import dataclass
import comfy.model_patcher
from torch import nn
from ..attention_coupling.family_admission import AttentionCouplingFamilyAdmission
from ..regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from .comfy_adapter_resolver import COMFY_REGIONAL_ADAPTER_RESOLVER
from .execution_cache import RegionalLoraExecutionCache
from .resolved_operation_translator import COMFY_RESOLVED_OPERATION_TRANSLATOR
from .standard_unet_target_capabilities import (
STANDARD_UNET_TARGET_CAPABILITY_CLASSIFIER,
)
from .target_binder import REGIONAL_LORA_TARGET_BINDER
from .target_binding import (
BoundRegionalLoraSpatialCapability,
RegionalLoraBindingResult,
)
@dataclass(frozen=True, slots=True)
class StandardUnetOperationAdmission(AttentionCouplingFamilyAdmission):
"""Retain target bindings and exact runtime consumer-role evidence."""
binding: RegionalLoraBindingResult | None
module_roles: dict[str, BoundRegionalLoraSpatialCapability]
cache: RegionalLoraExecutionCache | None
def __post_init__(self) -> None:
"""Require either an empty admission or a complete executable surface."""
AttentionCouplingFamilyAdmission.__post_init__(self)
if not self.adaptation.plan.adapters:
if self.binding is not None or self.module_roles or self.cache is not None:
raise ValueError("Empty standard admission cannot retain operations.")
return
if (
not isinstance(self.binding, RegionalLoraBindingResult)
or not self.binding.admissible
or not self.binding.entries
):
raise ValueError("Standard admission requires complete target binding.")
if not self.module_roles or not isinstance(
self.cache,
RegionalLoraExecutionCache,
):
raise ValueError("Standard admission requires roles and execution cache.")
class StandardUnetOperationPreparation:
"""Resolve, translate, bind, and classify regional LoRA operations."""
def admit(
self,
model: object,
adaptation: RegionalLoraPlanAdaptation,
) -> StandardUnetOperationAdmission:
"""Return complete immutable evidence before installing call-scoped work."""
if not isinstance(adaptation, RegionalLoraPlanAdaptation):
raise TypeError("Standard UNet operation admission requires adaptation.")
if not adaptation.plan.adapters:
return StandardUnetOperationAdmission(adaptation, None, {}, None)
if not isinstance(model, comfy.model_patcher.ModelPatcher):
raise TypeError("Standard UNet operation admission requires a MODEL.")
graph_root = model.model
if not isinstance(graph_root, nn.Module):
raise TypeError("Standard UNet MODEL graph must be an nn.Module.")
capabilities = STANDARD_UNET_TARGET_CAPABILITY_CLASSIFIER.classify(graph_root)
resolution = COMFY_REGIONAL_ADAPTER_RESOLVER.resolve(
adaptation,
model=model,
)
operations = COMFY_RESOLVED_OPERATION_TRANSLATOR.translate(resolution)
binding = REGIONAL_LORA_TARGET_BINDER.bind(
source=model,
candidate=model,
resolution=resolution,
operations=operations,
linear_spatial_capabilities=capabilities.linear_roles,
)
if not binding.admissible:
messages = tuple(issue.message for issue in binding.issues)
raise ValueError(
f"Standard UNet regional LoRA target admission failed: {messages!r}."
)
unavailable = tuple(
entry.descriptor.target.parameter_path
for entry in binding.entries
if entry.spatial_capability
in (
BoundRegionalLoraSpatialCapability.GLOBAL_ONLY,
BoundRegionalLoraSpatialCapability.UNSUPPORTED,
)
)
if unavailable:
raise ValueError(
"Standard UNet regional LoRA targets lack executable consumer roles: "
f"{unavailable!r}."
)
module_roles: dict[str, BoundRegionalLoraSpatialCapability] = {}
for entry in binding.entries:
path = entry.descriptor.target.model_target
role = entry.spatial_capability
previous = module_roles.setdefault(path, role)
if previous is not role:
raise ValueError(
f"Standard UNet operation {path!r} has conflicting roles."
)
return StandardUnetOperationAdmission(
adaptation,
binding,
module_roles,
RegionalLoraExecutionCache(),
)
STANDARD_UNET_OPERATION_PREPARATION = StandardUnetOperationPreparation()
@@ -0,0 +1,333 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Resolve spatial regional LoRA operations within one standard-UNet call."""
from __future__ import annotations
from collections.abc import Iterator, Mapping
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
import torch
from ...domain.regional_activation_geometry import (
RegionalActivationGeometry,
RegionalActivationLayout,
RegionalTemporalOwnership,
)
from ...domain.regional_attention_batch import BatchedRegionalAttentionContexts
from ...domain.regional_lora_plan import RegionalLoraPlan
from ...domain.regional_mask_bank import RegionalMaskBank
from ...domain.spatial_views import SpatialBatchLayout
from ...masking.regional_activation_mask_projection import (
REGIONAL_ACTIVATION_MASK_PROJECTOR,
)
from ...masking.regional_mask_projection import (
RegionalMaskForm,
RegionalMaskProjectionMode,
)
from ..attention_coupling.unet_attn2_execution import UnetAttn2Execution
from ..spatial_model_arguments import (
SIMPLE_SYRUP_TRANSFORMER_NAMESPACE,
SPATIAL_BATCH_LAYOUT_KEY,
)
from .activation_batch_alignment import (
REGIONAL_ACTIVATION_BATCH_ALIGNMENT_RESOLVER,
)
from .convolution_execution_plan import RegionalConvolutionExecutionPlan
from .convolution_rank_geometry import REGIONAL_CONVOLUTION_RANK_GEOMETRY_RESOLVER
from .linear_execution_plan import RegionalLinearExecutionPlan
from .operation_call_scope import RegionalOperationCallScope
from .operation_invocation import (
REGIONAL_OPERATION_INVOCATION_CONTEXT,
RegionalOperationExecutionPlan,
RegionalOperationInvocation,
RegionalOperationInvocationContext,
)
from .operation_mask_resolution import REGIONAL_OPERATION_MASK_RESOLVER
from .standard_unet_lora_schedule import StandardUnetLoraSchedule
from .standard_unet_packed_operation_masks import (
STANDARD_UNET_PACKED_OPERATION_MASK_RESOLVER,
)
from .target_binding import BoundRegionalLoraSpatialCapability
@dataclass(slots=True)
class _ActiveStandardUnetOperationCall:
"""Retain call authorities and optional compact attn2 execution state."""
contexts: BatchedRegionalAttentionContexts
transformer_options: dict[str, object]
schedule_strengths: tuple[float, ...]
packed_execution: UnetAttn2Execution | None = None
class StandardUnetRegionalOperationSession:
"""Own one shared UNet trajectory with spatial regional LoRA deltas."""
def __init__(
self,
plan: RegionalLoraPlan,
mask_bank: RegionalMaskBank,
module_roles: Mapping[str, BoundRegionalLoraSpatialCapability],
call_scope: RegionalOperationCallScope,
*,
invocation_context: RegionalOperationInvocationContext = (
REGIONAL_OPERATION_INVOCATION_CONTEXT
),
) -> None:
"""Retain immutable composition, mask, role, and operation authorities."""
if not isinstance(plan, RegionalLoraPlan) or not plan.adapters:
raise ValueError("Standard UNet operation session requires adapters.")
if not isinstance(mask_bank, RegionalMaskBank):
raise TypeError("Standard UNet operation session requires a mask bank.")
if not isinstance(module_roles, Mapping) or not module_roles:
raise ValueError("Standard UNet operation session requires module roles.")
roles = dict(module_roles)
if any(not isinstance(path, str) or not path for path in roles):
raise ValueError("Standard UNet operation paths must be nonempty.")
supported = (
BoundRegionalLoraSpatialCapability.SPATIAL_TOKENS,
BoundRegionalLoraSpatialCapability.PACKED_IMAGE_TOKENS,
BoundRegionalLoraSpatialCapability.PACKED_CONTEXT_TOKENS,
BoundRegionalLoraSpatialCapability.DIRECT,
)
if any(role not in supported for role in roles.values()):
raise ValueError("Standard UNet operation role is unsupported.")
if not isinstance(call_scope, RegionalOperationCallScope):
raise TypeError("Standard UNet operation session requires a call scope.")
if not isinstance(invocation_context, RegionalOperationInvocationContext):
raise TypeError(
"Standard UNet operation session requires invocation context."
)
self._mask_bank = mask_bank
self._module_roles = roles
self._call_scope = call_scope
self._invocation_context = invocation_context
self._schedule = StandardUnetLoraSchedule(plan)
self._active: ContextVar[_ActiveStandardUnetOperationCall | None] = ContextVar(
"simple_syrup_standard_unet_regional_operation_call",
default=None,
)
@contextmanager
def activate(
self,
contexts: BatchedRegionalAttentionContexts,
transformer_options: dict[str, object],
) -> Iterator[None]:
"""Publish operation masks and install wrappers for one model call."""
if not isinstance(contexts, BatchedRegionalAttentionContexts):
raise TypeError("Standard UNet operation call requires contexts.")
if not isinstance(transformer_options, dict):
raise TypeError("Standard UNet operation call requires options.")
active = _ActiveStandardUnetOperationCall(
contexts,
transformer_options,
self._schedule.resolve(transformer_options),
)
token = self._active.set(active)
try:
with (
self._invocation_context.activate(self),
self._call_scope.activate(),
):
yield
finally:
active.packed_execution = None
self._active.reset(token)
def begin_packed(self, execution: UnetAttn2Execution) -> None:
"""Publish compact attn2 execution until its paired output callback."""
active = self._require_active()
if active.packed_execution is not None:
raise ValueError("Standard UNet operation call already has packed state.")
if not isinstance(execution, UnetAttn2Execution):
raise TypeError("Standard UNet packed state requires attn2 execution.")
active.packed_execution = execution
def end_packed(self, execution: UnetAttn2Execution) -> None:
"""Clear only the compact execution opened by the input callback."""
active = self._require_active()
if active.packed_execution is not execution:
raise ValueError("Standard UNet packed output does not match input state.")
active.packed_execution = None
def resolve(
self,
module_path: str,
plan: RegionalOperationExecutionPlan,
inputs: torch.Tensor,
) -> RegionalOperationInvocation | None:
"""Resolve one installed operation from its declared consumer role."""
active = self._require_active()
role = self._module_roles.get(module_path)
if role is None:
return None
strengths = tuple(
active.schedule_strengths[use.composition_index] for use in plan.uses
)
if role is BoundRegionalLoraSpatialCapability.PACKED_IMAGE_TOKENS:
execution = self._require_packed(active)
if not isinstance(plan, RegionalLinearExecutionPlan):
raise TypeError("Packed image role requires a Linear plan.")
masks = STANDARD_UNET_PACKED_OPERATION_MASK_RESOLVER.resolve_image_tokens(
execution,
uses=plan.uses,
inputs=inputs,
)
elif role is BoundRegionalLoraSpatialCapability.PACKED_CONTEXT_TOKENS:
execution = self._require_packed(active)
if not isinstance(plan, RegionalLinearExecutionPlan):
raise TypeError("Packed context role requires a Linear plan.")
masks = STANDARD_UNET_PACKED_OPERATION_MASK_RESOLVER.resolve_context_tokens(
execution,
uses=plan.uses,
inputs=inputs,
)
else:
if active.packed_execution is not None:
raise ValueError("Ordinary regional operation ran inside packed attn2.")
geometry = self._ordinary_geometry(
role,
plan=plan,
inputs=inputs,
active=active,
)
spatial = REGIONAL_ACTIVATION_MASK_PROJECTOR.project(
bank=self._mask_bank,
geometry=geometry,
form=RegionalMaskForm.CONDITIONING,
mode=RegionalMaskProjectionMode.CONTINUOUS_COVERAGE,
device=inputs.device,
dtype=inputs.dtype,
)
masks = REGIONAL_OPERATION_MASK_RESOLVER.resolve(
spatial,
contexts=active.contexts,
uses=plan.uses,
)
return RegionalOperationInvocation(masks, strengths)
def clear(self, model: object, unpatch_all: bool) -> None:
"""Release retained sampling schedule state on model detach."""
del model, unpatch_all
self._schedule.clear()
def _ordinary_geometry(
self,
role: BoundRegionalLoraSpatialCapability,
*,
plan: RegionalOperationExecutionPlan,
inputs: torch.Tensor,
active: _ActiveStandardUnetOperationCall,
) -> RegionalActivationGeometry:
"""Resolve exact ordinary token or convolution activation geometry."""
layout = _spatial_layout(active.transformer_options)
alignment = REGIONAL_ACTIVATION_BATCH_ALIGNMENT_RESOLVER.resolve(
active.contexts,
spatial_layout=layout,
)
if role is BoundRegionalLoraSpatialCapability.SPATIAL_TOKENS:
if not isinstance(plan, RegionalLinearExecutionPlan) or inputs.ndim != 3:
raise ValueError("Spatial-token role requires B/S/C Linear inputs.")
activation_shape = _activation_shape(active.transformer_options)
if int(inputs.shape[0]) != activation_shape[0] or int(inputs.shape[1]) != (
activation_shape[2] * activation_shape[3]
):
raise ValueError("Spatial-token inputs must match live activation H/W.")
return RegionalActivationGeometry(
RegionalActivationLayout.CONSUMER_SPATIALIZED,
tuple(inputs.shape),
2,
activation_shape[2],
activation_shape[3],
alignment,
)
if role is not BoundRegionalLoraSpatialCapability.DIRECT or not isinstance(
plan,
RegionalConvolutionExecutionPlan,
):
raise ValueError("Standard UNet ordinary operation role is inconsistent.")
use = plan.uses[0]
spatial = REGIONAL_CONVOLUTION_RANK_GEOMETRY_RESOLVER.resolve(
tuple(int(value) for value in inputs.shape[2:]),
use,
)
rank_channels = int(use.preparation.down.shape[0]) * use.parameters.groups
layouts = {
1: RegionalActivationLayout.DIRECT_CONVOLUTION_1D,
2: RegionalActivationLayout.DIRECT_CONVOLUTION_2D,
3: RegionalActivationLayout.DIRECT_CONVOLUTION_3D,
}
return RegionalActivationGeometry(
layouts[use.parameters.dimension],
(int(inputs.shape[0]), rank_channels, *spatial),
1,
1 if use.parameters.dimension == 1 else spatial[-2],
spatial[-1],
alignment,
temporal_axis=2 if use.parameters.dimension == 3 else None,
temporal_ownership=(
RegionalTemporalOwnership.REPEAT_SPATIAL_MASK
if use.parameters.dimension == 3
else RegionalTemporalOwnership.NONE
),
)
def _require_active(self) -> _ActiveStandardUnetOperationCall:
"""Return the current call or reject execution outside its owner."""
active = self._active.get()
if active is None:
raise RuntimeError("Standard UNet regional operation ran outside a call.")
return active
@staticmethod
def _require_packed(
active: _ActiveStandardUnetOperationCall,
) -> UnetAttn2Execution:
"""Return the compact attn2 authority for the current projection."""
if active.packed_execution is None:
raise RuntimeError("Packed regional operation ran outside attn2 scope.")
return active.packed_execution
def _activation_shape(options: dict[str, object]) -> tuple[int, int, int, int]:
"""Narrow Comfy's live spatial-transformer BCHW metadata."""
value = options.get("activations_shape")
if not isinstance(value, list | tuple) or len(value) != 4:
raise TypeError("Standard UNet activations_shape must be a BCHW sequence.")
shape = tuple(value)
if any(
isinstance(item, bool) or not isinstance(item, int) or item < 1
for item in shape
):
raise ValueError("Standard UNet activation dimensions must be positive.")
return shape[0], shape[1], shape[2], shape[3]
def _spatial_layout(options: dict[str, object]) -> SpatialBatchLayout | None:
"""Return the optional authoritative full, tiled, or Contextual layout."""
namespace = options.get(SIMPLE_SYRUP_TRANSFORMER_NAMESPACE)
if namespace is None:
return None
if not isinstance(namespace, dict):
raise TypeError("Standard UNet SimpleSyrup namespace must be a dictionary.")
layout = namespace.get(SPATIAL_BATCH_LAYOUT_KEY)
if layout is not None and not isinstance(layout, SpatialBatchLayout):
raise TypeError("Standard UNet spatial layout has an invalid type.")
return layout
@@ -0,0 +1,147 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Resolve regional operation masks for compact standard-UNet attn2 rows."""
from __future__ import annotations
from collections.abc import Sequence
import torch
from ...domain.regional_activation_geometry import (
RegionalActivationBatchAlignment,
RegionalActivationGeometry,
RegionalActivationLayout,
)
from ..attention_coupling.unet_attn2_execution import UnetAttn2Execution
from .operation_mask_resolution import (
REGIONAL_OPERATION_BRANCH_GATE_RESOLVER,
RegionalOperationMaskBatch,
RegionalOperationMaskUse,
)
class StandardUnetPackedOperationMaskResolver:
"""Map regional uses onto exact compact attn2 image or context rows."""
def resolve_image_tokens(
self,
execution: UnetAttn2Execution,
*,
uses: Sequence[RegionalOperationMaskUse],
inputs: torch.Tensor,
) -> RegionalOperationMaskBatch:
"""Return query-grid masks for packed query and output projections."""
self._validate(execution, uses=uses, inputs=inputs)
if int(inputs.shape[1]) != execution.query_height * execution.query_width:
raise ValueError("Packed image tokens must match attn2 query H/W.")
use_masks = tuple(
self._packed_use_mask(execution, use=use, inputs=inputs, spatial=True)
for use in uses
)
geometry = RegionalActivationGeometry(
RegionalActivationLayout.CONSUMER_SPATIALIZED,
tuple(inputs.shape),
2,
execution.query_height,
execution.query_width,
RegionalActivationBatchAlignment(int(inputs.shape[0]), 1),
)
return RegionalOperationMaskBatch(
torch.stack(use_masks),
geometry,
tuple(use.composition_index for use in uses),
)
def resolve_context_tokens(
self,
execution: UnetAttn2Execution,
*,
uses: Sequence[RegionalOperationMaskUse],
inputs: torch.Tensor,
) -> RegionalOperationMaskBatch:
"""Return branch gates broadcast over untouched context tokens."""
self._validate(execution, uses=uses, inputs=inputs)
use_masks = tuple(
self._packed_use_mask(execution, use=use, inputs=inputs, spatial=False)
for use in uses
)
geometry = RegionalActivationGeometry(
RegionalActivationLayout.BRANCH_TOKENS,
tuple(inputs.shape),
2,
1,
int(inputs.shape[1]),
RegionalActivationBatchAlignment(int(inputs.shape[0]), 1),
)
return RegionalOperationMaskBatch(
torch.stack(use_masks),
geometry,
tuple(use.composition_index for use in uses),
)
@staticmethod
def _validate(
execution: object,
*,
uses: Sequence[RegionalOperationMaskUse],
inputs: object,
) -> None:
"""Require one exact packed B/S/C activation and ordered use sequence."""
if not isinstance(execution, UnetAttn2Execution):
raise TypeError("Packed operation masks require an attn2 execution.")
if not isinstance(inputs, torch.Tensor) or inputs.ndim != 3:
raise ValueError("Packed operation inputs must use B/S/C layout.")
if int(inputs.shape[0]) != execution.branches.packed_batch_size:
raise ValueError("Packed operation batch must match attn2 branches.")
if not isinstance(uses, Sequence) or not uses:
raise ValueError("Packed operation masks require target uses.")
composition = tuple(use.composition_index for use in uses)
if composition != tuple(sorted(composition)):
raise ValueError("Packed operation uses must follow composition order.")
@staticmethod
def _packed_use_mask(
execution: UnetAttn2Execution,
*,
use: RegionalOperationMaskUse,
inputs: torch.Tensor,
spatial: bool,
) -> torch.Tensor:
"""Return one use mask in exact compact branch-segment order."""
if use.region_index >= int(execution.query_masks.shape[0]):
raise ValueError("Packed operation use references an unavailable region.")
source_gate = REGIONAL_OPERATION_BRANCH_GATE_RESOLVER.resolve(
execution.contexts,
branch=use.branch,
authority=inputs,
)
segments: list[torch.Tensor] = []
for segment in execution.branches.segments:
count = int(segment.source_indices.shape[0])
if segment.key.region_index != use.region_index:
segments.append(inputs.new_zeros((count, int(inputs.shape[1]), 1)))
continue
gate = source_gate.index_select(0, segment.source_indices).reshape(
count,
1,
1,
)
if spatial:
mask = execution.query_masks[use.region_index].index_select(
0,
segment.source_indices,
)
segments.append(mask.unsqueeze(-1) * gate)
else:
segments.append(gate.expand(-1, int(inputs.shape[1]), -1))
return torch.cat(tuple(segments))
STANDARD_UNET_PACKED_OPERATION_MASK_RESOLVER = StandardUnetPackedOperationMaskResolver()
@@ -68,8 +68,7 @@ class RegionalLoraConditioningSourceCollector:
if not isinstance(plan, RawRegionalAttentionPlan):
raise TypeError("Regional LoRA source collection requires a plan.")
self._require_unhooked_base("positive", plan.positive)
self._require_unhooked_base("negative", plan.negative)
self._require_compatible_base_hooks(plan)
return (
*self._branch_sources(plan.positive, branch=RegionalLoraBranch.POSITIVE),
*self._branch_sources(plan.negative, branch=RegionalLoraBranch.NEGATIVE),
@@ -112,31 +111,49 @@ class RegionalLoraConditioningSourceCollector:
)
return tuple(sources)
def _require_unhooked_base(
def _require_compatible_base_hooks(
self,
plan: RawRegionalAttentionPlan,
) -> None:
"""Require one shared global model-hook schedule across CFG branches."""
positive = self._base_hook_signature("positive", plan.positive)
negative = self._base_hook_signature("negative", plan.negative)
if positive != negative:
raise ValueError(
"Attention Coupling global model hooks must match across positive "
"and negative conditioning. Encode both branches through the same "
"Prompt Control global segment."
)
def _base_hook_signature(
self,
branch_name: str,
branch: RawRegionalAttentionBranch,
) -> None:
"""Require only model-active global LoRAs to arrive on the input MODEL."""
) -> tuple[tuple[object, ...], ...]:
"""Return one uniform global model-hook signature for a CFG branch."""
groups = conditioning_hook_groups(branch.base_conditioning)
model_hook_count = sum(
len(
signatures = tuple(
self._group_signature(
self._model_hook_selection(
group,
source_label=(
f"Attention Coupling {branch_name} global conditioning"
),
).model_hooks
)
)
for group in groups
for group in conditioning_hook_groups(branch.base_conditioning)
)
if model_hook_count:
if not signatures:
return ()
authority = signatures[0]
if any(signature != authority for signature in signatures[1:]):
raise ValueError(
f"Attention Coupling {branch_name} global conditioning contains "
"model hooks. Apply global LoRAs to the input MODEL; reserve "
"conditioning hooks for masked regional entries."
f"Attention Coupling {branch_name} global conditioning uses "
"different model HookGroups across text schedule entries. Keep "
"model LoRA scheduling on one shared WeightHook schedule."
)
return authority
def _uniform_hooks(
self,
@@ -26,6 +26,7 @@ from ..runtime.comfy_conditioning_processing import (
ComfyRegionalConditioningProcessor,
)
from ..runtime.comfy_latent_normalization import ComfyLatentNormalizer
from ..runtime.global_hook_model_resolver import GlobalHookModelResolver
from ..runtime.regional_lora_conditioning_adapter import (
RegionalLoraConditioningAdapter,
)
@@ -82,6 +83,9 @@ class AttentionCouplingModelPreparationService:
latent_normalizer_class: ClassVar[type[ComfyLatentNormalizer]] = (
ComfyLatentNormalizer
)
global_hook_model_resolver_class: ClassVar[type[GlobalHookModelResolver]] = (
GlobalHookModelResolver
)
model_family_selector_class: ClassVar[
type[AttentionCouplingModelFamilySelector]
] = AttentionCouplingModelFamilySelector
@@ -116,6 +120,11 @@ class AttentionCouplingModelPreparationService:
)
interop_validator = self.interop_validator_class()
interop_report = interop_validator.validate(model, capabilities)
model = self.global_hook_model_resolver_class().resolve(
model,
positive=positive,
negative=negative,
)
model_family = self.model_family_selector_class().select(capabilities)
samples = self.latent_normalizer_class().normalize(
model=model,
@@ -9,6 +9,7 @@ from __future__ import annotations
from typing import Any, TypeAlias
import torch
from comfy.hooks import HookGroup
from ..domain.conditioning_batch import ConditioningBatch
from ..domain.regional_prompting import (
@@ -19,6 +20,9 @@ from ..masking.regional_prompt_masks import (
prepare_regional_mask_batch,
regional_mask,
)
from ..runtime.global_first_conditioning_hooks import (
GLOBAL_FIRST_CONDITIONING_HOOK_COMPOSER,
)
from ..runtime.regional_conditioning_companion import detach_global_companion
from ..shared.logging import get_logger
@@ -114,6 +118,11 @@ class RegionalConditioningService:
if not plan.pairs:
return self._copy_conditioning(global_conditioning)
global_hooks = GLOBAL_FIRST_CONDITIONING_HOOK_COMPOSER.global_hooks(
global_conditioning,
source_label=f"{input_name} global conditioning",
)
hook_cache: dict[tuple[HookGroup, HookGroup], HookGroup] = {}
assembled = self._as_default(global_conditioning)
for pair in plan.pairs:
conditioning = self._validate_conditioning(
@@ -121,6 +130,22 @@ class RegionalConditioningService:
input_name=(f"{input_name} regional entry {pair.conditioning_index}"),
)
conditioning, global_companion = detach_global_companion(conditioning)
conditioning = GLOBAL_FIRST_CONDITIONING_HOOK_COMPOSER.compose(
conditioning,
global_hooks,
source_label=(f"{input_name} regional entry {pair.conditioning_index}"),
cache=hook_cache,
)
if global_companion is not None:
global_companion = GLOBAL_FIRST_CONDITIONING_HOOK_COMPOSER.compose(
global_companion,
global_hooks,
source_label=(
f"{input_name} regional entry {pair.conditioning_index} "
"global companion"
),
cache=hook_cache,
)
mask = regional_mask(mask_batch, pair.mask_index)
if global_companion is not None and regional_prompt_weight < 1.0:
assembled.extend(
@@ -27,9 +27,9 @@ from ..runtime.attention_coupling.unet_context import (
from ..runtime.regional_attention_diagnostics import (
RegionalAttentionDiagnosticsBuilder,
)
from ..runtime.regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
StandardUnetNativeLoraAdmissionService,
from ..runtime.regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
StandardUnetOperationPreparation,
)
from ..runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from ..runtime.regional_model_patch_interop import RegionalModelPatchInteropReport
@@ -47,8 +47,8 @@ class StandardUnetAttentionCouplingModelFamily:
backend_class: ClassVar[type[StandardUnetAttentionBackend]] = (
StandardUnetAttentionBackend
)
native_admission_class: ClassVar[type[StandardUnetNativeLoraAdmissionService]] = (
StandardUnetNativeLoraAdmissionService
operation_preparation_class: ClassVar[type[StandardUnetOperationPreparation]] = (
StandardUnetOperationPreparation
)
@property
@@ -84,7 +84,7 @@ class StandardUnetAttentionCouplingModelFamily:
raise TypeError(
"Standard UNet Attention Coupling requires regional adaptation."
)
return self.native_admission_class().admit(model, adaptation)
return self.operation_preparation_class().admit(model, adaptation)
def prepare_sampler_conditioning(
self,
@@ -113,8 +113,8 @@ class StandardUnetAttentionCouplingModelFamily:
) -> object:
"""Build shared diagnostics state and derive the paired attn2 backend."""
if not isinstance(admission, StandardUnetNativeLoraAdmission):
raise TypeError("Standard UNet derivation requires native admission.")
if not isinstance(admission, StandardUnetOperationAdmission):
raise TypeError("Standard UNet derivation requires operation admission.")
if not isinstance(interop_report, RegionalModelPatchInteropReport):
raise TypeError("Standard UNet derivation requires interop evidence.")
if admission.adaptation.plan != processed_plan.lora_plan:
-219
View File
@@ -1,219 +0,0 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify exact static-global and regional Anima LoRA overlap rejection."""
from __future__ import annotations
from types import SimpleNamespace
import pytest
import torch
from comfy.weight_adapter.lora import LoRAAdapter
from simple_syrup.domain.regional_lora_plan import (
RegionalLoraAdapterIdentity,
RegionalLoraAdapterPlan,
RegionalLoraBranch,
RegionalLoraPlan,
RegionalLoraScheduleBoundary,
)
from simple_syrup.runtime.regional_lora.anima_global_lora_overlap import (
AnimaGlobalRegionalLoraOverlapError,
AnimaGlobalRegionalLoraOverlapValidator,
)
from simple_syrup.runtime.regional_lora.anima_plan_admission import (
AnimaRegionalLoraAdapterAdmission,
AnimaRegionalLoraPlanAdmission,
)
from simple_syrup.runtime.regional_lora.anima_targets import (
AnimaLoraAdmission,
AnimaLoraTarget,
AnimaLoraTargetFamily,
anima_lora_target_name,
expected_anima_lora_features,
)
from simple_syrup.runtime.regional_lora.standard_adapter import StandardLoraTarget
def test_overlap_validator_accepts_empty_distinct_and_partial_global_state() -> None:
"""Preserve unpatched and content-distinct global MODEL LoRAs."""
admission, targets = _admission(target_count=2)
validator = AnimaGlobalRegionalLoraOverlapValidator()
validator.validate(SimpleNamespace(patches={}), admission)
validator.validate(
SimpleNamespace(patches={_key(targets[0]): [_patch(targets[0])]}),
admission,
)
changed = targets[0].up.clone()
changed[0, 0] += 1.0
validator.validate(
SimpleNamespace(
patches={
_key(targets[0]): [_patch(targets[0], up=changed)],
_key(targets[1]): [_patch(targets[1])],
}
),
admission,
)
@pytest.mark.parametrize("clone_tensors", [False, True], ids=("identity", "content"))
def test_overlap_validator_rejects_complete_exact_global_content(
clone_tensors: bool,
) -> None:
"""Reject exact adapter content with identity and cloned-tensor paths."""
admission, targets = _admission(target_count=2)
patches = {
_key(target): [
_patch(
target,
down=target.down.clone() if clone_tensors else target.down,
up=target.up.clone() if clone_tensors else target.up,
)
]
for target in targets
}
with pytest.raises(
AnimaGlobalRegionalLoraOverlapError,
match="already applied globally.*regional.safetensors",
):
AnimaGlobalRegionalLoraOverlapValidator().validate(
SimpleNamespace(patches=patches),
admission,
)
@pytest.mark.parametrize(
"patch_factory",
(
lambda target: _patch(target, strength=0.0),
lambda target: (1.0, ("diff", (target.up,)), 1.0, None, None),
lambda target: _patch(target, alpha=1.0),
),
ids=("zero-strength", "non-lora", "alpha-form"),
)
def test_overlap_validator_preserves_noncomparable_global_patches(
patch_factory: object,
) -> None:
"""Keep inactive and other global patch formats unchanged."""
admission, targets = _admission(target_count=1)
factory = patch_factory
assert callable(factory)
AnimaGlobalRegionalLoraOverlapValidator().validate(
SimpleNamespace(patches={_key(targets[0]): [factory(targets[0])]}),
admission,
)
@pytest.mark.parametrize(
("patches", "message"),
(
(None, "requires MODEL patches"),
({"diffusion_model.blocks.0.self_attn.q_proj.weight": object()}, "list"),
({"diffusion_model.blocks.0.self_attn.q_proj.weight": [object()]}, "tuple"),
(
{
"diffusion_model.blocks.0.self_attn.q_proj.weight": [
(float("nan"), object(), 1.0)
]
},
"finite",
),
),
)
def test_overlap_validator_fails_closed_on_malformed_model_patch_state(
patches: object,
message: str,
) -> None:
"""Reject installed-host patch drift before regional execution setup."""
admission, _ = _admission(target_count=1)
with pytest.raises((TypeError, ValueError), match=message):
AnimaGlobalRegionalLoraOverlapValidator().validate(
SimpleNamespace(patches=patches),
admission,
)
def _admission(
*, target_count: int
) -> tuple[AnimaRegionalLoraPlanAdmission, tuple[StandardLoraTarget, ...]]:
"""Build one admitted regional adapter with small valid Anima targets."""
families = tuple(AnimaLoraTargetFamily)[:target_count]
targets = tuple(_target(family) for family in families)
plan_entry = RegionalLoraAdapterPlan(
adapter_identity=RegionalLoraAdapterIdentity("regional.safetensors"),
composition_index=0,
region_index=0,
branch=RegionalLoraBranch.POSITIVE,
model_strength=0.8,
schedule=(RegionalLoraScheduleBoundary(0.0, 1.0, 1.0, 0),),
)
plan = RegionalLoraPlan((plan_entry,))
admission = AnimaLoraAdmission(
tuple(
AnimaLoraTarget(0, family, target)
for family, target in zip(families, targets, strict=True)
)
)
return (
AnimaRegionalLoraPlanAdmission(
plan,
(AnimaRegionalLoraAdapterAdmission(plan_entry, admission),),
),
targets,
)
def _target(family: AnimaLoraTargetFamily) -> StandardLoraTarget:
"""Return one rank-one target with installed Anima feature dimensions."""
input_features, output_features = expected_anima_lora_features(family)
return StandardLoraTarget(
target=anima_lora_target_name(0, family),
down=torch.arange(input_features, dtype=torch.float32).reshape(1, -1),
up=torch.arange(output_features, dtype=torch.float32).reshape(-1, 1),
rank=1,
input_features=input_features,
output_features=output_features,
)
def _key(target: StandardLoraTarget) -> str:
"""Return the installed Comfy MODEL patch key for one target."""
return f"{target.target}.weight"
def _patch(
target: StandardLoraTarget,
*,
strength: float = 1.0,
down: torch.Tensor | None = None,
up: torch.Tensor | None = None,
alpha: float | None = None,
) -> tuple[object, ...]:
"""Return one installed-Comfy standard static LoRA patch entry."""
adapter = LoRAAdapter(
set(),
(
target.up if up is None else up,
target.down if down is None else down,
alpha,
None,
None,
None,
),
)
return (strength, adapter, 1.0, None, None)
@@ -140,6 +140,40 @@ def test_positive_conditioning_batch_selects_by_segment_index() -> None:
assert [call.negative for call in sampler.sample_calls] == [negative, negative]
def test_detailer_keeps_prompt_control_hooks_peer_scoped_by_segment() -> None:
"""Preserve each scheduled LoRA hook on only its selected face conditioning."""
sampler = _FakeSampler()
first = _segment(CropRegion(0, 0, 4, 4), BoundingBox(1, 1, 3, 3))
second = _segment(CropRegion(4, 4, 8, 8), BoundingBox(5, 5, 7, 7))
first_hooks = object()
second_hooks = object()
first_conditioning = [["first", {"hooks": first_hooks}]]
second_conditioning = [["second", {"hooks": second_hooks}]]
service = _service(sampler)
service.detail(
_image(),
_segs(first, second),
object(),
object(),
ConditioningBatch((first_conditioning, second_conditioning)),
[],
**_settings(),
)
assert sampler.sample_calls[0].positive is first_conditioning
assert sampler.sample_calls[1].positive is second_conditioning
selected_first = cast(list[list[object]], sampler.sample_calls[0].positive)
selected_second = cast(list[list[object]], sampler.sample_calls[1].positive)
first_metadata = selected_first[0][1]
second_metadata = selected_second[0][1]
assert isinstance(first_metadata, dict)
assert isinstance(second_metadata, dict)
assert first_metadata["hooks"] is first_hooks
assert second_metadata["hooks"] is second_hooks
def test_negative_conditioning_batch_selects_by_segment_index() -> None:
"""A negative batch varies by SEG while normal positive broadcasts."""
+94
View File
@@ -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 pre-derivation Comfy MODEL selection for global prompt hooks."""
from __future__ import annotations
from comfy.hooks import HookGroup
from simple_syrup.domain.conditioning_batch import ConditioningBatch
from simple_syrup.runtime.global_hook_model_resolver import GlobalHookModelResolver
class _Model:
"""Expose the dynamic-model boundary used by Comfy's CFG guider."""
def __init__(self, *, dynamic: bool, delegate: object | None = None) -> None:
"""Retain configured dynamic state and delegate result."""
self.dynamic = dynamic
self.delegate = delegate
self.delegate_calls = 0
def is_dynamic(self) -> bool:
"""Return the configured Comfy model mode."""
return self.dynamic
def get_non_dynamic_delegate(self) -> object:
"""Return and record the configured static delegate."""
self.delegate_calls += 1
return self.delegate
def test_unhooked_global_entry_preserves_dynamic_model() -> None:
"""Leave regional-only hooks to the existing custom regional runtime."""
model = _Model(dynamic=True)
conditioning = ConditioningBatch((_conditioning(), _conditioning(HookGroup())))
resolved = GlobalHookModelResolver().resolve(
model,
positive=conditioning,
negative=conditioning,
)
assert resolved is model
assert model.delegate_calls == 0
def test_global_hook_preserves_already_static_model() -> None:
"""Avoid unnecessary model replacement when hook execution is already static."""
model = _Model(dynamic=False)
resolved = GlobalHookModelResolver().resolve(
model,
positive=_conditioning(HookGroup()),
negative=_conditioning(HookGroup()),
)
assert resolved is model
assert model.delegate_calls == 0
def test_global_hook_selects_static_delegate_before_regional_derivation() -> None:
"""Bind regional wrappers to the same static graph Comfy samples with hooks."""
delegate = _Model(dynamic=False)
model = _Model(dynamic=True, delegate=delegate)
resolved = GlobalHookModelResolver().resolve(
model,
positive=ConditioningBatch(
(_conditioning(HookGroup()), _conditioning(HookGroup()))
),
negative=ConditioningBatch(
(_conditioning(HookGroup()), _conditioning(HookGroup()))
),
)
assert resolved is delegate
assert model.delegate_calls == 1
def _conditioning(hooks: HookGroup | None = None) -> list[list[object]]:
"""Return one standard conditioning with optional model hooks."""
metadata: dict[str, object] = {}
if hooks is not None:
metadata["hooks"] = hooks
return [["embedding", metadata]]
-7
View File
@@ -29,13 +29,6 @@ def test_matrix_covers_global_strength_regional_and_duplicate_placement() -> Non
1.0,
1.0,
]
assert [case.expect_overlap_rejection for case in definitions] == [
False,
False,
False,
False,
True,
]
def test_prompt_renderer_places_primary_adapter_only_in_declared_segments() -> None:
+143
View File
@@ -0,0 +1,143 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify exact user-prompt global LoRA proof graph mutations."""
from __future__ import annotations
from tools.run_global_prompt_lora_proof import (
ProofCase,
build_case_graph,
cases,
render_global_first_prompt,
)
def test_prompt_renderer_preserves_named_separators_and_places_both_scopes() -> None:
"""Keep global-first layout while preserving the user's named regions."""
rendered = render_global_first_prompt(
"global text[SEP|Taffy]left text[SEP|Anise]right text",
global_tag="<lora:shared:0.5>",
regional_tag="<lora:shared:0.8>",
)
assert rendered == (
"<lora:shared:0.5>\nglobal text"
"[SEP|Taffy]<lora:shared:0.8>\nleft text"
"[SEP|Anise]right text"
)
def test_matrix_covers_same_lora_turbo_tiled_and_contextual() -> None:
"""Keep all requested managed proof variants explicit and ordered."""
definitions = cases()
assert [case.case_id for case in definitions] == [
"sdxl-global-and-regional-same-lora",
"anima-global-and-regional-same-lora",
"anima-turbo-global-arcane-regional-tiled",
"anima-turbo-global-arcane-regional-contextual",
]
assert definitions[0].global_tag == definitions[0].regional_tag
assert "ArcaneViolet" in definitions[1].global_tag
assert "ArcaneViolet" in definitions[1].regional_tag
assert [case.turbo for case in definitions] == [False, False, True, True]
assert [case.contextual for case in definitions] == [False, False, False, True]
assert definitions[0].region_mask_feather == 10
assert {case.region_mask_feather for case in definitions[1:]} == {64}
def test_turbo_contextual_graph_uses_appropriate_sampling_contract() -> None:
"""Apply Turbo's low-step CFG-one contract to full and contextual stages."""
template = _anima_template()
case = ProofCase(
"fixture",
"anima",
"<lora:turbo:0.7>",
"<lora:regional:0.8>",
turbo=True,
contextual=True,
)
graph, save_ids = build_case_graph(template, case, run_id="run")
assert save_ids == ("proof:source", "proof:refinement")
for node_id in (
"anima-prompt-region:ksampler",
"anima-diffusion-upscale:ksampler",
):
inputs = graph[node_id]["inputs"]
assert isinstance(inputs, dict)
assert inputs["steps"] == 10
assert inputs["cfg"] == 1.0
assert inputs["sampler_name"] == "euler"
assert inputs["scheduler"] == "simple"
assert inputs["region_mask_feather"] == 64
refinement = graph["anima-diffusion-upscale:ksampler"]
assert refinement["class_type"] == "SimpleSyrup.KSamplerAttentionCouplingContextual"
refinement_inputs = refinement["inputs"]
assert isinstance(refinement_inputs, dict)
assert "latent_context_size" in refinement_inputs
assert "latent_tile_width" not in refinement_inputs
def test_replayed_multiselect_mask_values_are_literal_wrapped() -> None:
"""Keep executed multiselect lists from being reinterpreted as graph links."""
template = _anima_template()
template["mask-loader"] = {
"class_type": "SimpleSyrup.LoadMaskBatch",
"inputs": {"image": ["left.png", "right.png"], "channel": "red"},
}
case = ProofCase("fixture", "anima", "<lora:g:1>", "<lora:r:1>")
graph, _save_ids = build_case_graph(template, case, run_id="run")
assert graph["mask-loader"]["inputs"] == {
"image": {"__value__": ["left.png", "right.png"]},
"channel": "red",
}
def _anima_template() -> dict[str, dict[str, object]]:
"""Return one minimal expanded Anima graph accepted by the mutator."""
return {
"anima-prompt-region:positive_prompt": {
"class_type": "PrimitiveStringMultiline",
"inputs": {"value": "global[SEP|Taffy]left[SEP|Anise]right"},
},
"anima-diffusion-upscale:positive_prompt": {
"class_type": "PrimitiveStringMultiline",
"inputs": {"value": "global[SEP|Taffy]left[SEP|Anise]right"},
},
"anima-prompt-region:ksampler": {
"class_type": "SimpleSyrup.KSamplerAttentionCoupling",
"inputs": {},
},
"anima-diffusion-upscale:ksampler": {
"class_type": "SimpleSyrup.KSamplerAttentionCouplingTiled",
"inputs": {
"latent_tile_width": 128,
"latent_tile_height": 128,
"latent_tile_overlap": 16,
"latent_tile_batch_size": 4,
},
},
"anima-prompt-region:vae_decode": {
"class_type": "VAEDecode",
"inputs": {},
},
"anima-diffusion-upscale:vae_decode": {
"class_type": "VAEDecode",
"inputs": {},
},
"__sugarcubes_cube_output__:fixture": {
"class_type": "SugarCubes.CubeOutput",
"inputs": {},
},
}
+4 -29
View File
@@ -49,19 +49,12 @@ def test_p9_3_matrix_covers_global_regional_distinct_and_duplicate_cases() -> No
0,
1,
]
assert [case.expect_overlap_error for case in definitions] == [
False,
False,
False,
False,
True,
]
def test_p9_3_result_requires_success_outputs_and_exact_overlap_rejection(
def test_p9_3_result_requires_success_outputs_for_every_placement(
tmp_path: Path,
) -> None:
"""Persist four images and one pre-sampling rejection before completion."""
"""Persist all five global, regional, and additive placement images."""
recorder = GlobalRegionalLoraResultRecorder(tmp_path)
workflow = BuiltAnimaAttentionCouplingWorkflow(
@@ -77,7 +70,7 @@ def test_p9_3_result_requires_success_outputs_and_exact_overlap_rejection(
"status": {"status_str": "success", "completed": True},
"outputs": {"metrics": {"benchmark_metrics": [{"model_call_count": STEPS}]}},
}
for case in definitions[:-1]:
for case in definitions:
color = (
(240, 10, 10)
if "global-global_adapter-regional" in case.case_id
@@ -93,24 +86,6 @@ def test_p9_3_result_requires_success_outputs_and_exact_overlap_rejection(
reference=ImageReference("image.png", "", "output"),
image_bytes=_png(color),
)
rejection_history: JsonObject = {
"status": {
"status_str": "error",
"completed": False,
"messages": [
"Regional Anima LoRA content is already applied globally to the "
"input MODEL: 'adapter-a.safetensors'."
],
},
"outputs": {},
}
recorder.record_overlap_rejection(
definitions[-1],
workflow,
prompt_id="prompt-duplicate",
history=rejection_history,
)
result_path = recorder.finalize(
definitions,
system_stats={"devices": []},
@@ -125,7 +100,7 @@ def test_p9_3_result_requires_success_outputs_and_exact_overlap_rejection(
"success",
"success",
"success",
"rejected_before_sampling",
"success",
]
assert result["transition"] == {
"distinct_before_after": {
@@ -152,8 +152,10 @@ def test_native_sampler_activates_hooks_from_masked_regional_conditioning(
) -> None:
"""Comfy activates each preserved hook group through direct and tiled paths."""
global_hooks = comfy.hooks.HookGroup()
regional_hooks = comfy.hooks.HookGroup()
global_hooks = comfy.hooks.create_hook_lora({}, 0.7, 0.0)
regional_hooks = comfy.hooks.create_hook_lora({}, 0.9, 0.0)
global_hooks.get_type(comfy.hooks.EnumHookType.Weight)[0].hook_ref = "global"
regional_hooks.get_type(comfy.hooks.EnumHookType.Weight)[0].hook_ref = "regional"
assembled, _ = RegionalConditioningService().assemble(
positive=ConditioningBatch(
(
@@ -166,6 +168,12 @@ def test_native_sampler_activates_hooks_from_masked_regional_conditioning(
regional_prompt_weight=0.5,
region_mask_feather=0,
)
combined_hooks = assembled[1][1]["hooks"]
assert isinstance(combined_hooks, comfy.hooks.HookGroup)
assert [
hook.hook_ref
for hook in combined_hooks.get_type(comfy.hooks.EnumHookType.Weight)
] == ["global", "regional"]
converted = comfy.sampler_helpers.convert_cond(assembled)
for conditioning in converted:
cross_attn = conditioning.pop("cross_attn")
@@ -187,9 +195,9 @@ def test_native_sampler_activates_hooks_from_masked_regional_conditioning(
)
assert len(outputs) == 1
assert set(model.current_patcher.prepared) == {global_hooks, regional_hooks}
assert set(model.current_patcher.applied) == {global_hooks, regional_hooks}
assert set(model.model_calls) == {global_hooks, regional_hooks}
assert set(model.current_patcher.prepared) == {global_hooks, combined_hooks}
assert set(model.current_patcher.applied) == {global_hooks, combined_hooks}
assert set(model.model_calls) == {global_hooks, combined_hooks}
def _conditioning(
+78 -4
View File
@@ -8,6 +8,7 @@ from __future__ import annotations
import pytest
import torch
from comfy.hooks import EnumHookType, HookGroup, create_hook_lora
from simple_syrup.domain.conditioning_batch import ConditioningBatch
from simple_syrup.runtime.regional_conditioning_companion import (
@@ -141,10 +142,10 @@ def test_feathering_preserves_inputs_and_softens_regional_copy() -> None:
def test_mask_composition_preserves_segment_lora_hook_metadata() -> None:
"""Global and regional hook groups survive standard mask composition."""
"""Global hooks compose into every conventional regional model state."""
global_hooks = object()
regional_hooks = object()
global_hooks = _hooks("global")
regional_hooks = _hooks("regional")
positive = ConditioningBatch(
(
[["global", {"hooks": global_hooks, "other": "global metadata"}]],
@@ -162,10 +163,68 @@ def test_mask_composition_preserves_segment_lora_hook_metadata() -> None:
assert assembled[0][1]["hooks"] is global_hooks
assert assembled[0][1]["other"] == "global metadata"
assert assembled[1][1]["hooks"] is regional_hooks
combined = assembled[1][1]["hooks"]
assert isinstance(combined, HookGroup)
assert _hook_refs(combined) == ["global", "regional"]
assert assembled[1][1]["other"] == "region metadata"
def test_global_hooks_compose_with_regional_companion_and_both_cfg_sides() -> None:
"""Keep one global patch under every local and fallback regional prompt share."""
positive_global = _hooks("positive global")
positive_regional = _hooks("positive regional")
negative_global = _hooks("negative global")
negative_regional = _hooks("negative regional")
positive_region = attach_global_companion(
[["positive region", {"hooks": positive_regional}]],
[["positive fallback", {"hooks": positive_regional}]],
)
negative_region = attach_global_companion(
[["negative region", {"hooks": negative_regional}]],
[["negative fallback", {"hooks": negative_regional}]],
)
positive, negative = RegionalConditioningService().assemble(
positive=ConditioningBatch(
(
[["positive global", {"hooks": positive_global}]],
positive_region,
)
),
negative=ConditioningBatch(
(
[["negative global", {"hooks": negative_global}]],
negative_region,
)
),
masks=torch.ones((1, 2, 2)),
regional_prompt_weight=0.5,
region_mask_feather=0,
)
assert _hook_refs(positive[0][1]["hooks"]) == ["positive global"]
assert _hook_refs(positive[1][1]["hooks"]) == [
"positive global",
"positive regional",
]
assert _hook_refs(positive[2][1]["hooks"]) == [
"positive global",
"positive regional",
]
assert positive[1][1]["hooks"] is positive[2][1]["hooks"]
assert _hook_refs(negative[0][1]["hooks"]) == ["negative global"]
assert _hook_refs(negative[1][1]["hooks"]) == [
"negative global",
"negative regional",
]
assert _hook_refs(negative[2][1]["hooks"]) == [
"negative global",
"negative regional",
]
assert negative[1][1]["hooks"] is negative[2][1]["hooks"]
@pytest.mark.parametrize(
("regional_prompt_weight", "expected_sources", "expected_strengths"),
[
@@ -281,3 +340,18 @@ def test_invalid_regional_prompt_weight_fails_before_mask_processing(
regional_prompt_weight=weight,
region_mask_feather=0,
)
def _hooks(identity: str) -> HookGroup:
"""Return one recognizable model-active Prompt Control-style HookGroup."""
hooks = create_hook_lora({}, strength_model=1.0, strength_clip=0.0)
hooks.get_type(EnumHookType.Weight)[0].hook_ref = identity
return hooks
def _hook_refs(value: object) -> list[object]:
"""Return ordered WeightHook references from one asserted HookGroup."""
assert isinstance(value, HookGroup)
return [hook.hook_ref for hook in value.get_type(EnumHookType.Weight)]
@@ -22,6 +22,9 @@ from simple_syrup.masking.regional_prompt_masks import build_regional_mask_bank
from simple_syrup.runtime.regional_lora_conditioning_adapter import (
RegionalLoraConditioningAdapter,
)
from simple_syrup.runtime.regional_lora_conditioning_sources import (
conditioning_hook_groups,
)
class _Sampling:
@@ -201,17 +204,32 @@ def test_adapter_accepts_different_text_only_hooks_across_schedule_entries() ->
assert adaptation.adapter_payloads == ()
def test_adapter_rejects_global_hooks_and_opaque_regional_identity() -> None:
"""Require global MODEL ownership and a stable regional adapter identity."""
def test_adapter_preserves_global_hooks_and_rejects_opaque_regional_identity() -> None:
"""Leave compatible global hooks on sampler conditioning and adapt regions."""
hooks = _hooks(("pc-PRIMARY_ADAPTER-0.8-0.0",))
global_hooks = _hooks(("pc-GLOBAL_ADAPTER-0.7-0.0",))
regional_hooks = _hooks(("pc-PRIMARY_ADAPTER-0.8-0.0",))
global_plan = build_raw_regional_attention_plan(
positive=ConditioningBatch((_conditioning(hooks), _conditioning())),
negative=ConditioningBatch((_conditioning(), _conditioning())),
positive=ConditioningBatch(
(_conditioning(global_hooks), _conditioning(regional_hooks))
),
negative=ConditioningBatch((_conditioning(global_hooks), _conditioning())),
mask_bank=_mask_bank(1),
)
with pytest.raises(ValueError, match="Apply global LoRAs to the input MODEL"):
RegionalLoraConditioningAdapter().adapt(global_plan, model=_Model())
adaptation = RegionalLoraConditioningAdapter().adapt(
global_plan,
model=_Model(),
)
assert [item.adapter_identity.value for item in adaptation.plan.adapters] == [
"pc-PRIMARY_ADAPTER-0.8-0.0"
]
assert conditioning_hook_groups(global_plan.positive.base_conditioning) == (
global_hooks,
)
assert conditioning_hook_groups(global_plan.negative.base_conditioning) == (
global_hooks,
)
opaque = comfy.hooks.create_hook_lora({}, 1.0, 0.0)
opaque_plan = build_raw_regional_attention_plan(
@@ -223,6 +241,21 @@ def test_adapter_rejects_global_hooks_and_opaque_regional_identity() -> None:
RegionalLoraConditioningAdapter().adapt(opaque_plan, model=_Model())
def test_adapter_rejects_mismatched_global_cfg_hook_schedules() -> None:
"""Keep Attention Coupling on one global hook state per packed model call."""
positive_global = _hooks(("pc-positive-global",))
negative_global = _hooks(("pc-negative-global",))
plan = build_raw_regional_attention_plan(
positive=ConditioningBatch((_conditioning(positive_global), _conditioning())),
negative=ConditioningBatch((_conditioning(negative_global), _conditioning())),
mask_bank=_mask_bank(1),
)
with pytest.raises(ValueError, match="must match across positive and negative"):
RegionalLoraConditioningAdapter().adapt(plan, model=_Model())
def _hooks(identities: tuple[str, ...]) -> comfy.hooks.HookGroup:
"""Return ordered WeightHooks with Prompt Control-style stable refs."""
@@ -0,0 +1,110 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify regional LoRA ownership inside compact standard-UNet attn2 rows."""
from __future__ import annotations
from dataclasses import dataclass
import torch
from simple_syrup.domain.regional_activation_geometry import RegionalActivationLayout
from simple_syrup.domain.regional_attention import RegionalAttentionBranch
from simple_syrup.domain.regional_attention_batch import (
BatchedRegionalAttentionContexts,
BatchedRegionalAttentionEntry,
BatchedRegionalAttentionRegion,
RegionalAttentionChunkBatch,
)
from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch
from simple_syrup.runtime.attention_coupling.unet_attn2_execution import (
UnetAttn2Execution,
)
from simple_syrup.runtime.regional_lora.standard_unet_packed_operation_masks import (
StandardUnetPackedOperationMaskResolver,
)
@dataclass(frozen=True, slots=True)
class _Use:
"""Expose operation ownership for one synthetic target use."""
composition_index: int
region_index: int
branch: RegionalLoraBranch
def test_packed_image_masks_exclude_base_and_preserve_cfg_ownership() -> None:
"""Apply spatial masks only to matching regional query and output rows."""
execution = _execution()
masks = StandardUnetPackedOperationMaskResolver().resolve_image_tokens(
execution,
uses=(
_Use(0, 0, RegionalLoraBranch.POSITIVE),
_Use(1, 0, RegionalLoraBranch.NEGATIVE),
),
inputs=torch.zeros((4, 2, 8)),
)
assert masks.geometry.layout is RegionalActivationLayout.CONSUMER_SPATIALIZED
assert masks.multipliers.shape == (2, 4, 2, 1)
torch.testing.assert_close(
masks.multipliers[0, :, :, 0],
torch.tensor([[0.0, 0.0], [0.0, 0.0], [1.0, 0.5], [0.0, 0.0]]),
)
torch.testing.assert_close(
masks.multipliers[1, :, :, 0],
torch.tensor([[0.0, 0.0], [0.0, 0.0], [0.0, 0.0], [0.25, 1.0]]),
)
def test_packed_context_masks_gate_rows_without_reshaping_tokens() -> None:
"""Broadcast branch ownership over the untouched context-token sequence."""
execution = _execution()
masks = StandardUnetPackedOperationMaskResolver().resolve_context_tokens(
execution,
uses=(
_Use(0, 0, RegionalLoraBranch.POSITIVE),
_Use(1, 0, RegionalLoraBranch.NEGATIVE),
),
inputs=torch.zeros((4, 77, 64)),
)
assert masks.geometry.layout is RegionalActivationLayout.BRANCH_TOKENS
assert masks.geometry.invocation_shape == (4, 77, 64)
assert masks.multipliers.shape == (2, 4, 77, 1)
assert masks.multipliers[0, :2].eq(0.0).all()
assert masks.multipliers[0, 2].eq(1.0).all()
assert masks.multipliers[0, 3].eq(0.0).all()
assert masks.multipliers[1, 3].eq(1.0).all()
def _execution() -> UnetAttn2Execution:
"""Return two CFG rows and one compact regional branch over both rows."""
base = torch.zeros((2, 77, 64))
contexts = BatchedRegionalAttentionContexts(
latent_batch_size=1,
chunks=(
RegionalAttentionChunkBatch(0, RegionalAttentionBranch.POSITIVE, 0, 1),
RegionalAttentionChunkBatch(1, RegionalAttentionBranch.NEGATIVE, 1, 2),
),
base_context=base,
regions=(
BatchedRegionalAttentionRegion(
0,
(BatchedRegionalAttentionEntry(0, base.clone(), (1.0, 1.0)),),
),
),
)
return UnetAttn2Execution(
contexts,
torch.tensor([[[1.0, 0.5], [0.25, 1.0]]]),
(0.75,),
1,
2,
)
@@ -0,0 +1,399 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Prove spatial standard-UNet LoRA execution keeps one model trajectory."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, cast
from uuid import UUID
import comfy.model_patcher
import pytest
import torch
from comfy.patcher_extension import WrapperExecutor, WrappersMP
from comfy.weight_adapter.lora import LoRAAdapter
from torch import nn
from simple_syrup.domain.conditioning_schedule import ConditioningScheduleRange
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionBranch,
ProcessedRegionalAttentionContext,
ProcessedRegionalAttentionEntry,
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.regional_lora_plan import (
RegionalLoraAdapterIdentity,
RegionalLoraAdapterPlan,
RegionalLoraBranch,
RegionalLoraPlan,
RegionalLoraScheduleBoundary,
)
from simple_syrup.domain.regional_mask_bank import RegionalMaskBank
from simple_syrup.domain.spatial_views import (
SpatialBatchLayout,
SpatialView,
SpatialViewKind,
)
from simple_syrup.runtime.attention_coupling.unet import StandardUnetAttentionBackend
from simple_syrup.runtime.attention_coupling.unet_attention_state import (
StandardUnetAttentionState,
)
from simple_syrup.runtime.regional_attention_diagnostics import (
RegionalAttentionDiagnosticsBuilder,
)
from simple_syrup.runtime.regional_lora.comfy_adapter_resolution import (
ComfyAdapterTargetPath,
ComfyNormalizedAdapterTarget,
ComfyRegionalAdapterResolution,
ComfyRegionalLoraResolution,
)
from simple_syrup.runtime.regional_lora.execution_cache import (
RegionalLoraExecutionCache,
)
from simple_syrup.runtime.regional_lora.resolved_operation_translator import (
COMFY_RESOLVED_OPERATION_TRANSLATOR,
)
from simple_syrup.runtime.regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
)
from simple_syrup.runtime.regional_lora.target_binder import (
REGIONAL_LORA_TARGET_BINDER,
)
from simple_syrup.runtime.regional_lora.target_binding import (
BoundRegionalLoraSpatialCapability,
)
from simple_syrup.runtime.regional_lora_host_payload import RegionalLoraHostPayload
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from simple_syrup.runtime.spatial_model_arguments import (
SIMPLE_SYRUP_TRANSFORMER_NAMESPACE,
SPATIAL_BATCH_LAYOUT_KEY,
)
@dataclass(frozen=True, slots=True)
class _Mode:
"""Describe one exact model-call layout and expected regional ownership."""
input_shape: tuple[int, int, int, int]
layout: SpatialBatchLayout | None
expected_multiplier: torch.Tensor
class _OperationDiffusion(nn.Module):
"""Execute one spatial Linear target and count denoiser trajectories."""
def __init__(self) -> None:
"""Install one identity target and an empty call journal."""
super().__init__()
self.linear = nn.Linear(2, 2, bias=False)
self.linear.weight = nn.Parameter(torch.eye(2), requires_grad=False)
self.calls = 0
def forward(
self,
model_input: torch.Tensor,
timestep: torch.Tensor,
context: torch.Tensor,
y: object,
control: object,
transformer_options: dict[str, Any],
) -> torch.Tensor:
"""Run one B/S/C projection over the current activation grid."""
del timestep, context, y, control
self.calls += 1
transformer_options["activations_shape"] = list(model_input.shape)
tokens = model_input.permute(0, 2, 3, 1).reshape(
model_input.shape[0],
-1,
model_input.shape[1],
)
projected = cast(torch.Tensor, self.linear(tokens))
return projected.reshape(
model_input.shape[0],
model_input.shape[2],
model_input.shape[3],
model_input.shape[1],
).permute(0, 3, 1, 2)
@pytest.mark.parametrize("mode_id", ["full", "tiled", "contextual"])
def test_regional_operation_spatializes_one_shared_model_call(mode_id: str) -> None:
"""Apply the regional delta in place without sampling independent variants."""
mode = {
"full": _full_mode,
"tiled": _tiled_mode,
"contextual": _contextual_mode,
}[mode_id]()
source, diffusion = _patcher()
plan = _lora_plan()
admission = _admission(source, plan)
state = _state(plan)
built = StandardUnetAttentionBackend().derive(
model=source,
state=state,
admission=admission,
)
derived = built.model
assert isinstance(derived, comfy.model_patcher.ModelPatcher)
original = diffusion.linear
model_input = torch.ones(mode.input_shape)
context = torch.zeros((mode.input_shape[0], 2, 2))
options: dict[str, Any] = {
"cond_or_uncond": [0] * mode.input_shape[0],
"sigmas": torch.ones(mode.input_shape[0]),
"sample_sigmas": torch.tensor([1.0, 0.0]),
"wrappers": derived.wrappers,
}
if mode.layout is not None:
options[SIMPLE_SYRUP_TRANSFORMER_NAMESPACE] = {
SPATIAL_BATCH_LAYOUT_KEY: mode.layout
}
derived.patch_model(load_weights=False)
try:
output = WrapperExecutor.new_class_executor(
diffusion.forward,
diffusion,
derived.get_all_wrappers(WrappersMP.DIFFUSION_MODEL),
).execute(
model_input,
torch.ones(mode.input_shape[0]),
context,
None,
None,
options,
)
finally:
derived.unpatch_model(unpatch_weights=False)
torch.testing.assert_close(output, mode.expected_multiplier.expand_as(output))
assert diffusion.calls == 1
assert diffusion.linear is original
assert source.object_patches == {}
assert admission.cache is not None
derived.detach(unpatch_all=True)
assert admission.cache.size == 0
def test_regional_delta_composes_with_an_already_global_patched_weight() -> None:
"""Add the same eligible LoRA regionally over its global native weight."""
source, diffusion = _patcher(base_scale=2.0)
plan = _lora_plan()
admission = _admission(source, plan)
state = _state(plan)
derived = (
StandardUnetAttentionBackend()
.derive(
model=source,
state=state,
admission=admission,
)
.model
)
assert isinstance(derived, comfy.model_patcher.ModelPatcher)
options: dict[str, Any] = {
"cond_or_uncond": [0],
"sigmas": torch.ones(1),
"sample_sigmas": torch.tensor([1.0, 0.0]),
"wrappers": derived.wrappers,
}
derived.patch_model(load_weights=False)
try:
output = WrapperExecutor.new_class_executor(
diffusion.forward,
diffusion,
derived.get_all_wrappers(WrappersMP.DIFFUSION_MODEL),
).execute(
torch.ones((1, 2, 2, 4)),
torch.ones(1),
torch.zeros((1, 2, 2)),
None,
None,
options,
)
finally:
derived.unpatch_model(unpatch_weights=False)
expected = torch.tensor([[[[2.5, 2.5, 2.0, 2.0]]]])
torch.testing.assert_close(output, expected.expand_as(output))
assert diffusion.calls == 1
def _full_mode() -> _Mode:
"""Return one full 4x2 canvas with left-half regional ownership."""
multiplier = torch.tensor([[[[1.5, 1.5, 1.0, 1.0]]]])
return _Mode((1, 2, 2, 4), None, multiplier)
def _tiled_mode() -> _Mode:
"""Return left and right view-major tiles over the same canvas."""
layout = SpatialBatchLayout(
4,
2,
(
SpatialView(SpatialViewKind.TILE, 0, 0, 2, 2, 2, 2),
SpatialView(SpatialViewKind.TILE, 2, 0, 2, 2, 2, 2),
),
1,
)
multiplier = torch.tensor([1.5, 1.0]).reshape(2, 1, 1, 1)
return _Mode((2, 2, 2, 2), layout, multiplier)
def _contextual_mode() -> _Mode:
"""Return one downscaled full-canvas Contextual global view."""
layout = SpatialBatchLayout(
4,
2,
(
SpatialView(
SpatialViewKind.CONTEXTUAL_GLOBAL,
0,
0,
4,
2,
2,
2,
),
),
1,
)
multiplier = torch.tensor([[[[1.5, 1.0]]]])
return _Mode((1, 2, 2, 2), layout, multiplier)
def _patcher(
*,
base_scale: float = 1.0,
) -> tuple[comfy.model_patcher.ModelPatcher, _OperationDiffusion]:
"""Return one real CPU patcher around the test diffusion target."""
diffusion = _OperationDiffusion()
with torch.no_grad():
diffusion.linear.weight.mul_(base_scale)
root = nn.Module()
root.diffusion_model = diffusion
device = torch.device("cpu")
return comfy.model_patcher.ModelPatcher(root, device, device), diffusion
def _lora_plan() -> RegionalLoraPlan:
"""Return one positive adapter at half strength on the first region."""
return RegionalLoraPlan(
(
RegionalLoraAdapterPlan(
RegionalLoraAdapterIdentity("regional.safetensors"),
0,
0,
RegionalLoraBranch.POSITIVE,
0.5,
(RegionalLoraScheduleBoundary(0.0, 1.0, 1.0, 0),),
),
)
)
def _admission(
source: comfy.model_patcher.ModelPatcher,
plan: RegionalLoraPlan,
) -> StandardUnetOperationAdmission:
"""Build complete binder evidence for the declared regional adapter."""
adapter = plan.adapters[0]
payload = RegionalLoraHostPayload(False, None, {}, None)
resolution = ComfyRegionalLoraResolution(
(
ComfyRegionalAdapterResolution(
adapter,
payload,
(
ComfyNormalizedAdapterTarget(
ComfyAdapterTargetPath(
"diffusion_model.linear.weight",
None,
),
LoRAAdapter(
{"up", "down"},
(torch.eye(2), torch.eye(2), None, None, None, None),
),
"LoRAAdapter",
("down", "up"),
True,
),
),
(),
(),
),
),
(),
)
binding = REGIONAL_LORA_TARGET_BINDER.bind(
source=source,
candidate=source,
resolution=resolution,
operations=COMFY_RESOLVED_OPERATION_TRANSLATOR.translate(resolution),
linear_spatial_capabilities={
"diffusion_model.linear.weight": (
BoundRegionalLoraSpatialCapability.SPATIAL_TOKENS
)
},
)
return StandardUnetOperationAdmission(
RegionalLoraPlanAdaptation(plan, (payload,)),
binding,
{"diffusion_model.linear": BoundRegionalLoraSpatialCapability.SPATIAL_TOKENS},
RegionalLoraExecutionCache(),
)
def _state(plan: RegionalLoraPlan) -> StandardUnetAttentionState:
"""Return a left-half regional mask and matching processed plan."""
masks = torch.zeros((1, 2, 4))
masks[:, :, :2] = 1.0
bank = RegionalMaskBank(masks, masks.clone(), 4, 2)
regional = (_context(1, 0),)
processed = ProcessedRegionalAttentionPlan(
ProcessedRegionalAttentionBranch(_context(0, None), regional),
ProcessedRegionalAttentionBranch(_context(0, None), regional),
bank,
plan,
)
return StandardUnetAttentionState(
processed,
(1.0,),
RegionalAttentionDiagnosticsBuilder(bank, backend="test.shared-trajectory"),
)
def _context(
conditioning_index: int,
region_index: int | None,
) -> ProcessedRegionalAttentionContext:
"""Return one always-active compact conditioning context."""
return ProcessedRegionalAttentionContext(
conditioning_index,
region_index,
(
ProcessedRegionalAttentionEntry(
0,
UUID(int=conditioning_index + 1),
ConditioningScheduleRange(None, None, None, None),
torch.zeros((1, 2, 2)),
1.0,
),
),
)
-293
View File
@@ -1,293 +0,0 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify complete persistent-variant installation and detach cleanup."""
from __future__ import annotations
from typing import cast
from uuid import uuid4
import torch
from comfy.ldm.modules.attention import SpatialTransformer
from comfy.model_patcher import ModelPatcher
from comfy.patcher_extension import CallbacksMP, WrappersMP
from torch import nn
from simple_syrup.domain.conditioning_schedule import ConditioningScheduleRange
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionBranch,
ProcessedRegionalAttentionContext,
ProcessedRegionalAttentionEntry,
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.regional_lora_plan import (
RegionalLoraAdapterIdentity,
RegionalLoraAdapterPlan,
RegionalLoraBranch,
RegionalLoraPlan,
RegionalLoraScheduleBoundary,
)
from simple_syrup.domain.regional_mask_bank import RegionalMaskBank
from simple_syrup.runtime.attention_coupling.unet import StandardUnetAttentionBackend
from simple_syrup.runtime.attention_coupling.unet_attention_state import (
StandardUnetAttentionState,
)
from simple_syrup.runtime.regional_attention_diagnostics import (
RegionalAttentionDiagnosticsBuilder,
)
from simple_syrup.runtime.regional_lora.comfy_adapter_resolution import (
ComfyAdapterTargetPath,
ComfyNormalizedAdapterTarget,
ComfyRegionalAdapterResolution,
ComfyRegionalLoraResolution,
)
from simple_syrup.runtime.regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
)
from simple_syrup.runtime.regional_lora.standard_unet_variant_template import (
STANDARD_UNET_VARIANT_TEMPLATE_CACHE,
)
from simple_syrup.runtime.regional_lora_host_payload import RegionalLoraHostPayload
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
class _Diffusion(nn.Module):
"""Expose one conventional target and a standard-style private forward."""
def __init__(self) -> None:
"""Create one deterministic base projection."""
super().__init__()
self.layer = nn.Linear(1, 1, bias=False)
self.layer.weight.data.fill_(1.0)
self.input_blocks = nn.ModuleList(
[
nn.Sequential(
SpatialTransformer(
in_channels=32,
n_heads=1,
d_head=32,
depth=1,
context_dim=32,
use_checkpoint=False,
)
)
]
)
def _forward(
self, inputs: torch.Tensor, *args: object, **kwargs: object
) -> torch.Tensor:
"""Run the generic target projection."""
del args, kwargs
return cast(torch.Tensor, self.layer(inputs))
class _Root(nn.Module):
"""Expose the diffusion module through a Comfy patcher root."""
def __init__(self) -> None:
"""Create the source diffusion graph."""
super().__init__()
self.diffusion_model = _Diffusion()
def test_runtime_reuses_two_exact_template_variants_across_request_detach() -> None:
"""Retain exact shared variants while clearing only request-local state."""
source = _Root()
device = torch.device("cpu")
patcher = ModelPatcher(source, device, device)
admission, state = _admission_and_state()
def first_wrapper(function: object, arguments: object) -> object:
"""Represent one request-local upstream model wrapper."""
del function, arguments
return None
patcher.set_model_unet_function_wrapper(first_wrapper)
built = StandardUnetAttentionBackend().derive(
model=patcher,
state=state,
admission=admission,
)
derived = built.model
assert isinstance(derived, ModelPatcher)
assert "patches" not in derived.model_options["transformer_options"]
replacements = derived.model_options["transformer_options"].get(
"patches_replace",
{},
)
assert isinstance(replacements, dict)
assert "attn1" not in replacements
variant_root = derived.object_patches["diffusion_model"]
assert isinstance(variant_root, nn.Module)
container = variant_root.simple_syrup_regional_variants
assert isinstance(container, nn.ModuleDict)
assert len(container) == 2
assert (
len(
derived.get_wrappers(
WrappersMP.PREPARE_SAMPLING,
"simple_syrup.standard_unet_static_variant_residency",
)
)
== 1
)
weights = tuple(
float(cast(_Diffusion, shell).layer.weight.item())
for shell in container.values()
)
assert weights == (2.0, 3.0)
assert float(source.diffusion_model.layer.weight.item()) == 1.0
assert derived.model_options["model_function_wrapper"] is first_wrapper
def second_wrapper(function: object, arguments: object) -> object:
"""Represent a changed request-local upstream model wrapper."""
del function, arguments
return None
patcher.set_model_unet_function_wrapper(second_wrapper)
static_source = derived.model.diffusion_model
derived.object_patches_backup["diffusion_model"] = static_source
derived.model.diffusion_model = variant_root
second = (
StandardUnetAttentionBackend()
.derive(
model=patcher,
state=state,
admission=admission,
)
.model
)
assert isinstance(second, ModelPatcher)
assert second.model is derived.model
assert second.object_patches["diffusion_model"] is variant_root
assert second.model_options["model_function_wrapper"] is second_wrapper
callback = derived.get_callbacks(
CallbacksMP.ON_DETACH,
"simple_syrup.standard_unet_regional_variants",
)[0]
callback(derived, False)
assert len(container) == 2
callback(derived, True)
assert len(container) == 2
STANDARD_UNET_VARIANT_TEMPLATE_CACHE.clear()
def _admission_and_state() -> tuple[
StandardUnetNativeLoraAdmission,
StandardUnetAttentionState,
]:
"""Return paired generic adapter resolution and matching attention state."""
plans = tuple(
_adapter(index, region, branch)
for index, (region, branch) in enumerate(
(
(0, RegionalLoraBranch.POSITIVE),
(1, RegionalLoraBranch.POSITIVE),
(0, RegionalLoraBranch.NEGATIVE),
(1, RegionalLoraBranch.NEGATIVE),
)
)
)
payload_values = (object(), object())
payloads = (
RegionalLoraHostPayload.unresolved(payload_values[0]),
RegionalLoraHostPayload.unresolved(payload_values[1]),
RegionalLoraHostPayload.unresolved(payload_values[0]),
RegionalLoraHostPayload.unresolved(payload_values[1]),
)
adaptation = RegionalLoraPlanAdaptation(RegionalLoraPlan(plans), payloads)
results = tuple(
ComfyRegionalAdapterResolution(
plan,
payload,
(_target(1.0 if plan.region_index == 0 else 2.0),),
(),
(),
)
for plan, payload in zip(plans, payloads, strict=True)
)
admission = StandardUnetNativeLoraAdmission(
adaptation,
ComfyRegionalLoraResolution(results, ()),
)
masks = torch.tensor([[[1.0, 0.0]], [[0.0, 1.0]]])
processed = ProcessedRegionalAttentionPlan(
ProcessedRegionalAttentionBranch(
_context(0, None),
(_context(1, 0), _context(2, 1)),
),
ProcessedRegionalAttentionBranch(
_context(0, None),
(_context(1, 0), _context(2, 1)),
),
RegionalMaskBank(masks.clone(), masks.clone(), 2, 1),
adaptation.plan,
)
state = StandardUnetAttentionState(
processed,
(1.0, 1.0),
RegionalAttentionDiagnosticsBuilder(processed.mask_bank, backend="generic"),
)
return admission, state
def _adapter(
index: int,
region: int,
branch: RegionalLoraBranch,
) -> RegionalLoraAdapterPlan:
"""Return one fixed-strength generic adapter use."""
return RegionalLoraAdapterPlan(
RegionalLoraAdapterIdentity(f"adapter-{region}"),
index,
region,
branch,
1.0,
(RegionalLoraScheduleBoundary(0.0, 1.0, 1.0, 0),),
)
def _target(delta: float) -> ComfyNormalizedAdapterTarget:
"""Return one conventional additive target operation."""
return ComfyNormalizedAdapterTarget(
ComfyAdapterTargetPath("diffusion_model.layer.weight", None),
("diff", (torch.tensor([[delta]]),)),
"LegacyDiff",
("source",),
False,
)
def _context(index: int, region: int | None) -> ProcessedRegionalAttentionContext:
"""Return one always-active generic processed context."""
return ProcessedRegionalAttentionContext(
index,
region,
(
ProcessedRegionalAttentionEntry(
0,
uuid4(),
ConditioningScheduleRange(None, None, None, None),
torch.zeros((1, 1, 1)),
1.0,
),
),
)
+4 -4
View File
@@ -30,8 +30,8 @@ from simple_syrup.runtime.attention_coupling.unet_attention_state import (
from simple_syrup.runtime.regional_attention_diagnostics import (
RegionalAttentionDiagnosticsBuilder,
)
from simple_syrup.runtime.regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
from simple_syrup.runtime.regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
@@ -145,11 +145,11 @@ def test_standard_unet_backend_preserves_existing_attn1_patch() -> None:
assert source.object_patches == {}
def _empty_admission() -> StandardUnetNativeLoraAdmission:
def _empty_admission() -> StandardUnetOperationAdmission:
"""Return one prompt-only standard-family admission."""
adaptation = RegionalLoraPlanAdaptation(EMPTY_REGIONAL_LORA_PLAN, ())
return StandardUnetNativeLoraAdmission(adaptation, None)
return StandardUnetOperationAdmission(adaptation, None, {}, None)
def _state() -> StandardUnetAttentionState:
@@ -31,8 +31,8 @@ from simple_syrup.runtime.attention_coupling.unet_attention_state import (
from simple_syrup.runtime.regional_attention_diagnostics import (
RegionalAttentionDiagnosticsBuilder,
)
from simple_syrup.runtime.regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
from simple_syrup.runtime.regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
@@ -181,9 +181,11 @@ def test_backend_projects_unique_resolutions_once_in_one_native_trajectory(
diffusion_model = _ResolutionDiffusionModel(context_dimension)
source = _patcher(diffusion_model)
state = _state(context_dimension)
admission = StandardUnetNativeLoraAdmission(
admission = StandardUnetOperationAdmission(
RegionalLoraPlanAdaptation(EMPTY_REGIONAL_LORA_PLAN, ()),
None,
{},
None,
)
built = StandardUnetAttentionBackend().derive(
model=source,
@@ -41,9 +41,9 @@ from simple_syrup.runtime.attention_coupling.unet_attention_state import (
from simple_syrup.runtime.attention_coupling.unet_context import (
STANDARD_UNET_REGIONAL_CONTEXT_VALIDATOR,
)
from simple_syrup.runtime.regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
StandardUnetNativeLoraAdmissionService,
from simple_syrup.runtime.regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
StandardUnetOperationPreparation,
)
from simple_syrup.runtime.regional_lora_host_payload import RegionalLoraHostPayload
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
@@ -68,8 +68,8 @@ class _Backend:
return SimpleNamespace(model="derived-unet")
class _NativeAdmission:
"""Capture native adapter admission exactly once."""
class _OperationPreparation:
"""Capture operation admission exactly once."""
calls: ClassVar[list[tuple[object, RegionalLoraPlanAdaptation]]] = []
@@ -77,34 +77,34 @@ class _NativeAdmission:
self,
model: object,
adaptation: RegionalLoraPlanAdaptation,
) -> StandardUnetNativeLoraAdmission:
"""Return a recognizable typed empty native admission."""
) -> StandardUnetOperationAdmission:
"""Return a recognizable typed empty operation admission."""
type(self).calls.append((model, adaptation))
return StandardUnetNativeLoraAdmission(adaptation, None)
return StandardUnetOperationAdmission(adaptation, None, {}, None)
def test_unet_family_selects_native_variant_backend_boundaries() -> None:
"""Make the standard family own native admission and variant execution."""
def test_unet_family_selects_spatial_operation_backend_boundaries() -> None:
"""Make the standard family own operation admission and shared execution."""
family = StandardUnetAttentionCouplingModelFamily()
assert family.backend_class is StandardUnetAttentionBackend
assert family.native_admission_class is StandardUnetNativeLoraAdmissionService
assert family.operation_preparation_class is StandardUnetOperationPreparation
def test_unet_family_builds_shared_state_and_derives_paired_backend() -> None:
"""Bind native admission, state, diagnostics, and backend exactly once."""
"""Bind operation admission, state, diagnostics, and backend exactly once."""
family = StandardUnetAttentionCouplingModelFamily()
plan = _plan()
adaptation = RegionalLoraPlanAdaptation(EMPTY_REGIONAL_LORA_PLAN, ())
original_backend = family.backend_class
original_admission = family.native_admission_class
original_preparation = family.operation_preparation_class
type(family).backend_class = _Backend # type: ignore[assignment]
type(family).native_admission_class = _NativeAdmission # type: ignore[assignment]
type(family).operation_preparation_class = _OperationPreparation # type: ignore[assignment]
_Backend.calls = []
_NativeAdmission.calls = []
_OperationPreparation.calls = []
try:
family.validate_latent(torch.zeros(2, 4, 8, 8))
admission = family.admit_adaptation("model", adaptation)
@@ -118,12 +118,12 @@ def test_unet_family_builds_shared_state_and_derives_paired_backend() -> None:
)
finally:
type(family).backend_class = original_backend
type(family).native_admission_class = original_admission
type(family).operation_preparation_class = original_preparation
assert family.context_validator is STANDARD_UNET_REGIONAL_CONTEXT_VALIDATOR
assert derived == "derived-unet"
assert admission.adaptation is adaptation
assert _NativeAdmission.calls == [("model", adaptation)]
assert _OperationPreparation.calls == [("model", adaptation)]
assert len(_Backend.calls) == 1
assert _Backend.calls[0]["admission"] is admission
state = _Backend.calls[0]["state"]
@@ -155,12 +155,12 @@ def test_unet_family_preserves_prompt_only_base_sampler_conditioning() -> None:
assert built.negative is negative_base
def test_unet_family_derives_variant_capable_backend_for_prompt_coupling() -> None:
"""Keep standard prompt coupling inside the variant-capable backend."""
def test_unet_family_derives_shared_trajectory_backend_for_prompt_coupling() -> None:
"""Keep standard prompt coupling inside the shared-trajectory backend."""
family = StandardUnetAttentionCouplingModelFamily()
adaptation = RegionalLoraPlanAdaptation(EMPTY_REGIONAL_LORA_PLAN, ())
admission = StandardUnetNativeLoraAdmission(adaptation, None)
admission = StandardUnetOperationAdmission(adaptation, None, {}, None)
original_backend = family.backend_class
type(family).backend_class = _Backend # type: ignore[assignment]
_Backend.calls = []
@@ -187,7 +187,7 @@ def test_unet_family_requires_matching_family_admission() -> None:
adaptation = _regional_lora_adaptation()
_Backend.calls = []
with pytest.raises(TypeError, match="native admission"):
with pytest.raises(TypeError, match="operation admission"):
family.derive(
model="model",
processed_plan=_plan(),
@@ -200,9 +200,11 @@ def test_unet_family_requires_matching_family_admission() -> None:
family.derive(
model="model",
processed_plan=_plan(adaptation.plan),
admission=StandardUnetNativeLoraAdmission(
admission=StandardUnetOperationAdmission(
RegionalLoraPlanAdaptation(EMPTY_REGIONAL_LORA_PLAN, ()),
None,
{},
None,
),
interop_report=_interop_report(),
region_strengths=(1.0,),
+5 -3
View File
@@ -32,8 +32,8 @@ from simple_syrup.runtime.attention_coupling.unet_attention_state import (
from simple_syrup.runtime.regional_attention_diagnostics import (
RegionalAttentionDiagnosticsBuilder,
)
from simple_syrup.runtime.regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
from simple_syrup.runtime.regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
@@ -44,9 +44,11 @@ def test_prompt_only_backend_preserves_reference_self_attention() -> None:
built = StandardUnetAttentionBackend().derive(
model=_patcher(),
state=_state(),
admission=StandardUnetNativeLoraAdmission(
admission=StandardUnetOperationAdmission(
RegionalLoraPlanAdaptation(EMPTY_REGIONAL_LORA_PLAN, ()),
None,
{},
None,
),
)
derived: Any = built.model
+7 -27
View File
@@ -29,13 +29,6 @@ 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):
@@ -198,11 +191,8 @@ 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."""
def test_negpip_splits_exact_packed_regional_key_and_value_views() -> None:
"""Select even K and odd V tokens after regional branch packing."""
contexts = _negpip_contexts()
execution = UnetAttn2Execution(
@@ -224,22 +214,12 @@ def test_negpip_splits_exact_packed_regional_key_and_value_views(
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],
}
options = {
"patches": {
"attn2_patch": [pair.input_patch, split_negpip],
"attn2_output_patch": [pair.output_patch],
}
}
recording = _RecordingKeyValueAttention()
block = _block(recording)
@@ -26,8 +26,8 @@ from simple_syrup.runtime.attention_coupling.unet_attention_state import (
from simple_syrup.runtime.regional_attention_diagnostics import (
RegionalAttentionDiagnosticsBuilder,
)
from simple_syrup.runtime.regional_lora.standard_unet_native_admission import (
StandardUnetNativeLoraAdmission,
from simple_syrup.runtime.regional_lora.standard_unet_operation_preparation import (
StandardUnetOperationAdmission,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
@@ -71,9 +71,11 @@ class UnetAttentionCouplingLifecycleHarness(AttentionCouplingLifecycleHarness):
built = StandardUnetAttentionBackend().derive(
model=source,
state=state,
admission=StandardUnetNativeLoraAdmission(
admission=StandardUnetOperationAdmission(
RegionalLoraPlanAdaptation(plan.lora_plan, ()),
None,
{},
None,
),
)
derived: Any = built.model
+1 -1
View File
@@ -101,7 +101,7 @@ def default_manifest(repo_root: Path) -> AnimaRegressionManifest:
"tests/test_anima_multi_lora_fidelity.py",
"tests/test_anima_multi_lora_composition.py",
"tests/test_anima_full_tile_lora_equivalence.py",
"tests/test_anima_global_lora_overlap.py",
"tests/test_regional_lora_conditioning_adapter.py",
"tests/test_anima_lora_block.py",
"tests/test_anima_lora_block_inactive_schedule.py",
"tests/test_anima_projection_batch.py",
-2
View File
@@ -21,7 +21,6 @@ class GlobalLoraVisualCase:
label: str
global_strength: float | None = None
regional_strength: float | None = None
expect_overlap_rejection: bool = False
def cases() -> tuple[GlobalLoraVisualCase, ...]:
@@ -49,7 +48,6 @@ def cases() -> tuple[GlobalLoraVisualCase, ...]:
"GLOBAL + LEFT PRIMARY_ADAPTER 1.0 — duplicate",
global_strength=1.0,
regional_strength=1.0,
expect_overlap_rejection=True,
),
)
+7 -61
View File
@@ -17,8 +17,6 @@ from tools.comfy_integration.portable_font import load_label_font
from .matrix import GlobalLoraVisualCase
_OVERLAP_MESSAGE = "Regional Anima LoRA content is already applied globally"
class GlobalLoraVisualProofRecorder:
"""Own ordered proof observations, full images, and the comparison sheet."""
@@ -44,8 +42,6 @@ class GlobalLoraVisualProofRecorder:
) -> Path:
"""Persist one successful full-resolution image and exact sidecars."""
if case.expect_overlap_rejection:
raise ValueError("Duplicate LoRA case cannot be recorded as success.")
_require_status(history, "success")
path = self._root / f"{case.case_id}.png"
path.write_bytes(image_bytes)
@@ -68,36 +64,6 @@ class GlobalLoraVisualProofRecorder:
)
return path
def record_overlap_rejection(
self,
case: GlobalLoraVisualCase,
*,
workflow: dict[str, JsonObject],
history: JsonObject,
prompt_id: str,
) -> None:
"""Persist the intentional pre-sampling duplicate rejection."""
if not case.expect_overlap_rejection:
raise ValueError("Success case cannot be recorded as a rejection.")
_require_status(history, "error")
serialized = json.dumps(history, sort_keys=True)
if (
_OVERLAP_MESSAGE not in serialized
or "adapter-a.safetensors" not in serialized
):
raise ValueError("Duplicate case lacks the exact overlap diagnostic.")
self._write_sidecars(case, workflow, history)
self._observations.append(
{
"case_id": case.case_id,
"label": case.label,
"status": "rejected_before_sampling",
"prompt_id": prompt_id,
"diagnostic": _OVERLAP_MESSAGE,
}
)
def finalize(
self,
definitions: tuple[GlobalLoraVisualCase, ...],
@@ -137,14 +103,13 @@ class GlobalLoraVisualProofRecorder:
self,
definitions: tuple[GlobalLoraVisualCase, ...],
) -> Path:
"""Render four outputs and the duplicate diagnostic in a labeled grid."""
"""Render all five placement outputs in a labeled grid."""
panel = 768
header = 80
canvas = Image.new("RGB", (panel * 3, (panel + header) * 2), (20, 20, 20))
draw = ImageDraw.Draw(canvas)
title_font = load_label_font(24)
body_font = load_label_font(20)
for index, case in enumerate(definitions):
column = index % 3
row = index // 3
@@ -153,32 +118,13 @@ class GlobalLoraVisualProofRecorder:
draw.text((left + 18, top + 25), case.label, fill="white", font=title_font)
body_top = top + header
path = self._images.get(case.case_id)
if path is not None:
with Image.open(path) as source:
image = source.convert("RGB").resize(
(panel, panel), Image.Resampling.LANCZOS
)
canvas.paste(image, (left, body_top))
continue
draw.rectangle(
(left, body_top, left + panel - 1, body_top + panel - 1),
fill=(45, 28, 28),
outline=(220, 90, 90),
width=3,
)
lines = (
"REJECTED BEFORE SAMPLING",
"The same PRIMARY_ADAPTER weights were already",
"installed globally on the input model.",
"Current policy prevents double application.",
)
for line_index, line in enumerate(lines):
draw.text(
(left + 48, body_top + 230 + line_index * 42),
line,
fill=(255, 220, 220),
font=body_font,
if path is None:
raise ValueError(f"Global LoRA proof lacks image for {case.case_id!r}.")
with Image.open(path) as source:
image = source.convert("RGB").resize(
(panel, panel), Image.Resampling.LANCZOS
)
canvas.paste(image, (left, body_top))
path = self._root / "global-lora-placement__labeled-comparison.png"
canvas.save(path, format="PNG")
return path
@@ -24,7 +24,6 @@ class GlobalRegionalLoraCase:
"""Bind one public workflow case to its expected terminal status."""
integration: IntegrationCase
expect_overlap_error: bool = False
@property
def case_id(self) -> str:
@@ -84,13 +83,12 @@ def cases() -> tuple[GlobalRegionalLoraCase, ...]:
GlobalRegionalLoraCase(
IntegrationCase(
"duplicate-global-regional-primary_adapter",
"Exact global and regional PRIMARY_ADAPTER duplicate rejection",
"Same PRIMARY_ADAPTER globally and regionally with additive use",
1.0,
MASK_CASE_ID,
0,
global_loras=(global_primary_adapter,),
regional_loras=(regional_primary_adapter,),
),
expect_overlap_error=True,
)
),
)
@@ -19,8 +19,6 @@ from tools.comfy_api import ImageReference, JsonObject
from .image_validation import GLOBAL_REGIONAL_LORA_IMAGE_VALIDATOR
from .matrix import GlobalRegionalLoraCase
_OVERLAP_MESSAGE = "Regional Anima LoRA content is already applied globally"
class GlobalRegionalLoraResultRecorder:
"""Own terminal status, ordering, and durable P9.3 evidence persistence."""
@@ -46,8 +44,6 @@ class GlobalRegionalLoraResultRecorder:
) -> Path:
"""Require success and preserve its labeled image and sidecars."""
if case.expect_overlap_error:
raise ValueError("P9.3 overlap case cannot be recorded as success.")
_require_status(history, "success")
metrics = _metrics(history, workflow)
image_path = self._root / f"{case.case_id}.png"
@@ -72,38 +68,6 @@ class GlobalRegionalLoraResultRecorder:
)
return image_path
def record_overlap_rejection(
self,
case: GlobalRegionalLoraCase,
workflow: BuiltAnimaAttentionCouplingWorkflow,
*,
prompt_id: str,
history: JsonObject,
) -> None:
"""Require the exact pre-sampling duplicate diagnostic and no image."""
if not case.expect_overlap_error:
raise ValueError("P9.3 success case cannot be recorded as rejection.")
_require_status(history, "error")
serialized = json.dumps(history, sort_keys=True)
if (
_OVERLAP_MESSAGE not in serialized
or "adapter-a.safetensors" not in serialized
):
raise ValueError("P9.3 history lacks the exact overlap diagnostic.")
if workflow.save_node_id in _outputs(history):
raise ValueError("P9.3 overlap rejection unexpectedly saved an image.")
self._sidecars(case, workflow, history)
self._observations.append(
{
"case_id": case.case_id,
"label": case.integration.label,
"status": "rejected_before_sampling",
"prompt_id": prompt_id,
"diagnostic": _OVERLAP_MESSAGE,
}
)
def finalize(
self,
definitions: tuple[GlobalRegionalLoraCase, ...],
-8
View File
@@ -81,14 +81,6 @@ def main() -> int:
LOGGER.info("Starting %s", case.label)
prompt_id = running.client.submit(workflow)
history = running.client.wait_for_history(prompt_id, timeout=1200.0)
if case.expect_overlap_rejection:
recorder.record_overlap_rejection(
case,
workflow=workflow,
history=history,
prompt_id=prompt_id,
)
continue
reference = extract_saved_image(history, "11")
image_bytes = running.client.download_image(reference)
recorder.record_success(
+398
View File
@@ -0,0 +1,398 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Run exact user-prompt global and regional LoRA proof workflows."""
from __future__ import annotations
import argparse
import copy
import hashlib
import json
import logging
import re
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import cast
from PIL import Image
from tools.comfy_api import JsonObject
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.comfy_integration.default_paths import (
default_benchmark_artifact_root,
default_comfy_root,
)
from tools.comfy_integration.history_output import extract_saved_image
from tools.comfy_integration.loopback_port import is_loopback_port_available
from tools.comfy_integration.managed_server import ManagedComfyServer
DEFAULT_OUTPUT_ROOT = default_benchmark_artifact_root("global-prompt-lora-proof")
_SEPARATOR = re.compile(r"(\[SEP(?:\|[^\]]+)?\])")
_SDXL_GLOBAL = r"<lora:Illustrious\Style\IriaStyleIllustriousV1-000005.safetensors:0.5>"
_ANIMA_ARCANE_GLOBAL = r"<lora:Anima\style\ArcaneViolet_mpt_64.safetensors:0.6>"
_ANIMA_ARCANE_REGIONAL = r"<lora:Anima\style\ArcaneViolet_mpt_64.safetensors:0.8>"
_ANIMA_TURBO_GLOBAL = r"<lora:Anima\anima-turbo-lora-v0.2.safetensors:0.7>"
LOGGER = logging.getLogger(__name__)
@dataclass(frozen=True, slots=True)
class ProofCase:
"""Describe one exact template mutation and expected output pair."""
case_id: str
template: str
global_tag: str
regional_tag: str
region_mask_feather: int = 64
turbo: bool = False
contextual: bool = False
def cases() -> tuple[ProofCase, ...]:
"""Return SDXL, Anima duplicate, Turbo, and contextual proof coverage."""
return (
ProofCase(
"sdxl-global-and-regional-same-lora",
"sdxl",
_SDXL_GLOBAL,
_SDXL_GLOBAL,
region_mask_feather=10,
),
ProofCase(
"anima-global-and-regional-same-lora",
"anima",
_ANIMA_ARCANE_GLOBAL,
_ANIMA_ARCANE_REGIONAL,
),
ProofCase(
"anima-turbo-global-arcane-regional-tiled",
"anima",
_ANIMA_TURBO_GLOBAL,
_ANIMA_ARCANE_REGIONAL,
turbo=True,
),
ProofCase(
"anima-turbo-global-arcane-regional-contextual",
"anima",
_ANIMA_TURBO_GLOBAL,
_ANIMA_ARCANE_REGIONAL,
turbo=True,
contextual=True,
),
)
def render_global_first_prompt(
prompt: str,
*,
global_tag: str,
regional_tag: str,
) -> str:
"""Add one global tag and the same or distinct tag to the first region."""
pieces = _SEPARATOR.split(prompt)
segment_indices = tuple(range(0, len(pieces), 2))
if len(segment_indices) != 3:
raise ValueError("Global LoRA proof requires exactly three prompt segments.")
pieces[segment_indices[0]] = f"{global_tag}\n{pieces[segment_indices[0]].strip()}"
pieces[segment_indices[1]] = f"{regional_tag}\n{pieces[segment_indices[1]].strip()}"
return "".join(pieces)
def build_case_graph(
template: dict[str, JsonObject],
case: ProofCase,
*,
run_id: str,
) -> tuple[dict[str, JsonObject], tuple[str, str]]:
"""Return one isolated API graph plus its source and refinement save IDs."""
graph = copy.deepcopy(template)
_wrap_replayed_list_widgets(graph)
for node_id in tuple(graph):
if node_id.startswith("__sugarcubes_cube_output__"):
del graph[node_id]
prefix = "prompt-region" if case.template == "sdxl" else "anima-prompt-region"
refine_prefix = (
"tiled-upscale" if case.template == "sdxl" else "anima-diffusion-upscale"
)
for node_id in (
f"{prefix}:positive_prompt",
f"{refine_prefix}:positive_prompt",
):
inputs = _inputs(graph, node_id)
prompt = inputs.get("value")
if not isinstance(prompt, str):
raise TypeError(f"Proof prompt node {node_id!r} lacks string value.")
inputs["value"] = render_global_first_prompt(
prompt,
global_tag=case.global_tag,
regional_tag=case.regional_tag,
)
for sampler_id in (f"{prefix}:ksampler", f"{refine_prefix}:ksampler"):
_inputs(graph, sampler_id)["region_mask_feather"] = case.region_mask_feather
if case.turbo:
for sampler_id in (f"{prefix}:ksampler", f"{refine_prefix}:ksampler"):
sampler = _inputs(graph, sampler_id)
sampler.update(
{
"steps": 10,
"cfg": 1.0,
"sampler_name": "euler",
"scheduler": "simple",
}
)
if case.contextual:
sampler = graph[f"{refine_prefix}:ksampler"]
sampler["class_type"] = "SimpleSyrup.KSamplerAttentionCouplingContextual"
inputs = _inputs(graph, f"{refine_prefix}:ksampler")
for key in (
"latent_tile_width",
"latent_tile_height",
"latent_tile_overlap",
"latent_tile_batch_size",
):
inputs.pop(key, None)
inputs.update(
{
"diffusion_mode": "mixture_of_diffusers",
"latent_context_size": 128,
"latent_context_overlap": 32,
"latent_context_batch_size": 4,
"global_weight": 1.0,
"global_steps": 1,
"global_decay": 0.5,
}
)
source_save = "proof:source"
refinement_save = "proof:refinement"
output_prefix = f"simple_syrup_global_prompt_lora/{run_id}/{case.case_id}"
graph[source_save] = {
"class_type": "SaveImage",
"inputs": {
"images": [f"{prefix}:vae_decode", 0],
"filename_prefix": f"{output_prefix}-source",
},
}
graph[refinement_save] = {
"class_type": "SaveImage",
"inputs": {
"images": [f"{refine_prefix}:vae_decode", 0],
"filename_prefix": f"{output_prefix}-refinement",
},
}
return graph, (source_save, refinement_save)
def execute(
*,
sdxl_template_path: Path,
anima_template_path: Path,
output_root: Path,
comfy_root: Path,
case_ids: frozenset[str] | None = None,
) -> Path:
"""Run all proof cases in one managed Comfy process and persist evidence."""
artifacts = IntegrationArtifacts(output_root)
templates = {
"sdxl": _load_execution_prompt(sdxl_template_path),
"anima": _load_execution_prompt(anima_template_path),
}
available = cases()
definitions = tuple(
case for case in available if case_ids is None or case.case_id in case_ids
)
if not definitions:
raise ValueError("Global LoRA proof selection contains no known cases.")
if case_ids is not None:
unknown = case_ids - frozenset(case.case_id for case in available)
if unknown:
raise ValueError(f"Unknown global LoRA proof cases: {sorted(unknown)!r}.")
built = tuple(
(
case,
*build_case_graph(
templates[case.template],
case,
run_id=artifacts.run_id,
),
)
for case in definitions
)
required = frozenset(
cast(str, node["class_type"])
for _case, graph, _save_ids in built
for node in graph.values()
)
observations: list[JsonObject] = []
port = 0
process = None
with ManagedComfyServer(
comfy_root=comfy_root,
artifacts=artifacts,
required_node_ids=required,
readiness_timeout=300.0,
) as running:
for case, graph, save_ids in built:
prompt_id = running.client.submit(graph)
history = running.client.wait_for_history(prompt_id, timeout=1800.0)
status = history.get("status")
if not isinstance(status, dict) or status.get("status_str") != "success":
raise ValueError(f"Proof case {case.case_id!r} did not succeed.")
files: list[JsonObject] = []
for label, save_id in zip(
("source", "refinement"),
save_ids,
strict=True,
):
reference = extract_saved_image(history, save_id)
image_bytes = running.client.download_image(reference)
path = artifacts.root / f"{case.case_id}--{label}.png"
path.write_bytes(image_bytes)
with Image.open(path) as image:
width, height = image.size
files.append(
{
"label": label,
"file": path.name,
"width": width,
"height": height,
"sha256": hashlib.sha256(image_bytes).hexdigest(),
}
)
workflow_path = artifacts.root / f"{case.case_id}.workflow.json"
history_path = artifacts.root / f"{case.case_id}.history.json"
_write_json(workflow_path, graph)
_write_json(history_path, history)
observations.append(
{
"case_id": case.case_id,
"execution_status": "success",
"visual_review": "pending",
"prompt_id": prompt_id,
"global_tag": case.global_tag,
"regional_tag": case.regional_tag,
"region_mask_feather": case.region_mask_feather,
"turbo_settings": (
{
"steps": 10,
"cfg": 1.0,
"sampler": "euler",
"scheduler": "simple",
}
if case.turbo
else None
),
"refinement_mode": "contextual" if case.contextual else "tiled",
"files": files,
"workflow_file": workflow_path.name,
"history_file": history_path.name,
}
)
system_stats = running.system_stats
port = running.port
process = running.process
if process is None:
raise RuntimeError("Managed proof process did not reach ready state.")
port_available = is_loopback_port_available(port)
artifacts.record_cleanup(
process_running=process.is_running,
port_available=port_available,
)
if process.is_running or not port_available:
raise RuntimeError("Managed proof process or port cleanup failed.")
result = artifacts.root / "global-prompt-lora-proof.json"
_write_json(
result,
{
"schema_version": 2,
"execution_status": "completed",
"visual_review": "pending",
"observations": observations,
"system_stats": system_stats,
"cleanup": {"process_stopped": True, "port_available": True},
},
)
return result
def _load_execution_prompt(path: Path) -> dict[str, JsonObject]:
"""Load one SugarCubes queue response's expanded execution prompt."""
decoded = json.loads(path.read_text(encoding="utf-8"))
if not isinstance(decoded, dict):
raise TypeError("Proof template response must be an object.")
prompt = decoded.get("execution_prompt")
if not isinstance(prompt, dict):
raise TypeError("Proof template response lacks execution_prompt.")
return cast(dict[str, JsonObject], prompt)
def _inputs(graph: dict[str, JsonObject], node_id: str) -> JsonObject:
"""Return one mutable node input mapping from a proof graph."""
node = graph.get(node_id)
if not isinstance(node, dict):
raise KeyError(f"Proof graph lacks node {node_id!r}.")
inputs = node.get("inputs")
if not isinstance(inputs, dict):
raise TypeError(f"Proof node {node_id!r} inputs must be an object.")
return inputs
def _wrap_replayed_list_widgets(graph: dict[str, JsonObject]) -> None:
"""Restore Comfy's literal wrapper around executed multiselect list values."""
for node in graph.values():
if node.get("class_type") != "SimpleSyrup.LoadMaskBatch":
continue
inputs = node.get("inputs")
if not isinstance(inputs, dict):
raise TypeError("Replayed mask loader inputs must be an object.")
image = inputs.get("image")
if isinstance(image, list):
inputs["image"] = {"__value__": image}
def _write_json(path: Path, value: object) -> None:
"""Write deterministic UTF-8 proof evidence."""
path.write_text(
json.dumps(value, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
def main(argv: Sequence[str] | None = None) -> int:
"""Parse exact templates, run the proof, and report its manifest path."""
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--sdxl-template", type=Path, required=True)
parser.add_argument("--anima-template", type=Path, required=True)
parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT)
parser.add_argument("--comfy-root", type=Path, default=default_comfy_root())
parser.add_argument("--case-id", action="append", default=None)
args = parser.parse_args(argv)
try:
result = execute(
sdxl_template_path=args.sdxl_template,
anima_template_path=args.anima_template,
output_root=args.output_root,
comfy_root=args.comfy_root,
case_ids=(frozenset(args.case_id) if args.case_id else None),
)
except BaseException:
LOGGER.exception("Global prompt LoRA proof failed.")
return 1
LOGGER.info("Global prompt LoRA proof completed: %s", result)
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -44,7 +44,7 @@ def execute_matrix(
readiness_timeout: float,
prompt_timeout: float,
) -> Path:
"""Execute all success and intentional-rejection cases."""
"""Execute all global, regional, and additive-placement cases."""
definitions = cases()
manifest_case = next(
@@ -84,14 +84,6 @@ def execute_matrix(
prompt_id,
timeout=prompt_timeout,
)
if case.expect_overlap_error:
recorder.record_overlap_rejection(
case,
workflow,
prompt_id=prompt_id,
history=history,
)
continue
reference = extract_saved_image(history, workflow.save_node_id)
image = running.client.download_image(reference)
path = recorder.record_success(