Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9948cb3433 | ||
|
|
6cb9bbe868 |
@@ -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)
|
||||
|
||||
|
||||
|
||||
Generated
+2
-2
@@ -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
@@ -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
@@ -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"
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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]]
|
||||
@@ -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:
|
||||
|
||||
@@ -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": {},
|
||||
},
|
||||
}
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
),
|
||||
)
|
||||
@@ -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,
|
||||
),
|
||||
),
|
||||
)
|
||||
@@ -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,),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -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, ...],
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user