refactor(regional): reduce standard UNet cold startup

This commit is contained in:
Artificial Sweetener
2026-08-16 20:53:50 -04:00
parent 9de6033505
commit 24878eef89
83 changed files with 6422 additions and 210 deletions
@@ -12,6 +12,7 @@ from comfy.weight_adapter.base import WeightAdapterBase
from comfy.weight_adapter.lora import LoRAAdapter
from ...domain.regional_lora_plan import RegionalLoraAdapterPlan
from .comfy_adapter_identity_index import ComfyAdapterIdentityIndex
from .comfy_adapter_resolution import (
ComfyAdapterResolutionIssue,
ComfyAdapterResolutionIssueCode,
@@ -154,11 +155,11 @@ def _normalized_source_keys(
"""Recover consumed keys from adapter metadata or retained value identity."""
exposed = _operation_mapping_source_keys(normalized)
normalized_values = tuple(normalized.values())
identities = ComfyAdapterIdentityIndex.build(normalized.values())
return exposed | {
source_key
for source_key, source_value in raw_weights.items()
if any(_contains_identity(value, source_value) for value in normalized_values)
if identities.contains(source_value)
}
@@ -180,18 +181,6 @@ def _operation_source_keys(operation: object) -> tuple[str, ...]:
return tuple(sorted(operation.loaded_keys))
def _contains_identity(container: object, sought: object) -> bool:
"""Find a retained source value without tensor equality or device movement."""
if container is sought:
return True
if isinstance(container, (tuple, list)):
return any(_contains_identity(item, sought) for item in container)
if isinstance(container, WeightAdapterBase):
return _contains_identity(container.weights, sought)
return False
def _target_path(value: object) -> ComfyAdapterTargetPath | None:
"""Retain string and sliced target paths exactly as Comfy emits them."""
@@ -0,0 +1,55 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Index exact source-object identities retained by normalized Comfy adapters."""
from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass
from comfy.weight_adapter.base import WeightAdapterBase
@dataclass(frozen=True, slots=True)
class ComfyAdapterIdentityIndex:
"""Retain reachable object identities without equality or strong references."""
_identities: frozenset[int]
@classmethod
def build(cls, operations: Iterable[object]) -> ComfyAdapterIdentityIndex:
"""Walk each admitted host container edge once and tolerate cycles."""
identities: set[int] = set()
expanded: set[int] = set()
for operation in operations:
_collect_identities(operation, identities, expanded)
return cls(frozenset(identities))
def contains(self, value: object) -> bool:
"""Report whether the exact live object occurs in normalized evidence."""
return id(value) in self._identities
def _collect_identities(
value: object,
identities: set[int],
expanded: set[int],
) -> None:
"""Traverse only containers admitted by the established evidence contract."""
identity = id(value)
identities.add(identity)
if identity in expanded:
return
if isinstance(value, WeightAdapterBase):
expanded.add(identity)
_collect_identities(value.weights, identities, expanded)
return
if isinstance(value, tuple | list):
expanded.add(identity)
for item in value:
_collect_identities(item, identities, expanded)
@@ -8,6 +8,7 @@ from __future__ import annotations
import logging
from collections.abc import Mapping
from dataclasses import dataclass
from typing import cast
import comfy.lora
@@ -34,6 +35,32 @@ from .comfy_adapter_resolution import (
LOGGER = logging.getLogger(__name__)
@dataclass(slots=True)
class _ResolutionKeyMaps:
"""Build each successful host key map once per complete resolution."""
model: object
clip_model: object | None
_model_map: Mapping[object, object] | None = None
_clip_map: Mapping[object, object] | None = None
def model_map(self) -> Mapping[object, object]:
"""Return the one cached UNet key map after successful construction."""
if self._model_map is None:
self._model_map = comfy.lora.model_lora_keys_unet(self.model, {})
return self._model_map
def clip_map(self) -> Mapping[object, object]:
"""Return the one cached CLIP key map after successful construction."""
if self.clip_model is None:
return {}
if self._clip_map is None:
self._clip_map = comfy.lora.model_lora_keys_clip(self.clip_model, {})
return self._clip_map
class ComfyRegionalAdapterResolver:
"""Resolve regional model LoRAs without model, hook, or tensor mutation."""
@@ -48,22 +75,37 @@ class ComfyRegionalAdapterResolver:
"""Resolve every ordered payload and aggregate adapter-scoped failures."""
base_model = _base_model(model)
results = tuple(
self._resolve_adapter(
adapter,
payload,
model=base_model,
clip_model=clip_model,
vae_key_map=vae_key_map,
key_maps = _ResolutionKeyMaps(base_model, clip_model)
interned: list[
tuple[RegionalLoraHostPayload, ComfyRegionalAdapterResolution]
] = []
results: list[ComfyRegionalAdapterResolution] = []
for adapter, payload in zip(
adaptation.plan.adapters,
adaptation.adapter_payloads,
strict=True,
):
cached = next(
(
result
for existing_payload, result in interned
if _same_payload(existing_payload, payload)
),
None,
)
for adapter, payload in zip(
adaptation.plan.adapters,
adaptation.adapter_payloads,
strict=True,
)
)
if cached is None:
resolved = self._resolve_adapter(
adapter,
payload,
key_maps=key_maps,
vae_key_map=vae_key_map,
)
interned.append((payload, resolved))
else:
resolved = _rebind_resolution(cached, adapter, payload)
results.append(resolved)
return ComfyRegionalLoraResolution(
adapters=results,
adapters=tuple(results),
issues=tuple(issue for result in results for issue in result.issues),
)
@@ -72,8 +114,7 @@ class ComfyRegionalAdapterResolver:
adapter: RegionalLoraAdapterPlan,
payload: RegionalLoraHostPayload,
*,
model: object,
clip_model: object | None,
key_maps: _ResolutionKeyMaps,
vae_key_map: Mapping[str, object] | None,
) -> ComfyRegionalAdapterResolution:
"""Resolve one payload while converting exceptions into scoped evidence."""
@@ -82,8 +123,7 @@ class ComfyRegionalAdapterResolver:
return self._resolve_adapter_or_raise(
adapter,
payload,
model=model,
clip_model=clip_model,
key_maps=key_maps,
vae_key_map=vae_key_map,
)
except Exception as error:
@@ -120,8 +160,7 @@ class ComfyRegionalAdapterResolver:
adapter: RegionalLoraAdapterPlan,
payload: RegionalLoraHostPayload,
*,
model: object,
clip_model: object | None,
key_maps: _ResolutionKeyMaps,
vae_key_map: Mapping[str, object] | None,
) -> ComfyRegionalAdapterResolution:
"""Use only Comfy key maps and decoding for one valid host payload."""
@@ -130,16 +169,16 @@ class ComfyRegionalAdapterResolver:
raw_weights = _require_source_mapping(payload.raw_weights)
model_weights = comfy.lora.load_lora(
raw_weights,
comfy.lora.model_lora_keys_unet(model, {}),
key_maps.model_map(),
log_missing=False,
)
clip_weights = (
comfy.lora.load_lora(
raw_weights,
comfy.lora.model_lora_keys_clip(clip_model, {}),
key_maps.clip_map(),
log_missing=False,
)
if clip_model is not None
if key_maps.clip_model is not None
else {}
)
vae_weights = (
@@ -206,4 +245,37 @@ def _optional_mapping(value: object | None) -> Mapping[object, object]:
return {} if value is None else _require_mapping(value)
def _same_payload(
left: RegionalLoraHostPayload,
right: RegionalLoraHostPayload,
) -> bool:
"""Match only exact host payload references and resolution representation."""
return (
left.needs_resolution is right.needs_resolution
and left.raw_weights is right.raw_weights
and left.model_weights is right.model_weights
and left.clip_weights is right.clip_weights
)
def _rebind_resolution(
cached: ComfyRegionalAdapterResolution,
adapter: RegionalLoraAdapterPlan,
payload: RegionalLoraHostPayload,
) -> ComfyRegionalAdapterResolution:
"""Reuse immutable payload evidence while preserving authored issue scope."""
return ComfyRegionalAdapterResolution(
adapter=adapter,
payload=payload,
model_targets=cached.model_targets,
source_entries=cached.source_entries,
issues=tuple(
issue(adapter, observed.code, observed.message)
for observed in cached.issues
),
)
COMFY_REGIONAL_ADAPTER_RESOLVER = ComfyRegionalAdapterResolver()
@@ -0,0 +1,104 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Emit opt-in timing evidence for standard-UNet cold-path stages."""
from __future__ import annotations
import logging
import time
from collections.abc import Callable, Iterator
from contextlib import contextmanager
from enum import StrEnum
import torch
ColdMetadataValue = bool | float | int | str
ColdMetadata = dict[str, ColdMetadataValue]
Clock = Callable[[], int]
Synchronizer = Callable[[torch.device | None], None]
_LOGGER = logging.getLogger(
"simple_syrup.runtime.regional_lora.standard_unet_cold_path"
)
class StandardUnetColdStage(StrEnum):
"""Identify non-overlapping or explicitly aggregate first-use stages."""
ADMISSION_RESOLUTION = "admission_resolution"
VARIANT_MATERIALIZATION = "variant_materialization"
VARIANT_SHELL = "variant_shell"
TEMPLATE_PREPARATION = "template_preparation"
MODEL_RESIDENCY = "model_residency"
SAMPLING = "sampling"
class StandardUnetColdPathDiagnosticsEmitter:
"""Measure exact stage boundaries only while a DEBUG capture is active."""
def __init__(
self,
*,
logger: logging.Logger | None = None,
clock_ns: Clock = time.perf_counter_ns,
synchronize: Synchronizer | None = None,
) -> None:
"""Retain injected timing boundaries for deterministic characterization."""
if logger is not None and not isinstance(logger, logging.Logger):
raise TypeError("Standard UNet cold diagnostics logger is invalid.")
if not callable(clock_ns):
raise TypeError("Standard UNet cold diagnostics clock is invalid.")
if synchronize is not None and not callable(synchronize):
raise TypeError("Standard UNet cold diagnostics synchronizer is invalid.")
self._logger = logger or _LOGGER
self._clock_ns = clock_ns
self._synchronize = synchronize or _synchronize_device
@contextmanager
def measure(
self,
stage: StandardUnetColdStage,
*,
device: torch.device | None = None,
) -> Iterator[ColdMetadata]:
"""Yield mutable bounded metadata and emit one synchronized observation."""
if not isinstance(stage, StandardUnetColdStage):
raise TypeError("Standard UNet cold diagnostic stage is invalid.")
if device is not None and not isinstance(device, torch.device):
raise TypeError("Standard UNet cold diagnostic device is invalid.")
metadata: ColdMetadata = {}
if not self._logger.isEnabledFor(logging.DEBUG):
yield metadata
return
self._synchronize(device)
started_at_ns = self._clock_ns()
try:
yield metadata
finally:
self._synchronize(device)
elapsed_ms = (self._clock_ns() - started_at_ns) / 1_000_000.0
self._logger.debug(
"Measured standard UNet cold-path stage",
extra={
"operation": "standard_unet_cold_path.measure",
"cold_path_diagnostics": {
"stage": stage.value,
"elapsed_ms": elapsed_ms,
**metadata,
},
},
)
def _synchronize_device(device: torch.device | None) -> None:
"""Synchronize only an available CUDA boundary selected by the stage owner."""
if device is not None and device.type == "cuda" and torch.cuda.is_available():
torch.cuda.synchronize(device)
STANDARD_UNET_COLD_PATH_DIAGNOSTICS = StandardUnetColdPathDiagnosticsEmitter()
@@ -0,0 +1,73 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Measure opt-in denoising time for persistent standard-UNet variants."""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
import torch
from comfy.model_patcher import ModelPatcher
from comfy.patcher_extension import WrappersMP
from ..model_patcher_mutations import ModelKeyedWrapperMutation
from .standard_unet_cold_diagnostics import (
STANDARD_UNET_COLD_PATH_DIAGNOSTICS,
StandardUnetColdStage,
)
_WRAPPER_KEY = "simple_syrup.standard_unet_cold_sampling"
@dataclass(frozen=True, slots=True)
class StandardUnetColdSamplingDiagnosticsMutation:
"""Install one transparent sampler wrapper for opt-in cold attribution."""
def apply(self, model: object) -> None:
"""Attach the keyed wrapper without changing ordinary sampling behavior."""
if not isinstance(model, ModelPatcher):
raise TypeError("Standard UNet cold sampling requires a MODEL.")
ModelKeyedWrapperMutation(
WrappersMP.SAMPLER_SAMPLE,
_WRAPPER_KEY,
self.measure_sampling,
).apply(model)
def measure_sampling(
self,
executor: Callable[..., object],
guider: object,
sigmas: object,
extra_args: object,
callback: object,
noise: object,
*args: object,
**kwargs: object,
) -> object:
"""Measure one complete sampler call while forwarding exact arguments."""
if not callable(executor):
raise TypeError("Standard UNet cold sampling executor is invalid.")
device = noise.device if isinstance(noise, torch.Tensor) else None
with STANDARD_UNET_COLD_PATH_DIAGNOSTICS.measure(
StandardUnetColdStage.SAMPLING,
device=device,
) as metadata:
result = executor(
guider,
sigmas,
extra_args,
callback,
noise,
*args,
**kwargs,
)
if isinstance(sigmas, torch.Tensor):
metadata["sigma_count"] = sigmas.numel()
if isinstance(noise, torch.Tensor):
metadata["latent_batch_size"] = int(noise.shape[0])
return result
@@ -9,6 +9,8 @@ from __future__ import annotations
from dataclasses import dataclass
from typing import Protocol
import torch
from ..attention_coupling.family_admission import AttentionCouplingFamilyAdmission
from ..regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from .comfy_adapter_resolution import (
@@ -16,6 +18,10 @@ from .comfy_adapter_resolution import (
ComfyRegionalLoraResolution,
)
from .comfy_adapter_resolver import COMFY_REGIONAL_ADAPTER_RESOLVER
from .standard_unet_cold_diagnostics import (
STANDARD_UNET_COLD_PATH_DIAGNOSTICS,
StandardUnetColdStage,
)
_BLOCKING_ISSUE_CODES = frozenset(
{
@@ -78,8 +84,18 @@ class StandardUnetNativeLoraAdmissionService:
raise TypeError("Native standard-UNet admission requires adaptation.")
if not adaptation.plan.adapters:
return StandardUnetNativeLoraAdmission(adaptation, None)
resolution = self._resolver.resolve(adaptation, model=model)
self._validate(adaptation, resolution)
device = getattr(model, "load_device", None)
measured_device = device if isinstance(device, torch.device) else None
with STANDARD_UNET_COLD_PATH_DIAGNOSTICS.measure(
StandardUnetColdStage.ADMISSION_RESOLUTION,
device=measured_device,
) as metadata:
resolution = self._resolver.resolve(adaptation, model=model)
self._validate(adaptation, resolution)
metadata["adapter_count"] = len(resolution.adapters)
metadata["target_count"] = sum(
len(result.model_targets) for result in resolution.adapters
)
return StandardUnetNativeLoraAdmission(adaptation, resolution)
@staticmethod
@@ -13,6 +13,13 @@ from comfy import float as comfy_float
from comfy import lora, model_management, utils
from comfy.model_patcher import ModelPatcher, get_key_weight
from .standard_unet_cold_diagnostics import (
STANDARD_UNET_COLD_PATH_DIAGNOSTICS,
StandardUnetColdStage,
)
from .standard_unet_variant_materialization_device import (
STANDARD_UNET_VARIANT_MATERIALIZATION_DEVICE,
)
from .standard_unet_variant_topology import StandardUnetRegionalVariant
_DIFFUSION_PREFIX = "diffusion_model."
@@ -54,6 +61,9 @@ class StandardUnetMaterializedVariant:
paths = tuple(parameter.path for parameter in self.parameters)
if paths != tuple(sorted(set(paths))):
raise ValueError("Materialized variant paths must be unique and sorted.")
devices = {parameter.tensor.device for parameter in self.parameters}
if len(devices) != 1:
raise ValueError("Materialized variant parameters must share one device.")
class StandardUnetVariantMaterializer:
@@ -72,6 +82,37 @@ class StandardUnetVariantMaterializer:
if not isinstance(variant, StandardUnetRegionalVariant):
raise TypeError("Standard UNet variant materialization requires a variant.")
self._validate_multipliers(variant, schedule_multipliers)
device = STANDARD_UNET_VARIANT_MATERIALIZATION_DEVICE.resolve(model)
with STANDARD_UNET_COLD_PATH_DIAGNOSTICS.measure(
StandardUnetColdStage.VARIANT_MATERIALIZATION,
device=device,
) as metadata:
result = self._materialize_validated(
model,
variant,
schedule_multipliers,
device=device,
)
metadata["region_index"] = variant.region_index
metadata["adapter_count"] = len(variant.adapters)
metadata["parameter_count"] = len(result.parameters)
metadata["device_type"] = device.type
metadata["parameter_bytes"] = sum(
parameter.tensor.numel() * parameter.tensor.element_size()
for parameter in result.parameters
)
return result
def _materialize_validated(
self,
model: ModelPatcher,
variant: StandardUnetRegionalVariant,
schedule_multipliers: tuple[float, ...],
*,
device: torch.device,
) -> StandardUnetMaterializedVariant:
"""Construct one validated variant through conventional Comfy math."""
patches_by_key: dict[str, list[tuple[object, ...]]] = {}
for adapter in variant.adapters:
strength = (
@@ -100,6 +141,7 @@ class StandardUnetVariantMaterializer:
key,
patches_by_key.get(key, []),
originals=originals,
device=device,
)
for key in target_keys
)
@@ -112,6 +154,7 @@ class StandardUnetVariantMaterializer:
patches: list[tuple[object, ...]],
*,
originals: dict[str, list[object]],
device: torch.device,
) -> StandardUnetVariantParameter:
"""Apply one ordered host patch list and conventional final rounding."""
@@ -140,7 +183,7 @@ class StandardUnetVariantMaterializer:
base_weight = base_entry[0]
temporary = model_management.cast_to_device(
base_weight,
weight.device,
device,
torch.float32,
copy=True,
)
@@ -0,0 +1,32 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Select the conventional Comfy device for standard-UNet patch calculation."""
from __future__ import annotations
import torch
from comfy.model_patcher import ModelPatcher
class StandardUnetVariantMaterializationDevice:
"""Resolve exact patch calculation placement from one Comfy MODEL contract."""
@staticmethod
def resolve(model: object) -> torch.device:
"""Return the validated device Comfy requests for loaded model weights."""
if not isinstance(model, ModelPatcher):
raise TypeError("Standard UNet materialization device requires a MODEL.")
device = model.load_device
if not isinstance(device, torch.device):
raise TypeError("Standard UNet MODEL load device must be a torch device.")
if device.type == "meta":
raise ValueError("Standard UNet variants cannot materialize on meta.")
return device
STANDARD_UNET_VARIANT_MATERIALIZATION_DEVICE = (
StandardUnetVariantMaterializationDevice()
)
@@ -22,6 +22,9 @@ from ..model_patcher_mutations import (
ModelKeyedCallbackMutation,
ModelKeyedWrapperMutation,
)
from .standard_unet_cold_sampling import (
StandardUnetColdSamplingDiagnosticsMutation,
)
from .standard_unet_native_admission import StandardUnetNativeLoraAdmission
from .standard_unet_variant_base_attention import StandardUnetVariantBaseAttention
from .standard_unet_variant_conditioning import (
@@ -70,6 +73,7 @@ class StandardUnetVariantRuntimeMutation:
),
)
execution.prime()
StandardUnetColdSamplingDiagnosticsMutation().apply(model)
ModelKeyedWrapperMutation(
WrappersMP.DIFFUSION_MODEL,
_EXECUTION_WRAPPER_KEY,
@@ -11,6 +11,10 @@ from copy import copy
import torch
from torch import nn
from .standard_unet_cold_diagnostics import (
STANDARD_UNET_COLD_PATH_DIAGNOSTICS,
StandardUnetColdStage,
)
from .standard_unet_variant_materialization import StandardUnetMaterializedVariant
@@ -30,10 +34,16 @@ class StandardUnetVariantShellBuilder:
raise TypeError(
"Standard UNet variant shell requires materialized weights."
)
replacements = {
parameter.path: parameter.tensor for parameter in variant.parameters
}
return self._clone_branch(diffusion_model, replacements, prefix="")
with STANDARD_UNET_COLD_PATH_DIAGNOSTICS.measure(
StandardUnetColdStage.VARIANT_SHELL,
) as metadata:
replacements = {
parameter.path: parameter.tensor for parameter in variant.parameters
}
shell = self._clone_branch(diffusion_model, replacements, prefix="")
metadata["region_index"] = variant.region_index
metadata["parameter_count"] = len(variant.parameters)
return shell
def _clone_branch(
self,
@@ -58,11 +68,7 @@ class StandardUnetVariantShellBuilder:
continue
if parameter is None or not isinstance(parameter, nn.Parameter):
raise TypeError(f"Variant target '{path}' must be a Parameter.")
if (
tensor.shape != parameter.shape
or tensor.dtype != parameter.dtype
or tensor.device != parameter.device
):
if tensor.shape != parameter.shape or tensor.dtype != parameter.dtype:
raise ValueError(f"Variant target '{path}' is incompatible.")
clone._parameters[parameter_name] = nn.Parameter(
tensor,
@@ -16,6 +16,10 @@ from comfy.patcher_extension import WrappersMP
from torch import nn
from ..model_patcher_mutations import ModelKeyedWrapperMutation
from .standard_unet_cold_diagnostics import (
STANDARD_UNET_COLD_PATH_DIAGNOSTICS,
StandardUnetColdStage,
)
from .standard_unet_variant_residency_handoff import (
STANDARD_UNET_VARIANT_RESIDENCY_HANDOFF,
StandardUnetVariantResidencyHandoff,
@@ -105,7 +109,17 @@ class StandardUnetStaticVariantResidency:
"device": str(model.load_device),
},
)
return executor(model, noise_shape, conds, *args, **forwarded)
device = model.load_device
measured_device = device if isinstance(device, torch.device) else None
with STANDARD_UNET_COLD_PATH_DIAGNOSTICS.measure(
StandardUnetColdStage.MODEL_RESIDENCY,
device=measured_device,
) as metadata:
result = executor(model, noise_shape, conds, *args, **forwarded)
metadata["required_bytes"] = self._required_bytes
metadata["force_full_load"] = force_full_load or eligible
metadata["force_offload"] = force_offload
return result
def _is_eligible(
self,
@@ -11,6 +11,7 @@ from dataclasses import dataclass
from itertools import product
from threading import Lock
import torch
from comfy.model_patcher import ModelPatcher
from torch import nn
@@ -21,6 +22,10 @@ from ..model_patcher_mutations import (
)
from ..patcher_lifecycle import PATCHER_LIFECYCLE
from .execution_cache import ModelCloneLineage
from .standard_unet_cold_diagnostics import (
STANDARD_UNET_COLD_PATH_DIAGNOSTICS,
StandardUnetColdStage,
)
from .standard_unet_native_admission import StandardUnetNativeLoraAdmission
from .standard_unet_variant_execution_session import (
StandardUnetVariantExecutionSession,
@@ -88,37 +93,46 @@ class StandardUnetVariantTemplate:
if self._frozen:
raise RuntimeError("Standard UNet variant template is already prepared.")
for variant in self.topology.variants:
multiplier_sets = tuple(
tuple(
dict.fromkeys(
boundary.strength_multiplier for boundary in adapter.schedule
device = getattr(self.model, "load_device", None)
measured_device = device if isinstance(device, torch.device) else None
with STANDARD_UNET_COLD_PATH_DIAGNOSTICS.measure(
StandardUnetColdStage.TEMPLATE_PREPARATION,
device=measured_device,
) as metadata:
for variant in self.topology.variants:
multiplier_sets = tuple(
tuple(
dict.fromkeys(
boundary.strength_multiplier
for boundary in adapter.schedule
)
)
for adapter in variant.adapters
)
for adapter in variant.adapters
)
for local_values in product(*multiplier_sets):
if not any(
adapter.model_strength * multiplier != 0.0
for local_values in product(*multiplier_sets):
if not any(
adapter.model_strength * multiplier != 0.0
for adapter, multiplier in zip(
variant.adapters,
local_values,
strict=True,
)
):
continue
complete = [0.0] * self._adapter_count
for adapter, multiplier in zip(
variant.adapters,
local_values,
strict=True,
)
):
continue
complete = [0.0] * self._adapter_count
for adapter, multiplier in zip(
variant.adapters,
local_values,
strict=True,
):
complete[adapter.composition_index] = multiplier
self._prepare_variant(variant, tuple(complete))
self.root.install_forward(
StandardUnetVariantForward(self.root.module, self.execution_session)
)
self._frozen = True
):
complete[adapter.composition_index] = multiplier
self._prepare_variant(variant, tuple(complete))
self.root.install_forward(
StandardUnetVariantForward(self.root.module, self.execution_session)
)
self._frozen = True
metadata["topology_variant_count"] = len(self.topology.variants)
metadata["materialized_variant_count"] = len(self._variants)
def bind_request(self, source: object) -> ModelPatcher:
"""Bind current request patcher state to the cache-stable static graph."""
+108
View File
@@ -0,0 +1,108 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify the benchmark sampler schema matches the production boundary."""
from __future__ import annotations
import asyncio
import torch
from simple_syrup.domain.conditioning_batch import ConditioningBatch
from simple_syrup.nodes_v3.ksampler_attention_coupling import (
KSamplerAttentionCouplingV3,
)
from tools.attention_coupling_benchmark.comfy_probe import (
attention_coupling_phase_node,
comfy_entrypoint,
)
ProfiledKSamplerAttentionCouplingV3 = (
attention_coupling_phase_node.ProfiledKSamplerAttentionCouplingV3
)
def test_profiled_sampler_preserves_input_and_output_schema() -> None:
"""Change only benchmark identity and dev-only visibility."""
production = KSamplerAttentionCouplingV3.define_schema()
profiled = ProfiledKSamplerAttentionCouplingV3.define_schema()
assert profiled.node_id == "SimpleSyrupBenchmark.ProfiledKSamplerAttentionCoupling"
assert profiled.is_dev_only is True
assert [value.id for value in profiled.inputs] == [
value.id for value in production.inputs
]
assert [value.io_type for value in profiled.inputs] == [
value.io_type for value in production.inputs
]
assert [value.io_type for value in profiled.outputs] == [
value.io_type for value in production.outputs
]
def test_profiled_sampler_is_registered_only_by_benchmark_extension() -> None:
"""Make the diagnostic graph executable without changing product exports."""
extension = asyncio.run(comfy_entrypoint())
nodes = asyncio.run(extension.get_node_list())
assert ProfiledKSamplerAttentionCouplingV3 in nodes
def test_profiled_sampler_bridges_host_conditioning_batches_before_delegation() -> None:
"""Normalize only canonical host values at the sibling-extension boundary."""
calls: list[dict[str, object]] = []
class _Service:
"""Capture inherited-node arguments and return the supplied latent."""
def sample(self, **arguments: object) -> dict[str, object]:
"""Retain exact arguments for namespace assertions."""
calls.append(arguments)
latent = arguments["latent_image"]
assert isinstance(latent, dict)
return latent
host_type = type(
"ConditioningBatch",
(),
{
"__module__": (
"custom_nodes.SimpleSyrup.simple_syrup.domain.conditioning_batch"
)
},
)
positive = host_type()
positive.entries = ("positive",)
negative = host_type()
negative.entries = ("negative",)
original = ProfiledKSamplerAttentionCouplingV3.sampling_service_class
ProfiledKSamplerAttentionCouplingV3.sampling_service_class = _Service # type: ignore[assignment]
latent = {"samples": torch.zeros((1, 4, 2, 2))}
try:
output = ProfiledKSamplerAttentionCouplingV3.execute(
model=object(),
seed=1,
steps=1,
cfg=1.0,
sampler_name="sampler",
scheduler="scheduler",
positive=positive,
negative=negative,
region_masks=object(),
regional_prompt_weight=1.0,
region_mask_feather=0,
latent_image=latent,
denoise=1.0,
)
finally:
ProfiledKSamplerAttentionCouplingV3.sampling_service_class = original
assert output == (latent,)
assert isinstance(calls[0]["positive"], ConditioningBatch)
assert isinstance(calls[0]["negative"], ConditioningBatch)
@@ -0,0 +1,123 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify benchmark-only outer Attention Coupling phase profiling."""
from __future__ import annotations
import logging
from types import SimpleNamespace
from typing import Any, ClassVar
import torch
from simple_syrup.domain.regional_attention_execution import (
RegionalAttentionExecutionMode,
)
from simple_syrup.domain.regional_mask_bank import RegionalMaskBank
from simple_syrup.services.attention_coupling_model_preparation_service import (
AttentionCouplingModelPreparationService,
)
from simple_syrup.services.prepared_attention_coupling_model import (
PreparedAttentionCouplingModel,
)
from tools.attention_coupling_benchmark.comfy_probe import (
attention_coupling_phase_profile,
)
ProfiledAttentionCouplingSamplingService = (
attention_coupling_phase_profile.ProfiledAttentionCouplingSamplingService
)
_LOGGER_NAME = "simple_syrup.runtime.regional_lora.standard_unet_cold_path"
class _SamplingService:
"""Record exact delegate inputs and return one recognizable latent."""
calls: ClassVar[list[dict[str, object]]] = []
output: ClassVar[dict[str, Any]] = {"samples": torch.ones((1, 4, 2, 2))}
def sample(self, **arguments: object) -> dict[str, Any]:
"""Retain all delegate arguments."""
type(self).calls.append(arguments)
return self.output
def test_profiled_service_preserves_two_call_delegation_and_emits_phases(
caplog: Any,
) -> None:
"""Add timing evidence without changing preparation or sampling values."""
preparation_calls: list[dict[str, object]] = []
class _PreparationService(AttentionCouplingModelPreparationService):
"""Return one fixed prepared value and retain regional arguments."""
def prepare(
self,
**arguments: object,
) -> PreparedAttentionCouplingModel:
"""Record exact arguments and return a derived model triple."""
preparation_calls.append(arguments)
masks = torch.ones((1, 2, 2))
return PreparedAttentionCouplingModel(
"derived",
"prepared-positive",
"prepared-negative",
RegionalMaskBank(masks, masks.clone(), 2, 2),
)
original_preparation = (
ProfiledAttentionCouplingSamplingService.model_preparation_service_class
)
original_sampling = ProfiledAttentionCouplingSamplingService.sampling_service_class
ProfiledAttentionCouplingSamplingService.model_preparation_service_class = (
_PreparationService
)
ProfiledAttentionCouplingSamplingService.sampling_service_class = _SamplingService # type: ignore[assignment]
_SamplingService.calls = []
latent = {"samples": torch.zeros((1, 4, 2, 2))}
model = SimpleNamespace(load_device=torch.device("cpu"))
try:
with caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME):
output = ProfiledAttentionCouplingSamplingService().sample(
model=model,
seed=7,
steps=30,
cfg=5.0,
sampler_name="sampler",
scheduler="scheduler",
positive="positive",
negative="negative",
region_masks="masks",
regional_prompt_weight=1.0,
region_mask_feather=0,
latent_image=latent,
denoise=1.0,
)
finally:
ProfiledAttentionCouplingSamplingService.model_preparation_service_class = (
original_preparation
)
ProfiledAttentionCouplingSamplingService.sampling_service_class = (
original_sampling
)
assert output is _SamplingService.output
assert preparation_calls[0]["execution_mode"] is RegionalAttentionExecutionMode.FULL
assert preparation_calls[0]["model"] is model
delegate = _SamplingService.calls[0]
assert delegate["model"] == "derived"
assert delegate["positive"] == "prepared-positive"
assert delegate["negative"] == "prepared-negative"
assert delegate["latent_image"] is latent
stages = [
record.cold_path_diagnostics["stage"]
for record in caplog.records
if hasattr(record, "cold_path_diagnostics")
]
assert stages == ["model_preparation_total", "ksampler_delegate_total"]
+93
View File
@@ -0,0 +1,93 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify exclusive benchmark capture of cold-stage diagnostics."""
from __future__ import annotations
import logging
import pytest
from tools.attention_coupling_benchmark.comfy_probe.cold_path_capture import (
CaptureColdPathDiagnosticsV3,
ReadColdPathDiagnosticsV3,
)
_LOGGER_NAME = "simple_syrup.runtime.regional_lora.standard_unet_cold_path"
_COMPOSITION_LOGGER_NAME = (
"simple_syrup.runtime.regional_lora.standard_unet_composition"
)
def test_capture_returns_every_ordered_stage_and_restores_logger() -> None:
"""Retain repeated regional stages without deduplication."""
logger = logging.getLogger(_LOGGER_NAME)
original_level = logger.level
model = object()
started = CaptureColdPathDiagnosticsV3.execute(model, "cold-1")
assert started.result[0] is model
logger.debug(
"materialized left",
extra={
"cold_path_diagnostics": {
"stage": "variant_materialization",
"elapsed_ms": 10.0,
}
},
)
logger.debug(
"materialized right",
extra={
"cold_path_diagnostics": {
"stage": "variant_materialization",
"elapsed_ms": 11.0,
}
},
)
logging.getLogger(_COMPOSITION_LOGGER_NAME).debug(
"model call",
extra={"regional_composition": {"stage": "regional"}},
)
result = ReadColdPathDiagnosticsV3.execute({}, "cold-1")
capture = result.ui["cold_path_diagnostics"][0]
assert capture["record_count"] == 2
assert capture["model_call_count"] == 1
assert [record["elapsed_ms"] for record in capture["records"]] == [10.0, 11.0]
assert logger.level == original_level
def test_capture_is_exclusive_and_missing_read_fails_closed() -> None:
"""Prevent overlapping workflows from contaminating unattributed records."""
CaptureColdPathDiagnosticsV3.execute(object(), "cold-exclusive")
with pytest.raises(RuntimeError, match="already active"):
CaptureColdPathDiagnosticsV3.execute(object(), "cold-other")
logger = logging.getLogger(_LOGGER_NAME)
logger.debug(
"sampling",
extra={
"cold_path_diagnostics": {
"stage": "sampling",
"elapsed_ms": 1.0,
}
},
)
ReadColdPathDiagnosticsV3.execute({}, "cold-exclusive")
with pytest.raises(ValueError, match="was not started"):
ReadColdPathDiagnosticsV3.execute({}, "cold-exclusive")
def test_capture_nodes_are_benchmark_only() -> None:
"""Keep cold attribution outside the public SimpleSyrup node surface."""
capture = CaptureColdPathDiagnosticsV3.define_schema()
read = ReadColdPathDiagnosticsV3.define_schema()
assert capture.node_id == "SimpleSyrupBenchmark.CaptureColdPathDiagnostics"
assert read.node_id == "SimpleSyrupBenchmark.ReadColdPathDiagnostics"
assert capture.is_dev_only is True
assert read.is_dev_only is True
+84
View File
@@ -0,0 +1,84 @@
# 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 Comfy regional adapter source classification."""
from __future__ import annotations
import torch
from comfy.weight_adapter.lora import LoRAAdapter
from simple_syrup.runtime.regional_lora.comfy_adapter_evidence import (
classify_raw_sources,
)
from simple_syrup.runtime.regional_lora.comfy_adapter_resolution import (
ComfyAdapterSourceScope,
)
def test_raw_source_classification_preserves_explicit_fallback_and_unused_order() -> (
None
):
"""Retain loaded-key evidence, identity fallback, and authored source order."""
explicit = torch.ones(1)
fallback = torch.full((1,), 2.0)
unused = torch.full((1,), 3.0)
raw = {
"explicit": explicit,
"fallback": fallback,
"unused": unused,
}
operation = LoRAAdapter({"explicit"}, (explicit, explicit, None, None, None, None))
entries = classify_raw_sources(
raw,
model_weights={
"diffusion_model.explicit.weight": operation,
"diffusion_model.fallback.weight": ("diff", (fallback,)),
},
clip_weights={},
vae_weights={},
)
assert tuple((entry.source_key, entry.scope) for entry in entries) == (
("explicit", ComfyAdapterSourceScope.MODEL),
("fallback", ComfyAdapterSourceScope.MODEL),
("unused", ComfyAdapterSourceScope.UNUSED),
)
def test_raw_source_classification_preserves_scope_precedence() -> None:
"""Keep model, text-encoder, VAE, and unused scope priority exact."""
model_and_clip = torch.ones(1)
clip_and_vae = torch.full((1,), 2.0)
vae = torch.full((1,), 3.0)
unused = torch.full((1,), 4.0)
raw = {
"model-and-clip": model_and_clip,
"clip-and-vae": clip_and_vae,
"vae": vae,
"unused": unused,
}
entries = classify_raw_sources(
raw,
model_weights={"model": ("diff", (model_and_clip,))},
clip_weights={
"model": ("diff", (model_and_clip,)),
"clip": ("diff", (clip_and_vae,)),
},
vae_weights={
"clip": ("diff", (clip_and_vae,)),
"vae": ("diff", (vae,)),
},
)
assert tuple(entry.scope for entry in entries) == (
ComfyAdapterSourceScope.MODEL,
ComfyAdapterSourceScope.TEXT_ENCODER,
ComfyAdapterSourceScope.VAE,
ComfyAdapterSourceScope.UNUSED,
)
@@ -0,0 +1,58 @@
# 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 normalized Comfy adapter identity indexing."""
from __future__ import annotations
import torch
from comfy.weight_adapter.lora import LoRAAdapter
from simple_syrup.runtime.regional_lora.comfy_adapter_identity_index import (
ComfyAdapterIdentityIndex,
)
def test_identity_index_finds_roots_nested_containers_and_adapter_weights() -> None:
"""Traverse exactly the established tuple, list, and adapter-weight edges."""
root = object()
nested = object()
adapter_weight = torch.ones(1)
adapter = LoRAAdapter(
{"adapter"},
(adapter_weight, adapter_weight, None, None, None, None),
)
index = ComfyAdapterIdentityIndex.build((root, ("diff", [nested]), adapter))
assert index.contains(root)
assert index.contains(nested)
assert index.contains(adapter)
assert index.contains(adapter_weight)
def test_identity_index_uses_identity_and_does_not_descend_into_mappings() -> None:
"""Reject equality substitution and preserve the established mapping boundary."""
retained = torch.ones(1)
equal_but_distinct = torch.ones(1)
nested_in_mapping = object()
mapping = {"nested": nested_in_mapping}
index = ComfyAdapterIdentityIndex.build((retained, mapping))
assert index.contains(retained)
assert not index.contains(equal_but_distinct)
assert index.contains(mapping)
assert not index.contains(nested_in_mapping)
def test_identity_index_handles_cyclic_admitted_containers() -> None:
"""Index cyclic host evidence once without recursive failure."""
cycle: list[object] = []
cycle.append(cycle)
index = ComfyAdapterIdentityIndex.build((cycle,))
assert index.contains(cycle)
+99
View File
@@ -0,0 +1,99 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify ordered loopback Comfy execution trace aggregation."""
from __future__ import annotations
import pytest
from tools.comfy_api import LoopbackComfyClient
from tools.comfy_integration.execution_trace import (
ComfyExecutionTraceAccumulator,
)
def test_client_exposes_its_matching_loopback_websocket_session() -> None:
"""Keep HTTP submission and WebSocket events on one generated client id."""
client = LoopbackComfyClient("http://127.0.0.1:8188")
assert client.websocket_url.startswith("ws://127.0.0.1:8188/ws?clientId=")
def test_accumulator_filters_foreign_events_and_aggregates_ordered_nodes() -> None:
"""Measure each node until the next start or terminal success."""
accumulator = _accumulator()
assert not accumulator.observe(
_message("executing", prompt_id="foreign", node="1"),
observed_at_ns=1_010_000_000,
)
assert not accumulator.observe(
_message("execution_cached", nodes=["0"]),
observed_at_ns=1_020_000_000,
)
assert not accumulator.observe(
_message("executing", node="1"),
observed_at_ns=1_100_000_000,
)
assert not accumulator.observe(
_message("executing", node="2"),
observed_at_ns=1_160_000_000,
)
assert accumulator.observe(
_message("execution_success"),
observed_at_ns=1_220_000_000,
)
result = accumulator.finish({"outputs": {}})
assert result.cached_node_ids == ("0",)
assert [
(node.node_id, node.class_type, node.elapsed_ms) for node in result.nodes
] == [
("1", "First", 60.0),
("2", "Second", 60.0),
]
assert result.class_totals_ms() == {"First": 60.0, "Second": 60.0}
assert result.elapsed_ms == 220.0
def test_accumulator_fails_closed_on_execution_error_and_missing_success() -> None:
"""Surface runtime errors and incomplete traces without partial evidence."""
accumulator = _accumulator()
with pytest.raises(RuntimeError, match="RuntimeError: failed"):
accumulator.observe(
_message(
"execution_error",
exception_type="RuntimeError",
exception_message="failed",
),
observed_at_ns=1_100_000_000,
)
with pytest.raises(TimeoutError, match="terminal success"):
accumulator.finish({})
def _accumulator() -> ComfyExecutionTraceAccumulator:
"""Return one identity-neutral two-node trace fixture."""
return ComfyExecutionTraceAccumulator(
prompt_id="prompt-1",
prompt={
"1": {"class_type": "First", "inputs": {}},
"2": {"class_type": "Second", "inputs": {}},
},
submission_started_at_ns=1_000_000_000,
)
def _message(event_type: str, **data: object) -> dict[str, object]:
"""Build one prompt-scoped WebSocket lifecycle message."""
return {
"type": event_type,
"data": {"prompt_id": data.pop("prompt_id", "prompt-1"), **data},
}
@@ -156,6 +156,120 @@ def test_resolver_consumes_initialized_payloads_by_exact_identity(
]
def test_resolver_interns_exact_shared_raw_payload_with_one_host_decode(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Resolve one CFG-shared payload once while preserving authored uses."""
raw_weights = {"up": object(), "down": object()}
payloads = (
RegionalLoraHostPayload.unresolved(raw_weights),
RegionalLoraHostPayload.unresolved(raw_weights),
)
operation = _lora("up", "down")
model_map_calls = 0
load_calls = 0
def model_map(_model: object, _key_map: object) -> dict[str, str]:
nonlocal model_map_calls
model_map_calls += 1
return {"adapter": "diffusion_model.layer.weight"}
def load_lora(
_weights: object,
_key_map: object,
*,
log_missing: bool,
) -> dict[str, LoRAAdapter]:
nonlocal load_calls
assert log_missing is False
load_calls += 1
return {"diffusion_model.layer.weight": operation}
monkeypatch.setattr(comfy.lora, "model_lora_keys_unet", model_map)
monkeypatch.setattr(comfy.lora, "load_lora", load_lora)
result = ComfyRegionalAdapterResolver().resolve(
_adaptation(payloads),
model=_patcher(torch.nn.Linear(1, 1)),
)
assert model_map_calls == 1
assert load_calls == 1
assert len(result.adapters) == 2
assert result.adapters[0].adapter.composition_index == 0
assert result.adapters[1].adapter.composition_index == 1
assert result.adapters[0].payload is payloads[0]
assert result.adapters[1].payload is payloads[1]
assert result.adapters[0].model_targets is result.adapters[1].model_targets
def test_resolver_reuses_host_key_map_across_distinct_raw_payloads(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Scan the runtime model once while decoding each distinct payload once."""
payloads = (
RegionalLoraHostPayload.unresolved({"left": object()}),
RegionalLoraHostPayload.unresolved({"right": object()}),
)
model_map_calls = 0
load_calls = 0
def model_map(_model: object, _key_map: object) -> dict[str, str]:
nonlocal model_map_calls
model_map_calls += 1
return {"adapter": "diffusion_model.layer.weight"}
def load_lora(
_weights: object,
_key_map: object,
*,
log_missing: bool,
) -> dict[str, LoRAAdapter]:
nonlocal load_calls
assert log_missing is False
load_calls += 1
return {"diffusion_model.layer.weight": _lora("up", "down")}
monkeypatch.setattr(comfy.lora, "model_lora_keys_unet", model_map)
monkeypatch.setattr(comfy.lora, "load_lora", load_lora)
ComfyRegionalAdapterResolver().resolve(
_adaptation(payloads),
model=_patcher(torch.nn.Linear(1, 1)),
)
assert model_map_calls == 1
assert load_calls == 2
def test_resolver_rescopes_shared_payload_issues_for_every_authored_use() -> None:
"""Preserve canonical composition identity when cached evidence is invalid."""
model_weights = {"diffusion_model.layer.weight": ("set", (object(),))}
payloads = (
RegionalLoraHostPayload(False, None, model_weights, None),
RegionalLoraHostPayload(False, None, model_weights, None),
)
result = ComfyRegionalAdapterResolver().resolve(
_adaptation(payloads),
model=_patcher(torch.nn.Linear(1, 1)),
)
assert result.adapters[0].model_targets is result.adapters[1].model_targets
assert [observed.composition_index for observed in result.issues] == [0, 1]
assert [observed.adapter_identity for observed in result.issues] == [
"adapter-0.safetensors",
"adapter-1.safetensors",
]
assert [observed.code for observed in result.issues] == [
ComfyAdapterResolutionIssueCode.UNSUPPORTED_OPERATION,
ComfyAdapterResolutionIssueCode.UNSUPPORTED_OPERATION,
]
def test_resolver_preserves_normalized_lora_metadata_by_identity() -> None:
"""Retain alpha, middle, reshape, and tensor identities for U3 translation."""
+91
View File
@@ -0,0 +1,91 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify canonical conditioning-batch normalization across host namespaces."""
from __future__ import annotations
import pytest
from simple_syrup.domain.conditioning_batch import ConditioningBatch
from tools.attention_coupling_benchmark.comfy_probe.conditioning_batch_bridge import (
normalize_conditioning_batch,
)
def test_local_batch_is_preserved_and_plain_conditioning_passes_through() -> None:
"""Avoid copying local domain values or reinterpreting ordinary conditioning."""
batch = ConditioningBatch((object(),))
conditioning = object()
assert normalize_conditioning_batch(batch) is batch
assert normalize_conditioning_batch(conditioning) is conditioning
def test_canonical_host_namespace_is_rebuilt_as_the_local_domain_type() -> None:
"""Bridge only the same immutable domain surface loaded by Comfy's host."""
host_type = type(
"ConditioningBatch",
(),
{
"__module__": (
"custom_nodes.SimpleSyrup.simple_syrup.domain.conditioning_batch"
)
},
)
host_value = host_type()
entries = (object(), object())
host_value.entries = entries
normalized = normalize_conditioning_batch(host_value)
assert isinstance(normalized, ConditioningBatch)
assert normalized.entries == entries
@pytest.mark.parametrize(
("name", "module", "entries", "message"),
[
(
"OtherBatch",
"simple_syrup.domain.conditioning_batch",
(object(),),
"runtime type",
),
(
"ConditioningBatch",
"example.conditioning_batch",
(object(),),
"runtime type",
),
(
"ConditioningBatch",
"simple_syrup.domain.conditioning_batch",
[object()],
"immutable tuple",
),
(
"ConditioningBatch",
"simple_syrup.domain.conditioning_batch",
(),
"must not be empty",
),
],
)
def test_malformed_or_noncanonical_batch_surfaces_fail_closed(
name: str,
module: str,
entries: object,
message: str,
) -> None:
"""Reject lookalikes and mutable or empty canonical surfaces."""
value_type = type(name, (), {"__module__": module})
value = value_type()
value.entries = entries
with pytest.raises((TypeError, ValueError), match=message):
normalize_conditioning_batch(value)
@@ -68,65 +68,3 @@ def test_snapshot_node_is_output_only_and_dev_only() -> None:
assert schema.node_id == "SimpleSyrupBenchmark.SnapshotConditioningBatch"
assert schema.is_output_node is True
assert schema.is_dev_only is True
def test_snapshot_accepts_the_canonical_host_namespace_surface() -> None:
"""Recognize the immutable domain value when Comfy prefixes its module name."""
module_name = "custom_nodes.SimpleSyrup.simple_syrup.domain.conditioning_batch"
host_type = type(
"ConditioningBatch",
(),
{"__module__": module_name},
)
host_value = host_type()
host_value.entries = ([[torch.ones((1, 1, 1)), {}]],)
snapshot = snapshot_conditioning_value(host_value)
assert snapshot["kind"] == "conditioning_batch"
@pytest.mark.parametrize(
("name", "module", "entries", "message"),
[
(
"OtherBatch",
"simple_syrup.domain.conditioning_batch",
(object(),),
"runtime type",
),
(
"ConditioningBatch",
"example.conditioning_batch",
(object(),),
"runtime type",
),
(
"ConditioningBatch",
"simple_syrup.domain.conditioning_batch",
[object()],
"immutable tuple",
),
(
"ConditioningBatch",
"simple_syrup.domain.conditioning_batch",
(),
"must not be empty",
),
],
)
def test_snapshot_rejects_noncanonical_or_malformed_batch_surfaces(
name: str,
module: str,
entries: object,
message: str,
) -> None:
"""Fail closed on lookalike types and mutable or empty batch state."""
value_type = type(name, (), {"__module__": module})
value = value_type()
value.entries = entries
with pytest.raises((TypeError, ValueError), match=message):
snapshot_conditioning_value(value)
+90
View File
@@ -0,0 +1,90 @@
# 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 benchmark-only interop validation profiling."""
from __future__ import annotations
import logging
from types import SimpleNamespace
from typing import Any, cast
import torch
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.regional_attention_execution import (
RegionalAttentionExecutionMode,
)
from simple_syrup.domain.regional_model_capabilities import RegionalModelCapabilities
from simple_syrup.runtime.regional_model_patch_interop import (
RegionalModelPatchInteropReport,
RegionalModelPatchInteropValidator,
)
from tools.attention_coupling_benchmark.comfy_probe.interop_validation_profile import (
ProfiledRegionalModelPatchInteropValidator,
)
_LOGGER_NAME = "simple_syrup.runtime.regional_lora.standard_unet_cold_path"
def test_profiled_interop_validator_preserves_both_exact_delegates(
monkeypatch: Any,
caplog: Any,
) -> None:
"""Retain arguments, report identity, result, and ordered timing stages."""
report = cast(RegionalModelPatchInteropReport, object())
capabilities = cast(RegionalModelCapabilities, object())
processed = cast(ProcessedRegionalAttentionPlan, object())
model = SimpleNamespace(load_device=torch.device("cpu"))
calls: list[tuple[str, tuple[object, ...]]] = []
def validate(
_owner: object,
supplied_model: object,
supplied_capabilities: RegionalModelCapabilities,
) -> RegionalModelPatchInteropReport:
calls.append(("validate", (supplied_model, supplied_capabilities)))
return report
def validate_execution(
_owner: object,
supplied_report: RegionalModelPatchInteropReport,
supplied_plan: ProcessedRegionalAttentionPlan,
supplied_mode: RegionalAttentionExecutionMode,
) -> None:
calls.append(
("validate_execution", (supplied_report, supplied_plan, supplied_mode))
)
monkeypatch.setattr(RegionalModelPatchInteropValidator, "validate", validate)
monkeypatch.setattr(
RegionalModelPatchInteropValidator,
"validate_execution",
validate_execution,
)
validator = ProfiledRegionalModelPatchInteropValidator()
with caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME):
actual = validator.validate(model, capabilities)
validator.validate_execution(
report,
processed,
RegionalAttentionExecutionMode.FULL,
)
assert actual is report
assert calls == [
("validate", (model, capabilities)),
(
"validate_execution",
(report, processed, RegionalAttentionExecutionMode.FULL),
),
]
assert [
record.cold_path_diagnostics["stage"]
for record in caplog.records
if hasattr(record, "cold_path_diagnostics")
] == ["interop_validation", "interop_execution_validation"]
+96
View File
@@ -0,0 +1,96 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify the benchmark-only materialization-parity node boundary."""
from __future__ import annotations
import asyncio
from dataclasses import dataclass
import pytest
from tools.attention_coupling_benchmark.comfy_probe import comfy_entrypoint
from tools.attention_coupling_benchmark.comfy_probe.materialization_parity_node import (
CompareMaterializationParityV3,
)
@dataclass(frozen=True)
class _Observation:
"""Return one stable fake probe payload."""
def as_json_object(self) -> dict[str, object]:
"""Return exact fake numerical evidence."""
return {"exact": True, "element_count": 4}
class _Probe:
"""Capture the node-to-probe call without model execution."""
def __init__(self) -> None:
"""Initialize empty captured keyword arguments."""
self.arguments: dict[str, object] | None = None
def compare(self, **arguments: object) -> _Observation:
"""Retain all arguments and return one fixed observation."""
self.arguments = arguments
return _Observation()
def test_node_delegates_exact_inputs_and_publishes_one_ui_object(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Keep graph adaptation in the probe and JSON publication in the node."""
probe = _Probe()
monkeypatch.setattr(CompareMaterializationParityV3, "probe", probe)
values = {
"model": object(),
"positive": object(),
"negative": object(),
"region_masks": object(),
"latent_image": object(),
"region_mask_feather": 12,
}
result = CompareMaterializationParityV3.execute(
model=values["model"],
positive=values["positive"],
negative=values["negative"],
region_masks=values["region_masks"],
latent_image=values["latent_image"],
region_mask_feather=12,
run_id="parity-1",
)
assert probe.arguments == values
assert result.ui["materialization_parity"] == [
{"exact": True, "element_count": 4, "run_id": "parity-1"}
]
assert '"run_id":"parity-1"' in result.result[0]
def test_node_schema_is_registered_only_for_benchmark_execution() -> None:
"""Expose the focused terminal through the dev-only probe extension."""
schema = CompareMaterializationParityV3.define_schema()
extension = asyncio.run(comfy_entrypoint())
nodes = asyncio.run(extension.get_node_list())
assert schema.node_id == "SimpleSyrupBenchmark.CompareMaterializationParity"
assert schema.is_output_node is True
assert schema.is_dev_only is True
assert tuple(item.io_type for item in schema.inputs[1].io_types) == (
"CONDITIONING",
"CONDITIONING_BATCH",
)
assert tuple(item.io_type for item in schema.inputs[2].io_types) == (
"CONDITIONING",
"CONDITIONING_BATCH",
)
assert CompareMaterializationParityV3 in nodes
+150
View File
@@ -0,0 +1,150 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify materialization-parity schedule and result aggregation."""
from __future__ import annotations
import pytest
from simple_syrup.domain.regional_lora_plan import (
RegionalLoraAdapterIdentity,
RegionalLoraAdapterPlan,
RegionalLoraBranch,
RegionalLoraScheduleBoundary,
)
from tools.attention_coupling_benchmark.comfy_probe import (
materialization_parity_probe,
materialized_variant_comparison,
)
MaterializationParityObservation = (
materialization_parity_probe.MaterializationParityObservation
)
MaterializationParityProbe = materialization_parity_probe.MaterializationParityProbe
MaterializationParityVariantObservation = (
materialization_parity_probe.MaterializationParityVariantObservation
)
MaterializedVariantComparison = (
materialized_variant_comparison.MaterializedVariantComparison
)
def test_time_invariant_schedule_values_follow_composition_order() -> None:
"""Preserve every canonical adapter's exact effective multiplier."""
adapters = (
_adapter(0, 0.75, starts=(0.0, 0.5)),
_adapter(1, 1.25, starts=(0.0,)),
)
observed = MaterializationParityProbe._time_invariant_multipliers(adapters)
assert observed == (0.75, 1.25)
def test_changing_schedule_fails_closed() -> None:
"""Reject a comparison that would conceal multiple materialized states."""
adapter = RegionalLoraAdapterPlan(
adapter_identity=RegionalLoraAdapterIdentity("fixture"),
composition_index=0,
region_index=0,
branch=RegionalLoraBranch.POSITIVE,
model_strength=1.0,
schedule=(
RegionalLoraScheduleBoundary(0.0, 10.0, 1.0, 0),
RegionalLoraScheduleBoundary(0.5, 5.0, 0.5, 0),
),
)
with pytest.raises(ValueError, match="time-invariant"):
MaterializationParityProbe._time_invariant_multipliers((adapter,))
def test_observation_aggregates_weighted_errors_and_exactness() -> None:
"""Publish totals without averaging per-variant means equally."""
first = _comparison(elements=2, differing=1, maximum=0.5, mean=0.25, rms=0.5)
second = _comparison(
elements=6,
differing=2,
maximum=0.25,
mean=0.125,
rms=0.25,
region_index=1,
)
observation = MaterializationParityObservation(
selected_device="cuda:0",
source_load_ms=1.0,
admission_ms=2.0,
peak_vram_bytes=3,
variants=(
MaterializationParityVariantObservation(0, 4.0, 5.0, 6.0, first),
MaterializationParityVariantObservation(1, 7.0, 8.0, 9.0, second),
),
)
payload = observation.as_json_object()
assert payload["element_count"] == 8
assert payload["differing_element_count"] == 3
assert payload["max_absolute_error"] == 0.5
assert payload["mean_absolute_error"] == pytest.approx(0.15625)
assert payload["root_mean_squared_error"] == pytest.approx((0.109375) ** 0.5)
assert payload["exact"] is False
variants = payload["variants"]
assert isinstance(variants, list)
assert variants[0]["comparison"]["exact"] is False
def _adapter(
composition_index: int,
multiplier: float,
*,
starts: tuple[float, ...],
) -> RegionalLoraAdapterPlan:
"""Build one time-invariant domain adapter fixture."""
return RegionalLoraAdapterPlan(
adapter_identity=RegionalLoraAdapterIdentity(f"fixture-{composition_index}"),
composition_index=composition_index,
region_index=composition_index,
branch=RegionalLoraBranch.POSITIVE,
model_strength=1.0,
schedule=tuple(
RegionalLoraScheduleBoundary(
start_percent=start,
start_sigma=10.0 - start,
strength_multiplier=multiplier,
guarantee_steps=0,
)
for start in starts
),
)
def _comparison(
*,
elements: int,
differing: int,
maximum: float,
mean: float,
rms: float,
region_index: int = 0,
) -> MaterializedVariantComparison:
"""Build one aggregate comparison fixture."""
return MaterializedVariantComparison(
region_index=region_index,
parameter_count=1,
element_count=elements,
differing_element_count=differing,
max_absolute_error=maximum,
mean_absolute_error=mean,
root_mean_squared_error=rms,
max_error_parameter_path="weight",
reference_sha256="reference",
candidate_sha256="candidate",
)
@@ -0,0 +1,131 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify bounded comparison of independently materialized variant banks."""
from __future__ import annotations
import math
import pytest
import torch
from simple_syrup.runtime.regional_lora.standard_unet_variant_materialization import (
StandardUnetMaterializedVariant,
StandardUnetVariantParameter,
)
from tools.attention_coupling_benchmark.comfy_probe import (
materialized_variant_comparison,
)
compare_materialized_variants = (
materialized_variant_comparison.compare_materialized_variants
)
def _variant(
region_index: int,
*parameters: tuple[str, torch.Tensor],
) -> StandardUnetMaterializedVariant:
"""Build one canonical detached test bank."""
return StandardUnetMaterializedVariant(
region_index,
tuple(
StandardUnetVariantParameter(path, tensor.detach())
for path, tensor in parameters
),
)
def test_identical_independent_banks_have_equal_hashes_and_zero_error() -> None:
"""Treat equal values as exact without relying on tensor identity."""
reference = _variant(
2,
("block.bias", torch.tensor([1.0, -2.0], dtype=torch.float16)),
("block.weight", torch.arange(12, dtype=torch.float16).reshape(3, 4)),
)
candidate = _variant(
2,
("block.bias", reference.parameters[0].tensor.clone()),
("block.weight", reference.parameters[1].tensor.clone()),
)
result = compare_materialized_variants(reference, candidate, chunk_elements=3)
assert result.region_index == 2
assert result.parameter_count == 2
assert result.element_count == 14
assert result.differing_element_count == 0
assert result.max_absolute_error == 0.0
assert result.mean_absolute_error == 0.0
assert result.root_mean_squared_error == 0.0
assert result.reference_sha256 == result.candidate_sha256
assert result.exact is True
assert result.max_error_parameter_path is None
def test_changed_element_reports_exact_aggregate_error() -> None:
"""Aggregate one changed value without hiding it behind an average."""
reference = _variant(
0,
("block.weight", torch.tensor([1.0, 2.0, 3.0, 4.0])),
)
candidate = _variant(
0,
("block.weight", torch.tensor([1.0, 2.5, 3.0, 4.0])),
)
result = compare_materialized_variants(reference, candidate, chunk_elements=2)
assert result.differing_element_count == 1
assert result.max_absolute_error == 0.5
assert result.mean_absolute_error == 0.125
assert result.root_mean_squared_error == 0.25
assert result.reference_sha256 != result.candidate_sha256
assert result.exact is False
assert result.max_error_parameter_path == "block.weight"
@pytest.mark.parametrize(
("reference", "candidate", "message"),
[
(
_variant(0, ("a", torch.ones(2))),
_variant(0, ("b", torch.ones(2))),
"paths",
),
(
_variant(0, ("a", torch.ones(2))),
_variant(0, ("a", torch.ones(3))),
"shape",
),
(
_variant(0, ("a", torch.ones(2, dtype=torch.float16))),
_variant(0, ("a", torch.ones(2, dtype=torch.float32))),
"dtype",
),
],
)
def test_structural_mismatch_fails_closed(
reference: StandardUnetMaterializedVariant,
candidate: StandardUnetMaterializedVariant,
message: str,
) -> None:
"""Reject banks whose semantic parameter structures do not align."""
with pytest.raises(ValueError, match=message):
compare_materialized_variants(reference, candidate)
def test_nonfinite_parameter_fails_closed() -> None:
"""Avoid publishing meaningless bounded-error metrics for nonfinite values."""
reference = _variant(0, ("a", torch.tensor([1.0, math.inf])))
candidate = _variant(0, ("a", torch.tensor([1.0, math.inf])))
with pytest.raises(ValueError, match="finite"):
compare_materialized_variants(reference, candidate)
+250
View File
@@ -0,0 +1,250 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify transparent benchmark-only model-family profiling."""
from __future__ import annotations
import logging
from types import SimpleNamespace
from typing import Any, cast
import pytest
import torch
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.raw_regional_attention import RawRegionalAttentionPlan
from simple_syrup.domain.regional_model_capabilities import RegionalModelCapabilities
from simple_syrup.runtime.attention_coupling.context_validation import (
RegionalContextValidator,
)
from simple_syrup.runtime.attention_coupling.family_admission import (
AttentionCouplingFamilyAdmission,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from simple_syrup.runtime.regional_model_patch_interop import (
RegionalModelPatchInteropReport,
)
from simple_syrup.services.attention_coupling_model_family import (
AttentionCouplingModelFamily,
AttentionCouplingPreparedModelReuse,
AttentionCouplingSamplerConditioning,
)
from simple_syrup.services.attention_coupling_model_family_selector import (
AttentionCouplingModelFamilySelector,
)
from tools.attention_coupling_benchmark.comfy_probe.model_family_profile import (
ProfiledAttentionCouplingModelFamily,
ProfiledAttentionCouplingModelFamilySelector,
)
_LOGGER_NAME = "simple_syrup.runtime.regional_lora.standard_unet_cold_path"
class _Family:
"""Return recognizable identities while recording exact protocol calls."""
def __init__(self) -> None:
"""Create immutable property identities and empty call storage."""
self.validator = cast(RegionalContextValidator, object())
self.admission = cast(AttentionCouplingFamilyAdmission, object())
self.conditioning = cast(AttentionCouplingSamplerConditioning, object())
self.derived = object()
self.calls: list[tuple[str, tuple[object, ...], dict[str, object]]] = []
@property
def context_validator(self) -> RegionalContextValidator:
"""Return one fixed validator identity."""
return self.validator
@property
def prepared_model_reuse(self) -> AttentionCouplingPreparedModelReuse:
"""Return one fixed reuse policy."""
return AttentionCouplingPreparedModelReuse.EXACT_REQUEST
def validate_latent(self, samples: torch.Tensor) -> None:
"""Record the supplied tensor identity."""
self.calls.append(("validate", (samples,), {}))
def admit_adaptation(
self,
model: object,
adaptation: RegionalLoraPlanAdaptation,
) -> AttentionCouplingFamilyAdmission:
"""Record and return fixed admission evidence."""
self.calls.append(("admit", (model, adaptation), {}))
return self.admission
def prepare_sampler_conditioning(
self,
plan: RawRegionalAttentionPlan,
region_strengths: tuple[float, ...],
) -> AttentionCouplingSamplerConditioning:
"""Record and return fixed sampler conditioning."""
self.calls.append(("conditioning", (plan, region_strengths), {}))
return self.conditioning
def derive(
self,
*,
model: object,
processed_plan: ProcessedRegionalAttentionPlan,
admission: AttentionCouplingFamilyAdmission,
interop_report: RegionalModelPatchInteropReport,
region_strengths: tuple[float, ...],
latent_batch_size: int,
) -> object:
"""Record every named argument and return one fixed model."""
self.calls.append(
(
"derive",
(),
{
"model": model,
"processed_plan": processed_plan,
"admission": admission,
"interop_report": interop_report,
"region_strengths": region_strengths,
"latent_batch_size": latent_batch_size,
},
)
)
return self.derived
def test_profiled_family_preserves_properties_arguments_results_and_order(
caplog: Any,
) -> None:
"""Keep the selected family authoritative behind four timed calls."""
family = _Family()
profiled = ProfiledAttentionCouplingModelFamily(family)
samples = torch.zeros((1, 4, 2, 2))
model = SimpleNamespace(load_device=torch.device("cpu"))
adaptation = cast(RegionalLoraPlanAdaptation, object())
plan = cast(RawRegionalAttentionPlan, object())
processed = cast(ProcessedRegionalAttentionPlan, object())
report = cast(RegionalModelPatchInteropReport, object())
strengths = (1.0, 0.5)
with caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME):
profiled.validate_latent(samples)
admission = profiled.admit_adaptation(model, adaptation)
conditioning = profiled.prepare_sampler_conditioning(plan, strengths)
derived = profiled.derive(
model=model,
processed_plan=processed,
admission=admission,
interop_report=report,
region_strengths=strengths,
latent_batch_size=1,
)
assert profiled.delegate is family
assert profiled.context_validator is family.validator
assert (
profiled.prepared_model_reuse
is AttentionCouplingPreparedModelReuse.EXACT_REQUEST
)
assert admission is family.admission
assert conditioning is family.conditioning
assert derived is family.derived
assert [call[0] for call in family.calls] == [
"validate",
"admit",
"conditioning",
"derive",
]
assert family.calls[0][1] == (samples,)
assert family.calls[1][1] == (model, adaptation)
assert family.calls[2][1] == (plan, strengths)
assert family.calls[3][2] == {
"model": model,
"processed_plan": processed,
"admission": admission,
"interop_report": report,
"region_strengths": strengths,
"latent_batch_size": 1,
}
assert _stages(caplog) == [
"family_validate_latent",
"family_admit_adaptation",
"family_sampler_conditioning",
"family_derive_total",
]
def test_profiled_selector_wraps_exact_production_selected_family(
monkeypatch: Any,
) -> None:
"""Retain capability routing in the production selector implementation."""
family = _Family()
capabilities = cast(RegionalModelCapabilities, object())
calls: list[RegionalModelCapabilities] = []
def select(
_owner: object,
supplied: RegionalModelCapabilities,
) -> AttentionCouplingModelFamily:
calls.append(supplied)
return family
monkeypatch.setattr(AttentionCouplingModelFamilySelector, "select", select)
selected = ProfiledAttentionCouplingModelFamilySelector().select(capabilities)
assert calls == [capabilities]
assert isinstance(selected, ProfiledAttentionCouplingModelFamily)
assert selected.delegate is family
def test_profiled_family_preserves_delegate_exception(caplog: Any) -> None:
"""Propagate an original family error while closing the timed boundary."""
failure = RuntimeError("family admission failed")
class _FailingFamily(_Family):
"""Fail from one exact protocol method."""
def admit_adaptation(
self,
model: object,
adaptation: RegionalLoraPlanAdaptation,
) -> AttentionCouplingFamilyAdmission:
"""Raise the fixed original failure."""
del model, adaptation
raise failure
profiled = ProfiledAttentionCouplingModelFamily(_FailingFamily())
with (
caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME),
pytest.raises(RuntimeError) as raised,
):
profiled.admit_adaptation(
SimpleNamespace(load_device=torch.device("cpu")),
cast(RegionalLoraPlanAdaptation, object()),
)
assert raised.value is failure
assert _stages(caplog) == ["family_admit_adaptation"]
def _stages(caplog: Any) -> list[str]:
"""Return ordered focused stage names from captured records."""
return [
record.cold_path_diagnostics["stage"]
for record in caplog.records
if hasattr(record, "cold_path_diagnostics")
]
@@ -0,0 +1,233 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify exact benchmark-only preparation collaborator profiling."""
from __future__ import annotations
import logging
from types import SimpleNamespace
from typing import Any, cast
import pytest
import torch
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.raw_regional_attention import RawRegionalAttentionPlan
from simple_syrup.runtime.attention_coupling.context_validation import (
RegionalContextValidator,
)
from simple_syrup.runtime.comfy_conditioning_model_loader import (
ComfyConditioningModelLoader,
)
from simple_syrup.runtime.comfy_conditioning_processing import (
ComfyRegionalConditioningProcessor,
)
from simple_syrup.runtime.regional_lora_conditioning_adapter import (
RegionalLoraConditioningAdapter,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from simple_syrup.runtime.regional_model_patch_interop import (
RegionalModelPatchInteropValidator,
)
from simple_syrup.services.attention_coupling_model_family_selector import (
AttentionCouplingModelFamilySelector,
)
from simple_syrup.services.attention_coupling_preparation_service import (
AttentionCouplingPreparation,
AttentionCouplingPreparationService,
)
from tools.attention_coupling_benchmark.comfy_probe import (
preparation_collaborator_profile,
)
from tools.attention_coupling_benchmark.comfy_probe.interop_validation_profile import (
ProfiledRegionalModelPatchInteropValidator,
)
from tools.attention_coupling_benchmark.comfy_probe.model_family_profile import (
ProfiledAttentionCouplingModelFamilySelector,
)
ProfiledAttentionCouplingModelPreparationService = (
preparation_collaborator_profile.ProfiledAttentionCouplingModelPreparationService
)
ProfiledAttentionCouplingPreparationService = (
preparation_collaborator_profile.ProfiledAttentionCouplingPreparationService
)
ProfiledComfyConditioningModelLoader = (
preparation_collaborator_profile.ProfiledComfyConditioningModelLoader
)
ProfiledComfyRegionalConditioningProcessor = (
preparation_collaborator_profile.ProfiledComfyRegionalConditioningProcessor
)
ProfiledRegionalLoraConditioningAdapter = (
preparation_collaborator_profile.ProfiledRegionalLoraConditioningAdapter
)
_LOGGER_NAME = "simple_syrup.runtime.regional_lora.standard_unet_cold_path"
def test_profiled_preparation_service_substitutes_only_exact_collaborators() -> None:
"""Keep production orchestration while replacing its four timed owners."""
service = ProfiledAttentionCouplingModelPreparationService
assert service.model_loader_class is ProfiledComfyConditioningModelLoader
assert service.lora_adapter_class is ProfiledRegionalLoraConditioningAdapter
assert (
service.preparation_service_class is ProfiledAttentionCouplingPreparationService
)
assert (
service.conditioning_processor_class
is ProfiledComfyRegionalConditioningProcessor
)
assert service.interop_validator_class is ProfiledRegionalModelPatchInteropValidator
assert (
service.model_family_selector_class
is ProfiledAttentionCouplingModelFamilySelector
)
assert issubclass(
service.interop_validator_class,
RegionalModelPatchInteropValidator,
)
assert issubclass(
service.model_family_selector_class,
AttentionCouplingModelFamilySelector,
)
def test_profiled_collaborators_preserve_arguments_results_and_stage_order(
monkeypatch: Any,
caplog: Any,
) -> None:
"""Delegate each collaborator exactly and retain every returned identity."""
calls: list[tuple[str, tuple[object, ...], dict[str, object]]] = []
adaptation = cast(RegionalLoraPlanAdaptation, object())
prepared = cast(AttentionCouplingPreparation, object())
processed = cast(ProcessedRegionalAttentionPlan, object())
def load(owner: object, model: object) -> None:
calls.append(("load", (owner, model), {}))
def adapt(
owner: object,
plan: RawRegionalAttentionPlan,
*,
model: object,
) -> RegionalLoraPlanAdaptation:
calls.append(("adapt", (owner, plan), {"model": model}))
return adaptation
def prepare(
owner: object,
plan: RawRegionalAttentionPlan,
) -> AttentionCouplingPreparation:
calls.append(("prepare", (owner, plan), {}))
return prepared
def process(
owner: object,
preparation: AttentionCouplingPreparation,
*,
model: object,
noise: torch.Tensor,
device: torch.device,
context_validator: RegionalContextValidator,
) -> ProcessedRegionalAttentionPlan:
calls.append(
(
"process",
(owner, preparation),
{
"model": model,
"noise": noise,
"device": device,
"context_validator": context_validator,
},
)
)
return processed
monkeypatch.setattr(ComfyConditioningModelLoader, "load", load)
monkeypatch.setattr(RegionalLoraConditioningAdapter, "adapt", adapt)
monkeypatch.setattr(AttentionCouplingPreparationService, "prepare", prepare)
monkeypatch.setattr(ComfyRegionalConditioningProcessor, "process", process)
model = SimpleNamespace(load_device=torch.device("cpu"))
base_model = SimpleNamespace()
plan = cast(RawRegionalAttentionPlan, object())
noise = torch.zeros((1, 4, 2, 2))
device = torch.device("cpu")
validator = cast(RegionalContextValidator, object())
with caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME):
ProfiledComfyConditioningModelLoader().load(model)
actual_adaptation = ProfiledRegionalLoraConditioningAdapter().adapt(
plan,
model=base_model,
)
actual_preparation = ProfiledAttentionCouplingPreparationService().prepare(plan)
actual_processed = ProfiledComfyRegionalConditioningProcessor().process(
prepared,
model=model,
noise=noise,
device=device,
context_validator=validator,
)
assert actual_adaptation is adaptation
assert actual_preparation is prepared
assert actual_processed is processed
assert [call[0] for call in calls] == ["load", "adapt", "prepare", "process"]
assert calls[0][1][1] is model
assert calls[1][1][1] is plan
assert calls[1][2]["model"] is base_model
assert calls[2][1][1] is plan
assert calls[3][1][1] is prepared
assert calls[3][2] == {
"model": model,
"noise": noise,
"device": device,
"context_validator": validator,
}
assert _stages(caplog) == [
"source_model_load",
"regional_lora_adaptation",
"regional_plan_preparation",
"conditioning_processing",
]
def test_profiled_collaborator_preserves_delegate_exception(
monkeypatch: Any,
caplog: Any,
) -> None:
"""Propagate the original failure while still closing its timing phase."""
failure = RuntimeError("source load failed")
def fail(_owner: object, _model: object) -> None:
raise failure
monkeypatch.setattr(ComfyConditioningModelLoader, "load", fail)
with (
caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME),
pytest.raises(RuntimeError) as raised,
):
ProfiledComfyConditioningModelLoader().load(
SimpleNamespace(load_device=torch.device("cpu"))
)
assert raised.value is failure
assert _stages(caplog) == ["source_model_load"]
def _stages(caplog: Any) -> list[str]:
"""Return ordered focused stage names from captured records."""
return [
record.cold_path_diagnostics["stage"]
for record in caplog.records
if hasattr(record, "cold_path_diagnostics")
]
@@ -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 selected post-optimization visual CLI delegation."""
from __future__ import annotations
from pathlib import Path
from typing import Any
from tools import run_sdxl_post_optimization_visual_proof
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
def test_post_optimization_cli_executes_only_selected_case(
monkeypatch: Any,
tmp_path: Path,
) -> None:
"""Pass one explicit case through the shared selector to the executor."""
inventory_path = tmp_path / "inventory.json"
prompts_path = tmp_path / "prompts.json"
inventory = object()
prompts = object()
first = _Case("first")
second = _Case("second")
calls: list[tuple[object, ...]] = []
monkeypatch.setattr(
SdxlVisualInventory,
"load",
lambda path: inventory if path == inventory_path else None,
)
monkeypatch.setattr(
run_sdxl_post_optimization_visual_proof,
"load_visual_prompt_set",
lambda path: prompts if path == prompts_path else None,
)
monkeypatch.setattr(
run_sdxl_post_optimization_visual_proof,
"post_optimization_visual_cases",
lambda supplied_inventory, supplied_prompts: (
(
first,
second,
)
if (supplied_inventory, supplied_prompts) == (inventory, prompts)
else ()
),
)
def execute(
_artifacts: object,
*,
inventory: object,
cases: tuple[object, ...],
comfy_root: Path,
readiness_timeout: float,
prompt_timeout: float,
seed: int,
) -> Path:
calls.append(
(
inventory,
cases,
comfy_root,
readiness_timeout,
prompt_timeout,
seed,
)
)
return tmp_path / "result.json"
monkeypatch.setattr(
run_sdxl_post_optimization_visual_proof,
"execute_visual_cases",
execute,
)
result = run_sdxl_post_optimization_visual_proof.main(
(
"--inventory",
str(inventory_path),
"--prompt-case",
str(prompts_path),
"--comfy-root",
str(tmp_path),
"--output-root",
str(tmp_path / "artifacts"),
"--case-id",
"second",
"--seed",
"7429113058",
)
)
assert result == 0
assert calls == [(inventory, (second,), tmp_path, 240.0, 1200.0, 7429113058)]
class _Case:
"""Expose only the case identity required by the shared selector."""
def __init__(self, case_id: str) -> None:
"""Retain one stable case identity."""
self.case_id = case_id
+57
View File
@@ -0,0 +1,57 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify historical-context SDXL cold-path primers."""
from __future__ import annotations
from tools.sdxl_attention_coupling_integration.visual_case_model import (
RegionalVisualAdapter,
SdxlVisualCase,
)
from tools.sdxl_regional_lora_performance.cold_path_priming import (
build_cold_path_primers,
)
def test_primers_preserve_global_left_right_order_without_images() -> None:
"""Warm only the dependencies present before the historical regional run."""
primers = build_cold_path_primers(
checkpoint_name="checkpoint.safetensors",
mask_names=("left.png", "right.png"),
case=_case(),
)
assert tuple(primer.label for primer in primers) == (
"global-two-lora-reference",
"conventional-left-variant",
"conventional-right-variant",
)
for primer in primers:
class_types = tuple(
node["class_type"] for node in primer.workflow.prompt.values()
)
assert "SimpleSyrupBenchmark.CompleteLatent" in class_types
assert "VAEDecode" not in class_types
assert "SaveImage" not in class_types
assert "SimpleSyrupBenchmark.CaptureColdPathDiagnostics" not in class_types
def _case() -> SdxlVisualCase:
"""Return two anonymous full-strength character adapters."""
return SdxlVisualCase(
"cold-primer",
"Cold primer",
base_positive_g="two subjects",
base_positive_l="two subjects",
left_g="left traits",
left_l="left traits",
right_g="right traits",
right_l="right traits",
left_adapters=(RegionalVisualAdapter("left.safetensors", 1.0, 1.0),),
right_adapters=(RegionalVisualAdapter("right.safetensors", 1.0, 1.0),),
regional_prompt_weight=1.0,
)
+100
View File
@@ -0,0 +1,100 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify cold-path history decoding and non-overlapping attribution."""
from __future__ import annotations
from tools.sdxl_regional_lora_performance.cold_path_results import (
decode_cold_path_timing,
summarize_cold_path,
)
def test_summary_accounts_top_level_stages_without_double_counting_variants() -> None:
"""Keep nested materialization and shell evidence out of the elapsed sum."""
cold = decode_cold_path_timing(
_history(
completed_at_ns=101_000_000,
run_id="cold",
records=[
_stage("admission_resolution", 1.0),
_stage(
"variant_materialization",
10.0,
parameter_count=2,
parameter_bytes=16,
),
_stage("variant_shell", 2.0),
_stage(
"variant_materialization",
11.0,
parameter_count=3,
parameter_bytes=24,
),
_stage("variant_shell", 3.0),
_stage("template_preparation", 30.0),
_stage("model_residency", 40.0),
_stage("sampling", 20.0),
],
),
terminal_node_id="9",
started_at_ns=1_000_000,
seed=1,
)
warm = decode_cold_path_timing(
_history(
completed_at_ns=112_000_000,
run_id="warm",
records=[
_stage("model_residency", 1.0),
_stage("sampling", 9.0),
],
),
terminal_node_id="9",
started_at_ns=102_000_000,
seed=2,
)
summary = summarize_cold_path(cold, warm)
assert summary.cold_runtime_ms == 100.0
assert summary.warm_runtime_ms == 10.0
assert summary.stage_totals_ms["variant_materialization"] == 21.0
assert summary.top_level_accounted_ms == 91.0
assert summary.unattributed_ms == 9.0
assert summary.materialized_parameter_count == 5
assert summary.materialized_parameter_bytes == 40
def _history(
*,
completed_at_ns: int,
run_id: str,
records: list[dict[str, object]],
) -> dict[str, object]:
"""Build one generic Comfy terminal history."""
return {
"outputs": {
"9": {
"cold_path_diagnostics": [
{
"run_id": run_id,
"completed_at_ns": completed_at_ns,
"records": records,
"model_call_count": 30,
"peak_vram_bytes": 1024,
}
]
}
}
}
def _stage(stage: str, elapsed_ms: float, **metadata: object) -> dict[str, object]:
"""Build one structured cold-stage record."""
return {"stage": stage, "elapsed_ms": elapsed_ms, **metadata}
+71
View File
@@ -0,0 +1,71 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify the focused SDXL cold-attribution graph."""
from __future__ import annotations
from tools.comfy_api import JsonObject
from tools.sdxl_attention_coupling_integration.visual_case_model import (
RegionalVisualAdapter,
SdxlVisualCase,
)
from tools.sdxl_regional_lora_performance.cold_path_workflow import (
build_sdxl_cold_path_workflow,
)
def test_cold_workflow_is_image_free_and_revises_only_seed_and_capture() -> None:
"""Keep the measured preparation graph stable across cold and warm runs."""
built = build_sdxl_cold_path_workflow(
checkpoint_name="checkpoint.safetensors",
mask_names=("left.png", "right.png"),
case=_case(),
)
first = built.prompt_for_execution(seed=101, run_id="cold")
second = built.prompt_for_execution(seed=202, run_id="warm")
assert _inputs(first, built.sampler_node_id)["seed"] == 101
assert _inputs(second, built.sampler_node_id)["seed"] == 202
assert _inputs(first, built.capture_node_id)["run_id"] == "cold"
assert _inputs(first, built.terminal_node_id)["run_id"] == "cold"
_inputs(first, built.sampler_node_id)["seed"] = 202
_inputs(first, built.capture_node_id)["run_id"] = "warm"
_inputs(first, built.terminal_node_id)["run_id"] = "warm"
assert first == second
class_types = tuple(node["class_type"] for node in built.prompt.values())
assert class_types.count("SimpleSyrup.KSamplerAttentionCoupling") == 1
assert class_types.count("SimpleSyrupBenchmark.CaptureColdPathDiagnostics") == 1
assert class_types.count("SimpleSyrupBenchmark.ReadColdPathDiagnostics") == 1
assert "VAEDecode" not in class_types
assert "SaveImage" not in class_types
assert "SimpleSyrupBenchmark.InstrumentModel" not in class_types
def _case() -> SdxlVisualCase:
"""Return two anonymous full-strength character adapters."""
return SdxlVisualCase(
"cold-path",
"Cold path",
base_positive_g="two subjects",
base_positive_l="two subjects",
left_g="left traits",
left_l="left traits",
right_g="right traits",
right_l="right traits",
left_adapters=(RegionalVisualAdapter("left.safetensors", 1.0, 1.0),),
right_adapters=(RegionalVisualAdapter("right.safetensors", 1.0, 1.0),),
regional_prompt_weight=1.0,
)
def _inputs(prompt: dict[str, JsonObject], node_id: str) -> dict[str, object]:
"""Narrow one test graph input mapping."""
inputs = prompt[node_id].get("inputs")
if not isinstance(inputs, dict):
raise TypeError("Cold-path test node inputs must be an object.")
return inputs
@@ -0,0 +1,60 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify the cold profile graph substitutes only its sampler identity."""
from __future__ import annotations
from tools.sdxl_attention_coupling_integration.visual_case_model import (
RegionalVisualAdapter,
SdxlVisualCase,
)
from tools.sdxl_cold_upstream_trace.profile_workflow import (
build_profiled_cold_path_workflow,
)
from tools.sdxl_regional_lora_performance.cold_path_workflow import (
build_sdxl_cold_path_workflow,
)
def test_profile_workflow_changes_only_the_sampler_class() -> None:
"""Preserve every graph input, edge, node id, and terminal unchanged."""
case = _case()
production = build_sdxl_cold_path_workflow(
checkpoint_name="checkpoint",
mask_names=("left-mask", "right-mask"),
case=case,
)
profiled = build_profiled_cold_path_workflow(
checkpoint_name="checkpoint",
mask_names=("left-mask", "right-mask"),
case=case,
)
assert production.sampler_node_id == profiled.sampler_node_id
assert production.capture_node_id == profiled.capture_node_id
assert production.terminal_node_id == profiled.terminal_node_id
profiled.prompt[profiled.sampler_node_id]["class_type"] = (
"SimpleSyrup.KSamplerAttentionCoupling"
)
assert profiled.prompt == production.prompt
def _case() -> SdxlVisualCase:
"""Return one model-identity-neutral two-adapter case."""
return SdxlVisualCase(
"profile",
"Profile",
base_positive_g="two subjects",
base_positive_l="two subjects",
left_g="left traits",
left_l="left traits",
right_g="right traits",
right_l="right traits",
left_adapters=(RegionalVisualAdapter("adapter-a", 1.0, 1.0),),
right_adapters=(RegionalVisualAdapter("adapter-b", 1.0, 1.0),),
regional_prompt_weight=1.0,
)
@@ -0,0 +1,81 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify cold upstream node/stage decomposition."""
from __future__ import annotations
import pytest
from tools.comfy_integration.execution_trace import (
ComfyExecutionTraceResult,
ComfyNodeExecutionTiming,
)
from tools.sdxl_cold_upstream_trace.results import attribute_cold_upstream
from tools.sdxl_regional_lora_performance.cold_path_results import (
ColdPathStageTiming,
SdxlColdPathTiming,
)
def test_attribution_separates_upstream_sampler_and_terminal_intervals() -> None:
"""Explain cold wall time without double-counting nested sampler stages."""
trace = ComfyExecutionTraceResult(
prompt_id="prompt",
submission_started_at_ns=1_000_000_000,
completed_at_ns=1_200_000_000,
cached_node_ids=("cached",),
nodes=(
ComfyNodeExecutionTiming("a", "Encode", 1_010_000_000, 40.0),
ComfyNodeExecutionTiming("s", "Sampler", 1_050_000_000, 120.0),
ComfyNodeExecutionTiming("t", "Terminal", 1_170_000_000, 30.0),
),
history={},
)
cold = SdxlColdPathTiming(
run_id="cold",
seed=1,
runtime_ms=200.0,
stages=(
ColdPathStageTiming("admission_resolution", 10.0, {}),
ColdPathStageTiming("template_preparation", 20.0, {}),
ColdPathStageTiming("model_residency", 10.0, {}),
ColdPathStageTiming("sampling", 60.0, {}),
ColdPathStageTiming("variant_materialization", 5.0, {}),
ColdPathStageTiming("variant_shell", 1.0, {}),
),
model_call_count=30,
peak_vram_bytes=1,
)
result = attribute_cold_upstream(trace, cold, sampler_node_id="s")
assert result.submission_to_first_node_ms == 10.0
assert result.pre_sampler_node_ms == 40.0
assert result.sampler_node_ms == 120.0
assert result.sampler_instrumented_ms == 100.0
assert result.sampler_unattributed_ms == 20.0
assert result.post_sampler_node_ms == 30.0
assert result.trace_residual_ms == 0.0
assert result.cold_unattributed_ms == 100.0
assert result.named_upstream_fraction == 0.4
assert result.upstream_class_totals_ms == {"Encode": 40.0}
def test_attribution_requires_one_traced_sampler() -> None:
"""Fail closed when graph identity cannot isolate the sampler interval."""
trace = ComfyExecutionTraceResult(
"prompt",
1,
2,
(),
(ComfyNodeExecutionTiming("other", "Other", 1, 0.0),),
{},
)
cold = SdxlColdPathTiming("cold", 1, 1.0, (), 30, 1)
with pytest.raises(ValueError, match="exactly one sampler"):
attribute_cold_upstream(trace, cold, sampler_node_id="sampler")
@@ -0,0 +1,64 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify fail-closed materialization parity history decoding."""
from __future__ import annotations
import pytest
from tools.sdxl_materialization_parity.results import (
decode_materialization_parity,
)
def _payload() -> dict[str, object]:
"""Return one model-identity-neutral parity result fixture."""
return {
"run_id": "run-1",
"variant_count": 1,
"element_count": 8,
"differing_element_count": 0,
"peak_vram_bytes": 1024,
"exact": True,
"variants": [{"region_index": 0}],
}
def test_decoder_returns_one_complete_matching_result() -> None:
"""Preserve the terminal object after validating its required evidence."""
payload = _payload()
result = decode_materialization_parity(
{"outputs": {"9": {"materialization_parity": [payload]}}},
terminal_node_id="9",
expected_run_id="run-1",
)
assert result == payload
@pytest.mark.parametrize(
"payload",
[
{**_payload(), "run_id": "wrong"},
{**_payload(), "exact": "yes"},
{**_payload(), "variant_count": 2},
{**_payload(), "differing_element_count": 9},
{**_payload(), "element_count": 0},
],
)
def test_decoder_rejects_incomplete_or_inconsistent_evidence(
payload: dict[str, object],
) -> None:
"""Fail before publishing malformed benchmark evidence."""
with pytest.raises(ValueError):
decode_materialization_parity(
{"outputs": {"9": {"materialization_parity": [payload]}}},
terminal_node_id="9",
expected_run_id="run-1",
)
@@ -0,0 +1,62 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify the focused sampler-free materialization parity graph."""
from __future__ import annotations
from tools.comfy_api import JsonObject
from tools.sdxl_attention_coupling_integration.visual_case_model import (
RegionalVisualAdapter,
SdxlVisualCase,
)
from tools.sdxl_materialization_parity.workflow import (
build_materialization_parity_workflow,
)
def test_workflow_reuses_regional_topology_without_sampling_or_decode() -> None:
"""Keep the comparison graph image-free and preserve its semantic inputs."""
case = SdxlVisualCase(
"materialization-parity",
"Materialization parity",
base_positive_g="two subjects",
base_positive_l="two subjects",
left_g="left traits",
left_l="left traits",
right_g="right traits",
right_l="right traits",
left_adapters=(RegionalVisualAdapter("adapter-a", 1.0, 1.0),),
right_adapters=(RegionalVisualAdapter("adapter-b", 1.0, 1.0),),
region_mask_feather=7,
regional_prompt_weight=1.0,
)
built = build_materialization_parity_workflow(
run_id="comparison-1",
checkpoint_name="checkpoint",
mask_names=("mask-a", "mask-b"),
case=case,
)
class_types = tuple(str(node["class_type"]) for node in built.prompt.values())
assert class_types.count("SimpleSyrupBenchmark.CompareMaterializationParity") == 1
assert not any("KSampler" in class_type for class_type in class_types)
assert "VAEDecode" not in class_types
assert "SaveImage" not in class_types
inputs = _inputs(built.prompt, built.terminal_node_id)
assert inputs["run_id"] == "comparison-1"
assert inputs["region_mask_feather"] == 7
assert inputs["positive"] != inputs["negative"]
assert inputs["region_masks"] == inputs["region_masks"]
def _inputs(prompt: dict[str, JsonObject], node_id: str) -> dict[str, object]:
"""Narrow one test graph input mapping."""
inputs = prompt[node_id].get("inputs")
if not isinstance(inputs, dict):
raise TypeError("Materialization parity node inputs must be an object.")
return inputs
+2
View File
@@ -63,6 +63,7 @@ def test_recorder_serializes_prompts_from_the_exact_case(tmp_path: Path) -> None
checkpoint_name=r"owned\checkpoint.safetensors",
mask_names=("left.png", "right.png"),
case=case,
seed=7_429_113_058,
)
history = _history(workflow, case)
recorder.record_case(
@@ -77,6 +78,7 @@ def test_recorder_serializes_prompts_from_the_exact_case(tmp_path: Path) -> None
decoded = json.loads((tmp_path / "u11-result.json").read_text(encoding="utf-8"))
prompts = decoded["observations"][0]["prompts"]
assert decoded["observations"][0]["seed"] == 7_429_113_058
assert prompts == {
"base_positive_g": case.base_positive_g,
"base_positive_l": case.base_positive_l,
+23 -27
View File
@@ -2,40 +2,36 @@
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Characterize the focused SDXL visual sampling controls."""
"""Verify the focused SDXL visual sampling controls."""
from __future__ import annotations
from typing import cast
import pytest
from tools.comfy_api import JsonObject
from tools.sdxl_attention_couple_parity.cases import parity_case
from tools.sdxl_attention_couple_parity.workflow import (
ParityBackend,
build_parity_workflow,
from tools.sdxl_attention_coupling_integration.sampling_controls import (
MAX_COMFY_SEED,
validate_sdxl_visual_seed,
)
def test_candidate_workflow_uses_characterized_sampling_controls() -> None:
"""Preserve the exact base sampler controls through ownership extraction."""
@pytest.mark.parametrize("seed", (0, 7_429_113_058, MAX_COMFY_SEED))
def test_visual_seed_accepts_comfy_sampler_range(seed: int) -> None:
"""Return every integer inside Comfy's declared seed range unchanged."""
workflow = build_parity_workflow(
backend=ParityBackend.CANDIDATE,
run_id="sampling-control-characterization",
checkpoint_name="checkpoint.safetensors",
mask_names=("left.png", "right.png"),
case=parity_case(),
)
assert validate_sdxl_visual_seed(seed) == seed
sampler = next(
node
for node in workflow.prompt.values()
if node["class_type"] == "SimpleSyrup.KSamplerAttentionCoupling"
)
inputs = cast(JsonObject, sampler["inputs"])
assert inputs["seed"] == 7_429_113_057
assert inputs["cfg"] == 5.0
assert inputs["sampler_name"] == "euler_ancestral"
assert inputs["scheduler"] == "karras"
assert inputs["steps"] == 30
@pytest.mark.parametrize("seed", (-1, MAX_COMFY_SEED + 1))
def test_visual_seed_rejects_values_outside_comfy_sampler_range(seed: int) -> None:
"""Reject integers Comfy's sampler schema cannot represent."""
with pytest.raises(ValueError, match="must be between"):
validate_sdxl_visual_seed(seed)
@pytest.mark.parametrize("seed", (True, 1.5, "1"))
def test_visual_seed_rejects_non_integer_values(seed: object) -> None:
"""Reject bool and dynamically supplied non-integer values explicitly."""
with pytest.raises(TypeError, match="must be an integer"):
validate_sdxl_visual_seed(seed)
+42
View File
@@ -51,6 +51,28 @@ def test_baseline_builds_only_one_native_full_sampler(tmp_path: Path) -> None:
assert sampler["region_mask_feather"] == 0
def test_visual_workflow_applies_one_explicit_seed_to_every_sampler(
tmp_path: Path,
) -> None:
"""Keep every mode and durable workflow identity on one requested seed."""
case = next(
item
for item in visual_cases(visual_inventory(tmp_path))
if item.case_id == "spatial-mode-global-style-regional-character"
)
built = build_sdxl_visual_workflow(
run_id="seed-control",
checkpoint_name=r"owned\checkpoint.safetensors",
mask_names=("left.png", "right.png"),
case=case,
seed=7_429_113_058,
)
assert built.seed == 7_429_113_058
assert _sampler_seeds(built.prompt) == {7_429_113_058}
def test_full_regional_control_changes_only_sampler_prompt_weight(
tmp_path: Path,
) -> None:
@@ -158,3 +180,23 @@ def _sampler_inputs(prompt: dict[str, JsonObject]) -> JsonObject:
if len(samplers) != 1 or not isinstance(samplers[0], dict):
raise AssertionError("Expected one full attention-coupling sampler.")
return samplers[0]
def _sampler_seeds(prompt: dict[str, JsonObject]) -> set[int]:
"""Return every explicitly emitted SimpleSyrup sampler seed."""
seeds: set[int] = set()
for node in prompt.values():
class_type = node["class_type"]
if not isinstance(class_type, str) or not class_type.startswith(
"SimpleSyrup.KSampler"
):
continue
inputs = node["inputs"]
if not isinstance(inputs, dict):
raise AssertionError("Sampler inputs must be a JSON object.")
seed = inputs.get("seed")
if not isinstance(seed, int):
raise AssertionError("Sampler seed must be an integer.")
seeds.add(seed)
return seeds
@@ -0,0 +1,88 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify opt-in standard-UNet cold-stage diagnostics."""
from __future__ import annotations
import logging
import torch
from simple_syrup.runtime.regional_lora.standard_unet_cold_diagnostics import (
StandardUnetColdPathDiagnosticsEmitter,
StandardUnetColdStage,
)
def test_enabled_measurement_emits_elapsed_time_and_bounded_metadata() -> None:
"""Publish one structured stage after its measured work completes."""
logger = logging.getLogger("tests.standard_unet_cold.enabled")
logger.setLevel(logging.DEBUG)
ticks = iter((1_000, 4_500))
synchronizations: list[torch.device | None] = []
emitter = StandardUnetColdPathDiagnosticsEmitter(
logger=logger,
clock_ns=lambda: next(ticks),
synchronize=synchronizations.append,
)
handler = _RecordHandler()
logger.addHandler(handler)
try:
with emitter.measure(
StandardUnetColdStage.VARIANT_MATERIALIZATION,
device=torch.device("cpu"),
) as metadata:
metadata["parameter_count"] = 2
metadata["parameter_bytes"] = 16
finally:
logger.removeHandler(handler)
assert synchronizations == [torch.device("cpu"), torch.device("cpu")]
assert len(handler.records) == 1
payload = getattr(handler.records[0], "cold_path_diagnostics", None)
assert payload == {
"stage": "variant_materialization",
"elapsed_ms": 0.0035,
"parameter_count": 2,
"parameter_bytes": 16,
}
def test_disabled_measurement_avoids_clock_and_synchronization() -> None:
"""Keep ordinary execution free of attribution work and device barriers."""
logger = logging.getLogger("tests.standard_unet_cold.disabled")
logger.setLevel(logging.INFO)
def fail_clock() -> int:
raise AssertionError("disabled diagnostics must not read the clock")
def fail_synchronize(_device: torch.device | None) -> None:
raise AssertionError("disabled diagnostics must not synchronize")
emitter = StandardUnetColdPathDiagnosticsEmitter(
logger=logger,
clock_ns=fail_clock,
synchronize=fail_synchronize,
)
with emitter.measure(StandardUnetColdStage.SAMPLING) as metadata:
metadata["model_call_count"] = 30
class _RecordHandler(logging.Handler):
"""Retain exact records from one focused emitter test."""
def __init__(self) -> None:
"""Create an empty record list."""
super().__init__(logging.DEBUG)
self.records: list[logging.LogRecord] = []
def emit(self, record: logging.LogRecord) -> None:
"""Append one emitted diagnostic record."""
self.records.append(record)
+67
View File
@@ -0,0 +1,67 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify transparent standard-UNet cold sampling instrumentation."""
from __future__ import annotations
import torch
from comfy.model_patcher import ModelPatcher
from comfy.patcher_extension import WrappersMP
from torch import nn
from simple_syrup.runtime.regional_lora.standard_unet_cold_sampling import (
StandardUnetColdSamplingDiagnosticsMutation,
)
def test_mutation_installs_one_keyed_sampler_wrapper() -> None:
"""Keep sampling attribution distinct from diffusion execution wrappers."""
patcher = ModelPatcher(nn.Linear(2, 2), torch.device("cpu"), torch.device("cpu"))
StandardUnetColdSamplingDiagnosticsMutation().apply(patcher)
wrappers = patcher.get_wrappers(
WrappersMP.SAMPLER_SAMPLE,
"simple_syrup.standard_unet_cold_sampling",
)
assert len(wrappers) == 1
def test_wrapper_forwards_exact_arguments_and_result() -> None:
"""Make disabled diagnostics observational for ordinary execution."""
received: tuple[object, ...] = ()
def executor(*args: object, **kwargs: object) -> str:
nonlocal received
received = (*args, kwargs)
return "sampled"
mutation = StandardUnetColdSamplingDiagnosticsMutation()
sigmas = torch.tensor([1.0, 0.0])
noise = torch.zeros((1, 4, 8, 8))
result = mutation.measure_sampling(
executor,
"guider",
sigmas,
{"seed": 1},
"callback",
noise,
"latent",
disable_pbar=True,
)
assert result == "sampled"
assert received == (
"guider",
sigmas,
{"seed": 1},
"callback",
noise,
"latent",
{"disable_pbar": True},
)
@@ -6,6 +6,8 @@
from __future__ import annotations
import comfy.model_management
import pytest
import torch
from comfy.model_patcher import ModelPatcher
from torch import nn
@@ -16,7 +18,9 @@ from simple_syrup.runtime.regional_lora.comfy_adapter_resolution import (
ComfyNormalizedAdapterTarget,
)
from simple_syrup.runtime.regional_lora.standard_unet_variant_materialization import (
StandardUnetMaterializedVariant,
StandardUnetVariantMaterializer,
StandardUnetVariantParameter,
)
from simple_syrup.runtime.regional_lora.standard_unet_variant_topology import (
StandardUnetRegionalVariant,
@@ -73,6 +77,52 @@ def test_materialization_applies_global_then_ordered_regional_strengths() -> Non
assert torch.equal(root.diffusion_model.layer.weight, torch.full((2, 2), 2.0))
def test_materialization_calculates_on_comfy_load_device(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Route the exact base calculation to Comfy's requested residency device."""
observed: list[torch.device] = []
def record_cast(
tensor: torch.Tensor,
device: torch.device,
dtype: torch.dtype,
copy: bool = False,
) -> torch.Tensor:
"""Record placement while keeping this deterministic test CPU-only."""
del copy
observed.append(device)
return tensor.to(dtype=dtype, copy=True)
monkeypatch.setattr(comfy.model_management, "cast_to_device", record_cast)
root = _Root()
patcher = ModelPatcher(root, torch.device("cuda"), torch.device("cpu"))
variant = StandardUnetRegionalVariant(0, (_adapter(0, torch.ones((2, 2))),))
StandardUnetVariantMaterializer().materialize(patcher, variant, (1.0,))
assert observed
assert observed[0] == torch.device("cuda")
def test_materialized_variant_rejects_mixed_parameter_devices() -> None:
"""Require one residency destination for every parameter in a variant bank."""
with pytest.raises(ValueError, match="must share one device"):
StandardUnetMaterializedVariant(
0,
(
StandardUnetVariantParameter("a", torch.ones(1)),
StandardUnetVariantParameter(
"b",
torch.empty(1, device="meta"),
),
),
)
def _adapter(index: int, delta: torch.Tensor) -> StandardUnetVariantAdapter:
"""Return one generic legacy-diff operation accepted by Comfy."""
@@ -0,0 +1,67 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify capability-routed standard-UNet materialization placement."""
from __future__ import annotations
import pytest
import torch
from comfy.model_patcher import ModelPatcher
from torch import nn
from simple_syrup.runtime.regional_lora import (
standard_unet_variant_materialization_device as materialization_device,
)
def test_materialization_device_uses_comfy_model_load_device() -> None:
"""Select the requested accelerator independently of source parameter device."""
model = ModelPatcher(
nn.Linear(2, 2, bias=False),
torch.device("cuda"),
torch.device("cpu"),
)
assert materialization_device.StandardUnetVariantMaterializationDevice.resolve(
model
) == torch.device("cuda")
assert next(model.model.parameters()).device == torch.device("cpu")
def test_materialization_device_preserves_cpu_execution() -> None:
"""Keep CPU-only Comfy models on their existing calculation device."""
model = ModelPatcher(
nn.Linear(2, 2, bias=False),
torch.device("cpu"),
torch.device("cpu"),
)
assert materialization_device.StandardUnetVariantMaterializationDevice.resolve(
model
) == torch.device("cpu")
def test_materialization_device_rejects_invalid_model_contract() -> None:
"""Fail closed for absent, malformed, and non-materializable devices."""
with pytest.raises(TypeError, match="requires a MODEL"):
materialization_device.StandardUnetVariantMaterializationDevice.resolve(
object()
)
model = ModelPatcher(
nn.Linear(2, 2, bias=False),
torch.device("cpu"),
torch.device("cpu"),
)
model.load_device = "cuda"
with pytest.raises(TypeError, match="must be a torch device"):
materialization_device.StandardUnetVariantMaterializationDevice.resolve(model)
model.load_device = torch.device("meta")
with pytest.raises(ValueError, match="cannot materialize on meta"):
materialization_device.StandardUnetVariantMaterializationDevice.resolve(model)
+79
View File
@@ -8,6 +8,7 @@ from __future__ import annotations
from typing import cast
import pytest
import torch
from torch import nn
@@ -79,3 +80,81 @@ def test_shell_replaces_only_target_parameter_and_preserves_source() -> None:
assert shell.branch.target.weight is not source.branch.target.weight
assert torch.equal(source.branch.target.weight, torch.ones((2, 2)))
assert torch.equal(shell(torch.ones((1, 2))), torch.full((1, 2), 6.0))
def test_shell_reports_every_missing_target_path() -> None:
"""Preserve complete canonical evidence when target branches are absent."""
source = _Diffusion()
materialized = StandardUnetMaterializedVariant(
0,
(
StandardUnetVariantParameter(
"absent.weight",
torch.ones_like(source.branch.target.weight),
),
StandardUnetVariantParameter(
"branch.absent.weight",
torch.ones_like(source.branch.target.weight),
),
),
)
with pytest.raises(
ValueError,
match=(
"Variant shell did not bind target paths "
"\\('absent.weight', 'branch.absent.weight'\\)"
),
):
StandardUnetVariantShellBuilder().build(source, materialized)
def test_shell_rejects_null_and_incompatible_parameters() -> None:
"""Preserve fail-closed parameter and tensor compatibility contracts."""
source = _Diffusion()
null_parameter = StandardUnetMaterializedVariant(
0,
(
StandardUnetVariantParameter(
"branch.target.bias",
torch.ones(2),
),
),
)
with pytest.raises(TypeError, match="must be a Parameter"):
StandardUnetVariantShellBuilder().build(source, null_parameter)
wrong_shape = StandardUnetMaterializedVariant(
0,
(
StandardUnetVariantParameter(
"branch.target.weight",
torch.ones(3, 2),
),
),
)
with pytest.raises(ValueError, match="is incompatible"):
StandardUnetVariantShellBuilder().build(source, wrong_shape)
def test_shell_binds_pre_resident_target_device() -> None:
"""Accept exact variant tensors already placed for Comfy model residency."""
source = _Diffusion()
materialized = StandardUnetMaterializedVariant(
0,
(
StandardUnetVariantParameter(
"branch.target.weight",
torch.empty_like(source.branch.target.weight, device="meta"),
),
),
)
shell = StandardUnetVariantShellBuilder().build(source, materialized)
assert isinstance(shell, _Diffusion)
assert shell.branch.target.weight.device == torch.device("meta")
assert source.branch.target.weight.device == torch.device("cpu")
+59
View File
@@ -0,0 +1,59 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify the focused benchmark synchronized phase owner."""
from __future__ import annotations
import logging
from types import SimpleNamespace
from typing import Any
import torch
from tools.attention_coupling_benchmark.comfy_probe.synchronized_phase_timing import (
measure_synchronized_phase,
model_device,
)
_LOGGER_NAME = "simple_syrup.runtime.regional_lora.standard_unet_cold_path"
def test_phase_timing_is_transparent_when_capture_is_disabled(caplog: Any) -> None:
"""Avoid logging or synchronization work outside a diagnostic capture."""
value = object()
with measure_synchronized_phase("disabled", device=torch.device("cpu")):
actual = value
assert actual is value
assert not caplog.records
def test_phase_timing_emits_one_structured_record(caplog: Any) -> None:
"""Publish one non-negative elapsed duration under the focused logger."""
with (
caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME),
measure_synchronized_phase("measured", device=torch.device("cpu")),
):
pass
payloads = [
record.cold_path_diagnostics
for record in caplog.records
if hasattr(record, "cold_path_diagnostics")
]
assert len(payloads) == 1
assert payloads[0]["stage"] == "measured"
assert payloads[0]["elapsed_ms"] >= 0.0
def test_model_device_accepts_only_explicit_torch_devices() -> None:
"""Never infer synchronization ownership from arbitrary device-like values."""
device = torch.device("cpu")
assert model_device(SimpleNamespace(load_device=device)) is device
assert model_device(SimpleNamespace(load_device="cuda")) is None
assert model_device(object()) is None
@@ -10,10 +10,16 @@ from importlib import import_module
from typing import TYPE_CHECKING, Any
from .anima_regional_profile_node import StaticAnimaRegionalProfileV3
from .attention_coupling_phase_node import ProfiledKSamplerAttentionCouplingV3
from .clip_schedule_snapshot import SnapshotClipScheduleV3
from .cold_path_capture import (
CaptureColdPathDiagnosticsV3,
ReadColdPathDiagnosticsV3,
)
from .conditioning_batch_snapshot import SnapshotConditioningBatchV3
from .latent_completion import CompleteLatentV3
from .lora_execution_probe import InstrumentLoraModelV3, ReadLoraMetricsV3
from .materialization_parity_node import CompareMaterializationParityV3
from .model_modifier_snapshot import SnapshotModelModifierV3
from .operator_profile import ProfileIndexedModelCallV3, ReadOperatorProfileV3
from .prompt_control_expansion import SnapshotPromptControlExpansionV3
@@ -50,15 +56,19 @@ class BenchmarkProbeExtension(_ComfyExtensionBase):
return [
StaticAnimaRegionalProfileV3,
ProfiledKSamplerAttentionCouplingV3,
SnapshotClipScheduleV3,
SnapshotConditioningBatchV3,
InstrumentModelV3,
ReadMetricsV3,
CompleteLatentV3,
CaptureColdPathDiagnosticsV3,
ReadColdPathDiagnosticsV3,
ProfileIndexedModelCallV3,
ReadOperatorProfileV3,
InstrumentLoraModelV3,
ReadLoraMetricsV3,
CompareMaterializationParityV3,
SnapshotModelModifierV3,
SnapshotPromptControlV3,
SnapshotPromptControlExpansionV3,
@@ -0,0 +1,90 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Expose phase-profiled Attention Coupling through a benchmark-only node."""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar
from simple_syrup.nodes_v3.ksampler_attention_coupling import (
KSamplerAttentionCouplingV3,
)
from simple_syrup.nodes_v3.ksampler_schema import (
attention_coupling_ksampler_inputs,
)
from .attention_coupling_phase_profile import (
ProfiledAttentionCouplingSamplingService,
)
from .conditioning_batch_bridge import normalize_conditioning_batch
if TYPE_CHECKING:
class _ComfyNodeBase(KSamplerAttentionCouplingV3):
"""Type-checking base for the benchmark-only sampler."""
pass
else:
_ComfyNodeBase = KSamplerAttentionCouplingV3
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
class ProfiledKSamplerAttentionCouplingV3(_ComfyNodeBase):
"""Run exact Attention Coupling while publishing two outer phase timings."""
sampling_service_class: ClassVar[type[ProfiledAttentionCouplingSamplingService]] = (
ProfiledAttentionCouplingSamplingService
)
@classmethod
def define_schema(cls) -> Any:
"""Declare the production inputs under one dev-only benchmark identity."""
return _comfy_io.Schema(
node_id="SimpleSyrupBenchmark.ProfiledKSamplerAttentionCoupling",
display_name="Benchmark Profiled KSampler Attention Coupling",
category="SimpleSyrup/Benchmark",
inputs=attention_coupling_ksampler_inputs(_comfy_io),
outputs=[_comfy_io.Latent.Output("latent")],
is_dev_only=True,
)
@classmethod
def execute(
cls,
model: Any,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
positive: object,
negative: object,
region_masks: object,
regional_prompt_weight: float,
region_mask_feather: int,
latent_image: dict[str, Any],
denoise: float,
) -> tuple[dict[str, Any]]:
"""Bridge canonical host batches, then run the inherited exact delegate."""
return super().execute(
model=model,
seed=seed,
steps=steps,
cfg=cfg,
sampler_name=sampler_name,
scheduler=scheduler,
positive=normalize_conditioning_batch(positive),
negative=normalize_conditioning_batch(negative),
region_masks=region_masks,
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=region_mask_feather,
latent_image=latent_image,
denoise=denoise,
)
@@ -0,0 +1,81 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Profile exact preparation and sampler-delegate boundaries without altering them."""
from __future__ import annotations
from typing import Any, ClassVar
from simple_syrup.domain.regional_attention_execution import (
RegionalAttentionExecutionMode,
)
from simple_syrup.services.attention_coupling_model_preparation_service import (
AttentionCouplingModelPreparationService,
)
from simple_syrup.services.attention_coupling_sampling_service import (
AttentionCouplingSamplingService,
)
from simple_syrup.services.ksampler_sampling_service import KSamplerSamplingService
from .preparation_collaborator_profile import (
ProfiledAttentionCouplingModelPreparationService,
)
from .synchronized_phase_timing import measure_synchronized_phase, model_device
class ProfiledAttentionCouplingSamplingService(AttentionCouplingSamplingService):
"""Delegate the production two-call sequence with synchronized phase timing."""
model_preparation_service_class: ClassVar[
type[AttentionCouplingModelPreparationService]
] = ProfiledAttentionCouplingModelPreparationService
sampling_service_class: ClassVar[type[KSamplerSamplingService]] = (
KSamplerSamplingService
)
def sample(
self,
*,
model: Any,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
positive: object,
negative: object,
region_masks: object,
regional_prompt_weight: float,
region_mask_feather: int,
latent_image: dict[str, Any],
denoise: float,
) -> dict[str, Any]:
"""Return the exact production result after timing its two owners."""
device = model_device(model)
with measure_synchronized_phase("model_preparation_total", device=device):
prepared = self.model_preparation_service_class().prepare(
model=model,
positive=positive,
negative=negative,
region_masks=region_masks,
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=region_mask_feather,
latent_image=latent_image,
execution_mode=RegionalAttentionExecutionMode.FULL,
)
with measure_synchronized_phase("ksampler_delegate_total", device=device):
return self.sampling_service_class().sample(
model=prepared.model,
seed=seed,
steps=steps,
cfg=cfg,
sampler_name=sampler_name,
scheduler=scheduler,
positive=prepared.positive,
negative=prepared.negative,
latent_image=latent_image,
denoise=denoise,
)
@@ -0,0 +1,199 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Capture standard-UNet cold-stage diagnostics for one benchmark request."""
from __future__ import annotations
import json
import logging
import threading
import time
from dataclasses import dataclass
from importlib import import_module
from typing import TYPE_CHECKING, Any, cast
import torch
_comfy_api: Any = None
if TYPE_CHECKING:
class _ComfyNodeBase:
"""Type-checking base for benchmark-only Comfy v3 nodes."""
pass
else:
_comfy_api = import_module("comfy_api.latest")
_ComfyNodeBase = _comfy_api.io.ComfyNode
_comfy_io: Any = None if TYPE_CHECKING else _comfy_api.io
_LOGGER_NAME = "simple_syrup.runtime.regional_lora.standard_unet_cold_path"
_COMPOSITION_LOGGER_NAME = (
"simple_syrup.runtime.regional_lora.standard_unet_composition"
)
class _ColdPathHandler(logging.Handler):
"""Retain every ordered JSON-safe cold-stage diagnostic."""
def __init__(self) -> None:
"""Initialize empty thread-safe record storage."""
super().__init__(logging.DEBUG)
self._lock = threading.Lock()
self._records: list[str] = []
self._model_call_count = 0
def emit(self, record: logging.LogRecord) -> None:
"""Capture only the focused structured payload."""
if record.name == _COMPOSITION_LOGGER_NAME:
if getattr(record, "regional_composition", None) is not None:
with self._lock:
self._model_call_count += 1
return
if record.name != _LOGGER_NAME:
return
payload = getattr(record, "cold_path_diagnostics", None)
if payload is None:
return
serialized = json.dumps(payload, sort_keys=True, separators=(",", ":"))
if not isinstance(json.loads(serialized), dict):
raise TypeError("Cold-path diagnostic payload must be an object.")
with self._lock:
self._records.append(serialized)
def records(self) -> list[dict[str, object]]:
"""Return detached ordered JSON objects."""
with self._lock:
return [
cast(dict[str, object], json.loads(serialized))
for serialized in self._records
]
@property
def model_call_count(self) -> int:
"""Return the exact observed persistent-composition call count."""
with self._lock:
return self._model_call_count
@dataclass(frozen=True, slots=True)
class _CaptureState:
"""Retain the sole active logger lease and handler."""
loggers: tuple[tuple[logging.Logger, int], ...]
handler: _ColdPathHandler
_STATES: dict[str, _CaptureState] = {}
_STATE_LOCK = threading.Lock()
class CaptureColdPathDiagnosticsV3(_ComfyNodeBase):
"""Start one exclusive cold-stage capture before sampler preparation."""
@classmethod
def define_schema(cls) -> Any:
"""Declare a benchmark-only transparent MODEL boundary."""
return _comfy_io.Schema(
node_id="SimpleSyrupBenchmark.CaptureColdPathDiagnostics",
display_name="Benchmark Capture Cold Path Diagnostics",
category="SimpleSyrup/Benchmark",
inputs=[
_comfy_io.Model.Input("model"),
_comfy_io.String.Input("run_id"),
],
outputs=[_comfy_io.Model.Output("model")],
is_dev_only=True,
)
@classmethod
def execute(cls, model: Any, run_id: str) -> Any:
"""Lease DEBUG capture without cloning or mutating the supplied MODEL."""
if not isinstance(run_id, str) or not run_id:
raise ValueError("Cold-path capture run id must be non-empty.")
loggers = tuple(
logging.getLogger(name) for name in (_LOGGER_NAME, _COMPOSITION_LOGGER_NAME)
)
handler = _ColdPathHandler()
with _STATE_LOCK:
if _STATES:
raise RuntimeError("Cold-path diagnostic capture is already active.")
state = _CaptureState(
tuple((logger, logger.level) for logger in loggers),
handler,
)
_STATES[run_id] = state
for logger in loggers:
logger.setLevel(logging.DEBUG)
logger.addHandler(handler)
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
return _comfy_io.NodeOutput(model)
class ReadColdPathDiagnosticsV3(_ComfyNodeBase):
"""Synchronize and publish one complete cold-stage capture."""
@classmethod
def define_schema(cls) -> Any:
"""Declare the benchmark-only latent and JSON terminal."""
return _comfy_io.Schema(
node_id="SimpleSyrupBenchmark.ReadColdPathDiagnostics",
display_name="Benchmark Read Cold Path Diagnostics",
category="SimpleSyrup/Benchmark",
inputs=[
_comfy_io.Latent.Input("latent"),
_comfy_io.String.Input("run_id"),
],
outputs=[
_comfy_io.Latent.Output("latent"),
_comfy_io.String.Output("diagnostics_json"),
],
is_output_node=True,
is_dev_only=True,
)
@classmethod
def execute(cls, latent: dict[str, Any], run_id: str) -> Any:
"""End the logger lease after all queued CUDA work completes."""
if torch.cuda.is_available():
torch.cuda.synchronize()
completed_at_ns = time.perf_counter_ns()
with _STATE_LOCK:
state = _STATES.pop(run_id, None)
if state is None:
raise ValueError(
f"Cold-path diagnostic capture was not started: {run_id!r}."
)
for logger, original_level in state.loggers:
logger.removeHandler(state.handler)
logger.setLevel(original_level)
records = state.handler.records()
if not records:
raise ValueError("Cold-path diagnostic capture observed no stages.")
result: dict[str, object] = {
"run_id": run_id,
"completed_at_ns": completed_at_ns,
"record_count": len(records),
"records": records,
"model_call_count": state.handler.model_call_count,
"peak_vram_bytes": (
torch.cuda.max_memory_allocated() if torch.cuda.is_available() else 0
),
}
encoded = json.dumps(result, sort_keys=True, separators=(",", ":"))
return _comfy_io.NodeOutput(
latent,
encoded,
ui={"cold_path_diagnostics": [result]},
)
@@ -0,0 +1,44 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Normalize canonical ConditioningBatch values across Comfy module namespaces."""
from __future__ import annotations
from typing import Protocol, runtime_checkable
from simple_syrup.domain.conditioning_batch import ConditioningBatch
@runtime_checkable
class _ConditioningBatchValue(Protocol):
"""Describe the immutable domain surface across Comfy loader namespaces."""
@property
def entries(self) -> tuple[object, ...]:
"""Return ordered conditioning values."""
...
def normalize_conditioning_batch(value: object) -> object:
"""Return one local batch for canonical host values or pass conditioning through."""
if isinstance(value, ConditioningBatch):
return value
if not hasattr(value, "entries"):
return value
value_type = type(value)
if value_type.__name__ != "ConditioningBatch" or not value_type.__module__.endswith(
".domain.conditioning_batch"
):
raise TypeError("Unsupported conditioning batch runtime type.")
if not isinstance(value, _ConditioningBatchValue):
raise TypeError("ConditioningBatch must expose immutable entries.")
entries = value.entries
if not isinstance(entries, tuple):
raise TypeError("ConditioningBatch entries must be an immutable tuple.")
if not entries:
raise ValueError("ConditioningBatch entries must not be empty.")
return ConditioningBatch(entries)
@@ -9,10 +9,13 @@ from __future__ import annotations
import json
from collections.abc import Sequence
from importlib import import_module
from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable
from typing import TYPE_CHECKING, Any
import torch
from simple_syrup.domain.conditioning_batch import ConditioningBatch
from .conditioning_batch_bridge import normalize_conditioning_batch
from .tensor_snapshot import snapshot_tensor
_comfy_api: Any = None
@@ -33,17 +36,6 @@ _mixed_conditioning_io: Any = (
)
@runtime_checkable
class _ConditioningBatchValue(Protocol):
"""Describe the immutable domain surface across Comfy loader namespaces."""
@property
def entries(self) -> tuple[object, ...]:
"""Return ordered conditioning values."""
...
class SnapshotConditioningBatchV3(_ComfyNodeBase):
"""Observe paired conditioning tensors without changing graph values."""
@@ -86,13 +78,13 @@ class SnapshotConditioningBatchV3(_ComfyNodeBase):
def snapshot_conditioning_value(value: object) -> dict[str, object]:
"""Record one conditioning or ordered ConditioningBatch tensor structure."""
batch_entries = _conditioning_batch_entries(value)
if batch_entries is not None:
normalized = normalize_conditioning_batch(value)
if isinstance(normalized, ConditioningBatch):
kind = "conditioning_batch"
batches = batch_entries
batches = normalized.entries
else:
kind = "conditioning"
batches = (value,)
batches = (normalized,)
return {
"kind": kind,
"batch_entries": [
@@ -108,26 +100,6 @@ def snapshot_conditioning_value(value: object) -> dict[str, object]:
}
def _conditioning_batch_entries(value: object) -> tuple[object, ...] | None:
"""Narrow the canonical batch contract across Comfy package namespaces."""
if not hasattr(value, "entries"):
return None
value_type = type(value)
if value_type.__name__ != "ConditioningBatch" or not value_type.__module__.endswith(
".domain.conditioning_batch"
):
raise TypeError("Unsupported conditioning batch runtime type.")
if not isinstance(value, _ConditioningBatchValue):
raise TypeError("ConditioningBatch must expose immutable entries.")
entries = value.entries
if not isinstance(entries, tuple):
raise TypeError("ConditioningBatch entries must be an immutable tuple.")
if not entries:
raise ValueError("ConditioningBatch entries must not be empty.")
return entries
def _snapshot_conditioning(
conditioning: object,
*,
@@ -0,0 +1,56 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Profile exact MODEL interoperability validation boundaries."""
from __future__ import annotations
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.regional_attention_execution import (
RegionalAttentionExecutionMode,
)
from simple_syrup.domain.regional_model_capabilities import RegionalModelCapabilities
from simple_syrup.runtime.regional_model_patch_interop import (
RegionalModelPatchInteropReport,
RegionalModelPatchInteropValidator,
)
from .synchronized_phase_timing import measure_synchronized_phase, model_device
class ProfiledRegionalModelPatchInteropValidator(RegionalModelPatchInteropValidator):
"""Time exact production interop admission and execution validation."""
def validate(
self,
model: object,
capabilities: RegionalModelCapabilities,
) -> RegionalModelPatchInteropReport:
"""Delegate admission validation with synchronized model-device timing."""
with measure_synchronized_phase(
"interop_validation",
device=model_device(model),
):
return super().validate(model, capabilities)
def validate_execution(
self,
report: RegionalModelPatchInteropReport,
processed_plan: ProcessedRegionalAttentionPlan,
execution_mode: RegionalAttentionExecutionMode,
) -> None:
"""Delegate execution validation while recording its exact duration."""
with measure_synchronized_phase(
"interop_execution_validation",
device=None,
):
return super().validate_execution(
report,
processed_plan,
execution_mode,
)
@@ -0,0 +1,100 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Expose image-free materialization parity as one benchmark-only node."""
from __future__ import annotations
import json
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar
from .materialization_parity_probe import (
MATERIALIZATION_PARITY_PROBE,
MaterializationParityProbe,
)
_comfy_api: Any = None
if TYPE_CHECKING:
class _ComfyNodeBase:
"""Type-checking base for the benchmark-only node."""
pass
else:
_comfy_api = import_module("comfy_api.latest")
_ComfyNodeBase = _comfy_api.io.ComfyNode
_comfy_io: Any = None if TYPE_CHECKING else _comfy_api.io
class CompareMaterializationParityV3(_ComfyNodeBase):
"""Publish CPU-versus-selected-device regional parameter-bank evidence."""
probe: ClassVar[MaterializationParityProbe] = MATERIALIZATION_PARITY_PROBE
@classmethod
def define_schema(cls) -> Any:
"""Declare one terminal benchmark comparison without sampler outputs."""
conditioning_batch = _comfy_io.Custom("CONDITIONING_BATCH")
return _comfy_io.Schema(
node_id="SimpleSyrupBenchmark.CompareMaterializationParity",
display_name="Benchmark Compare Materialization Parity",
category="SimpleSyrup/Benchmark",
inputs=[
_comfy_io.Model.Input("model"),
_comfy_io.MultiType.Input(
"positive",
[_comfy_io.Conditioning, conditioning_batch],
),
_comfy_io.MultiType.Input(
"negative",
[_comfy_io.Conditioning, conditioning_batch],
),
_comfy_io.Mask.Input("region_masks"),
_comfy_io.Latent.Input("latent_image"),
_comfy_io.Int.Input(
"region_mask_feather",
default=0,
min=0,
max=4096,
),
_comfy_io.String.Input("run_id"),
],
outputs=[_comfy_io.String.Output("comparison_json")],
is_output_node=True,
is_dev_only=True,
)
@classmethod
def execute(
cls,
model: object,
positive: object,
negative: object,
region_masks: object,
latent_image: object,
region_mask_feather: int,
run_id: str,
) -> Any:
"""Run one synchronized comparison and publish JSON-safe evidence."""
if not isinstance(run_id, str) or not run_id:
raise ValueError("Materialization parity run id must be non-empty.")
payload = cls.probe.compare(
model=model,
positive=positive,
negative=negative,
region_masks=region_masks,
latent_image=latent_image,
region_mask_feather=region_mask_feather,
).as_json_object()
payload["run_id"] = run_id
encoded = json.dumps(payload, sort_keys=True, separators=(",", ":"))
return _comfy_io.NodeOutput(
encoded,
ui={"materialization_parity": [payload]},
)
@@ -0,0 +1,295 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Compare CPU and selected-device regional materialization without sampling."""
from __future__ import annotations
import gc
import math
import time
from dataclasses import asdict, dataclass
import torch
from comfy.model_patcher import ModelPatcher
from simple_syrup.domain.raw_regional_attention import (
build_raw_regional_attention_plan,
)
from simple_syrup.masking.regional_prompt_masks import build_regional_mask_bank
from simple_syrup.runtime.comfy_conditioning_model_loader import (
COMFY_CONDITIONING_MODEL_LOADER,
)
from simple_syrup.runtime.regional_lora.standard_unet_native_admission import (
STANDARD_UNET_NATIVE_LORA_ADMISSION_SERVICE,
)
from simple_syrup.runtime.regional_lora.standard_unet_variant_materialization import (
STANDARD_UNET_VARIANT_MATERIALIZER,
)
from simple_syrup.runtime.regional_lora.standard_unet_variant_topology import (
STANDARD_UNET_VARIANT_TOPOLOGY_BUILDER,
StandardUnetRegionalVariant,
)
from simple_syrup.runtime.regional_lora_conditioning_adapter import (
REGIONAL_LORA_CONDITIONING_ADAPTER,
)
from .conditioning_batch_bridge import normalize_conditioning_batch
from .materialized_variant_comparison import (
MaterializedVariantComparison,
compare_materialized_variants,
)
@dataclass(frozen=True, slots=True)
class MaterializationParityVariantObservation:
"""Retain one region's materialization timing and numerical comparison."""
region_index: int
selected_device_materialization_ms: float
cpu_materialization_ms: float
comparison_ms: float
comparison: MaterializedVariantComparison
@dataclass(frozen=True, slots=True)
class MaterializationParityObservation:
"""Retain one complete architecture-neutral materialization comparison."""
selected_device: str
source_load_ms: float
admission_ms: float
peak_vram_bytes: int
variants: tuple[MaterializationParityVariantObservation, ...]
@property
def exact(self) -> bool:
"""Report whether every compared variant bank is byte-identical."""
return all(variant.comparison.exact for variant in self.variants)
def as_json_object(self) -> dict[str, object]:
"""Return one stable JSON-safe diagnostic object."""
variant_payloads = []
for variant in self.variants:
payload = asdict(variant)
comparison = dict(payload["comparison"])
comparison["exact"] = variant.comparison.exact
payload["comparison"] = comparison
variant_payloads.append(payload)
element_count = sum(
variant.comparison.element_count for variant in self.variants
)
differing_count = sum(
variant.comparison.differing_element_count for variant in self.variants
)
absolute_error_sum = sum(
variant.comparison.mean_absolute_error * variant.comparison.element_count
for variant in self.variants
)
squared_error_sum = sum(
variant.comparison.root_mean_squared_error**2
* variant.comparison.element_count
for variant in self.variants
)
return {
"selected_device": self.selected_device,
"source_load_ms": self.source_load_ms,
"admission_ms": self.admission_ms,
"peak_vram_bytes": self.peak_vram_bytes,
"variant_count": len(self.variants),
"element_count": element_count,
"differing_element_count": differing_count,
"max_absolute_error": max(
(variant.comparison.max_absolute_error for variant in self.variants),
default=0.0,
),
"mean_absolute_error": (
absolute_error_sum / element_count if element_count else 0.0
),
"root_mean_squared_error": (
math.sqrt(squared_error_sum / element_count) if element_count else 0.0
),
"exact": self.exact,
"variants": variant_payloads,
}
class MaterializationParityProbe:
"""Resolve once and compare each complete regional bank on CPU and device."""
def compare(
self,
*,
model: object,
positive: object,
negative: object,
region_masks: object,
latent_image: object,
region_mask_feather: int,
) -> MaterializationParityObservation:
"""Return bounded parity evidence without sampling or mutating the source."""
patcher = self._require_patcher(model)
samples = self._latent_samples(latent_image)
selected_device = patcher.load_device
if selected_device.type == "cpu":
raise ValueError(
"Materialization parity requires a non-CPU selected load device."
)
self._synchronize(selected_device)
torch.cuda.reset_peak_memory_stats(selected_device)
mask_bank = build_regional_mask_bank(
region_masks,
feather=region_mask_feather,
canvas_height=int(samples.shape[-2]),
canvas_width=int(samples.shape[-1]),
)
raw = build_raw_regional_attention_plan(
positive=normalize_conditioning_batch(positive),
negative=normalize_conditioning_batch(negative),
mask_bank=mask_bank,
)
started = time.perf_counter_ns()
COMFY_CONDITIONING_MODEL_LOADER.load(patcher)
self._synchronize(selected_device)
source_load_ms = self._elapsed_ms(started)
started = time.perf_counter_ns()
adaptation = REGIONAL_LORA_CONDITIONING_ADAPTER.adapt(
raw,
model=patcher.model,
)
admission = STANDARD_UNET_NATIVE_LORA_ADMISSION_SERVICE.admit(
patcher,
adaptation,
)
resolution = admission.resolution
if resolution is None:
raise ValueError("Materialization parity requires regional LoRA targets.")
topology = STANDARD_UNET_VARIANT_TOPOLOGY_BUILDER.build(resolution)
multipliers = self._time_invariant_multipliers(adaptation.plan.adapters)
admission_ms = self._elapsed_ms(started)
observations = tuple(
self._compare_variant(
patcher,
variant,
multipliers,
selected_device=selected_device,
)
for variant in topology.variants
)
return MaterializationParityObservation(
selected_device=str(selected_device),
source_load_ms=source_load_ms,
admission_ms=admission_ms,
peak_vram_bytes=torch.cuda.max_memory_allocated(selected_device),
variants=observations,
)
@staticmethod
def _compare_variant(
patcher: ModelPatcher,
variant: StandardUnetRegionalVariant,
multipliers: tuple[float, ...],
*,
selected_device: torch.device,
) -> MaterializationParityVariantObservation:
"""Materialize and release one aligned bank pair before advancing."""
selected_patcher = patcher.clone()
cpu_patcher = patcher.clone()
cpu_patcher.load_device = torch.device("cpu")
started = time.perf_counter_ns()
selected_bank = STANDARD_UNET_VARIANT_MATERIALIZER.materialize(
selected_patcher,
variant,
multipliers,
)
MaterializationParityProbe._synchronize(selected_device)
selected_ms = MaterializationParityProbe._elapsed_ms(started)
started = time.perf_counter_ns()
cpu_bank = STANDARD_UNET_VARIANT_MATERIALIZER.materialize(
cpu_patcher,
variant,
multipliers,
)
cpu_ms = MaterializationParityProbe._elapsed_ms(started)
started = time.perf_counter_ns()
comparison = compare_materialized_variants(cpu_bank, selected_bank)
comparison_ms = MaterializationParityProbe._elapsed_ms(started)
del selected_bank, cpu_bank, selected_patcher, cpu_patcher
gc.collect()
torch.cuda.empty_cache()
return MaterializationParityVariantObservation(
region_index=comparison.region_index,
selected_device_materialization_ms=selected_ms,
cpu_materialization_ms=cpu_ms,
comparison_ms=comparison_ms,
comparison=comparison,
)
@staticmethod
def _time_invariant_multipliers(adapters: tuple[object, ...]) -> tuple[float, ...]:
"""Return one exact effective value for every canonical adapter use."""
values: list[float] = []
for adapter in adapters:
schedule = getattr(adapter, "schedule", None)
if not isinstance(schedule, tuple) or not schedule:
raise ValueError("Materialization parity requires adapter schedules.")
multipliers = tuple(
getattr(boundary, "strength_multiplier", None) for boundary in schedule
)
first = multipliers[0]
if not isinstance(first, float) or any(
multiplier != first for multiplier in multipliers
):
raise ValueError(
"Materialization parity requires time-invariant schedules."
)
values.append(first)
if not values:
raise ValueError("Materialization parity requires regional adapters.")
return tuple(values)
@staticmethod
def _require_patcher(model: object) -> ModelPatcher:
"""Narrow one Comfy MODEL and its selected CUDA-capable device."""
if not isinstance(model, ModelPatcher):
raise TypeError("Materialization parity requires a Comfy MODEL.")
if not isinstance(model.load_device, torch.device):
raise TypeError("Materialization parity load device must be a device.")
if model.load_device.type != "cuda" or not torch.cuda.is_available():
raise ValueError("Materialization parity requires available CUDA.")
return model
@staticmethod
def _latent_samples(latent_image: object) -> torch.Tensor:
"""Return one valid floating latent tensor for mask geometry."""
if not isinstance(latent_image, dict):
raise TypeError("Materialization parity latent must be a dictionary.")
samples = latent_image.get("samples")
if not isinstance(samples, torch.Tensor) or not samples.is_floating_point():
raise TypeError("Materialization parity latent must contain samples.")
if samples.ndim not in (4, 5):
raise ValueError("Materialization parity latent samples must be 4D or 5D.")
return samples
@staticmethod
def _synchronize(device: torch.device) -> None:
"""Complete selected-device work before recording a boundary."""
torch.cuda.synchronize(device)
@staticmethod
def _elapsed_ms(started_at_ns: int) -> float:
"""Return monotonic milliseconds since one captured boundary."""
return (time.perf_counter_ns() - started_at_ns) / 1_000_000.0
MATERIALIZATION_PARITY_PROBE = MaterializationParityProbe()
@@ -0,0 +1,214 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Compare exact regional variant banks with bounded temporary memory."""
from __future__ import annotations
import hashlib
import math
import struct
from dataclasses import dataclass
from typing import Protocol
import torch
from simple_syrup.runtime.regional_lora.standard_unet_variant_materialization import (
StandardUnetMaterializedVariant,
)
_DEFAULT_CHUNK_ELEMENTS = 1_048_576
class _HashSink(Protocol):
"""Accept ordered bytes for one deterministic digest."""
def update(self, value: bytes) -> None:
"""Add one byte segment to the digest."""
@dataclass(frozen=True, slots=True)
class MaterializedVariantComparison:
"""Summarize exact structure, bytes, and bounded numerical divergence."""
region_index: int
parameter_count: int
element_count: int
differing_element_count: int
max_absolute_error: float
mean_absolute_error: float
root_mean_squared_error: float
max_error_parameter_path: str | None
reference_sha256: str
candidate_sha256: str
@property
def exact(self) -> bool:
"""Report byte-identical structure and tensor content."""
return self.reference_sha256 == self.candidate_sha256
def compare_materialized_variants(
reference: StandardUnetMaterializedVariant,
candidate: StandardUnetMaterializedVariant,
*,
chunk_elements: int = _DEFAULT_CHUNK_ELEMENTS,
) -> MaterializedVariantComparison:
"""Compare aligned parameters incrementally after bounded CPU transfer."""
_validate_inputs(reference, candidate, chunk_elements=chunk_elements)
reference_digest = hashlib.sha256()
candidate_digest = hashlib.sha256()
_update_bank_metadata(reference_digest, reference)
_update_bank_metadata(candidate_digest, candidate)
element_count = 0
differing_element_count = 0
maximum_error = 0.0
absolute_error_sum = 0.0
squared_error_sum = 0.0
maximum_error_path: str | None = None
with torch.no_grad():
for reference_parameter, candidate_parameter in zip(
reference.parameters,
candidate.parameters,
strict=True,
):
_validate_parameter_pair(reference_parameter, candidate_parameter)
reference_cpu = reference_parameter.tensor.detach().to("cpu")
candidate_cpu = candidate_parameter.tensor.detach().to("cpu")
_update_parameter_metadata(
reference_digest,
reference_parameter.path,
reference_cpu,
)
_update_parameter_metadata(
candidate_digest,
candidate_parameter.path,
candidate_cpu,
)
reference_flat = reference_cpu.contiguous().reshape(-1)
candidate_flat = candidate_cpu.contiguous().reshape(-1)
for start in range(0, reference_flat.numel(), chunk_elements):
stop = min(start + chunk_elements, reference_flat.numel())
reference_chunk = reference_flat[start:stop]
candidate_chunk = candidate_flat[start:stop]
_require_finite(reference_chunk, reference_parameter.path)
_require_finite(candidate_chunk, candidate_parameter.path)
reference_digest.update(_raw_tensor_bytes(reference_chunk))
candidate_digest.update(_raw_tensor_bytes(candidate_chunk))
differing_element_count += int(
torch.count_nonzero(reference_chunk != candidate_chunk).item()
)
delta = candidate_chunk.to(torch.float64) - reference_chunk.to(
torch.float64
)
absolute_delta = delta.abs()
local_maximum = float(absolute_delta.max().item())
if local_maximum > maximum_error:
maximum_error = local_maximum
maximum_error_path = reference_parameter.path
absolute_error_sum += float(absolute_delta.sum().item())
squared_error_sum += float(torch.square(delta).sum().item())
element_count += reference_chunk.numel()
if element_count < 1:
raise ValueError("Materialized variant comparison requires tensor elements.")
return MaterializedVariantComparison(
region_index=reference.region_index,
parameter_count=len(reference.parameters),
element_count=element_count,
differing_element_count=differing_element_count,
max_absolute_error=maximum_error,
mean_absolute_error=absolute_error_sum / element_count,
root_mean_squared_error=math.sqrt(squared_error_sum / element_count),
max_error_parameter_path=maximum_error_path,
reference_sha256=reference_digest.hexdigest(),
candidate_sha256=candidate_digest.hexdigest(),
)
def _validate_inputs(
reference: object,
candidate: object,
*,
chunk_elements: int,
) -> None:
"""Require comparable bank values and a bounded positive chunk size."""
if not isinstance(reference, StandardUnetMaterializedVariant) or not isinstance(
candidate, StandardUnetMaterializedVariant
):
raise TypeError("Materialized comparison requires two variant banks.")
if reference.region_index != candidate.region_index:
raise ValueError("Materialized variant region indices do not match.")
if isinstance(chunk_elements, bool) or not isinstance(chunk_elements, int):
raise TypeError("Materialized comparison chunk size must be an integer.")
if chunk_elements < 1:
raise ValueError("Materialized comparison chunk size must be positive.")
reference_paths = tuple(parameter.path for parameter in reference.parameters)
candidate_paths = tuple(parameter.path for parameter in candidate.parameters)
if reference_paths != candidate_paths:
raise ValueError("Materialized variant parameter paths do not match.")
def _validate_parameter_pair(reference: object, candidate: object) -> None:
"""Require one aligned shape and dtype without device restrictions."""
reference_tensor = getattr(reference, "tensor", None)
candidate_tensor = getattr(candidate, "tensor", None)
if not isinstance(reference_tensor, torch.Tensor) or not isinstance(
candidate_tensor, torch.Tensor
):
raise TypeError("Materialized comparison parameters must contain tensors.")
if reference_tensor.shape != candidate_tensor.shape:
raise ValueError("Materialized variant parameter shapes do not match.")
if reference_tensor.dtype != candidate_tensor.dtype:
raise ValueError("Materialized variant parameter dtypes do not match.")
def _update_bank_metadata(
digest: _HashSink,
variant: StandardUnetMaterializedVariant,
) -> None:
"""Frame one bank identity independently of object or device identity."""
digest.update(struct.pack("<qQ", variant.region_index, len(variant.parameters)))
def _update_parameter_metadata(
digest: _HashSink,
path: str,
tensor: torch.Tensor,
) -> None:
"""Frame path, dtype, shape, and element count before raw tensor bytes."""
_update_string(digest, path)
_update_string(digest, str(tensor.dtype))
digest.update(struct.pack("<Q", tensor.ndim))
for dimension in tensor.shape:
digest.update(struct.pack("<q", dimension))
digest.update(struct.pack("<Q", tensor.numel()))
def _update_string(digest: _HashSink, value: str) -> None:
"""Add one length-prefixed UTF-8 string to a bank digest."""
encoded = value.encode("utf-8")
digest.update(struct.pack("<Q", len(encoded)))
digest.update(encoded)
def _raw_tensor_bytes(tensor: torch.Tensor) -> bytes:
"""Return device-neutral contiguous storage bytes for one CPU chunk."""
return tensor.contiguous().view(torch.uint8).numpy().tobytes()
def _require_finite(tensor: torch.Tensor, path: str) -> None:
"""Reject nonfinite weights whose error metrics would be undefined."""
if not bool(torch.isfinite(tensor).all().item()):
raise ValueError(
f"Materialized variant parameter {path!r} must contain finite values."
)
@@ -0,0 +1,139 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Profile exact model-family boundaries behind a transparent decorator."""
from __future__ import annotations
import torch
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.raw_regional_attention import RawRegionalAttentionPlan
from simple_syrup.domain.regional_model_capabilities import RegionalModelCapabilities
from simple_syrup.runtime.attention_coupling.context_validation import (
RegionalContextValidator,
)
from simple_syrup.runtime.attention_coupling.family_admission import (
AttentionCouplingFamilyAdmission,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from simple_syrup.runtime.regional_model_patch_interop import (
RegionalModelPatchInteropReport,
)
from simple_syrup.services.attention_coupling_model_family import (
AttentionCouplingModelFamily,
AttentionCouplingPreparedModelReuse,
AttentionCouplingSamplerConditioning,
)
from simple_syrup.services.attention_coupling_model_family_selector import (
AttentionCouplingModelFamilySelector,
)
from .synchronized_phase_timing import measure_synchronized_phase, model_device
class ProfiledAttentionCouplingModelFamily:
"""Time exact family calls while preserving the selected family instance."""
def __init__(self, delegate: AttentionCouplingModelFamily) -> None:
"""Retain the production-selected family as the sole behavior owner."""
self._delegate = delegate
@property
def delegate(self) -> AttentionCouplingModelFamily:
"""Expose the exact wrapped family for benchmark verification."""
return self._delegate
@property
def context_validator(self) -> RegionalContextValidator:
"""Return the production family's exact context validator."""
return self._delegate.context_validator
@property
def prepared_model_reuse(self) -> AttentionCouplingPreparedModelReuse:
"""Return the production family's exact prepared-result policy."""
return self._delegate.prepared_model_reuse
def validate_latent(self, samples: torch.Tensor) -> None:
"""Delegate latent validation while timing the family boundary."""
with measure_synchronized_phase(
"family_validate_latent",
device=samples.device,
):
return self._delegate.validate_latent(samples)
def admit_adaptation(
self,
model: object,
adaptation: RegionalLoraPlanAdaptation,
) -> AttentionCouplingFamilyAdmission:
"""Delegate admission while timing the exact selected family."""
with measure_synchronized_phase(
"family_admit_adaptation",
device=model_device(model),
):
return self._delegate.admit_adaptation(model, adaptation)
def prepare_sampler_conditioning(
self,
plan: RawRegionalAttentionPlan,
region_strengths: tuple[float, ...],
) -> AttentionCouplingSamplerConditioning:
"""Delegate sampler conditioning while timing the family boundary."""
with measure_synchronized_phase(
"family_sampler_conditioning",
device=None,
):
return self._delegate.prepare_sampler_conditioning(
plan,
region_strengths,
)
def derive(
self,
*,
model: object,
processed_plan: ProcessedRegionalAttentionPlan,
admission: AttentionCouplingFamilyAdmission,
interop_report: RegionalModelPatchInteropReport,
region_strengths: tuple[float, ...],
latent_batch_size: int,
) -> object:
"""Delegate complete derivation with synchronized model-device timing."""
with measure_synchronized_phase(
"family_derive_total",
device=model_device(model),
):
return self._delegate.derive(
model=model,
processed_plan=processed_plan,
admission=admission,
interop_report=interop_report,
region_strengths=region_strengths,
latent_batch_size=latent_batch_size,
)
class ProfiledAttentionCouplingModelFamilySelector(
AttentionCouplingModelFamilySelector
):
"""Decorate the exact family selected by production capability routing."""
def select(
self,
capabilities: RegionalModelCapabilities,
) -> AttentionCouplingModelFamily:
"""Return a transparent profiler around the production-selected family."""
return ProfiledAttentionCouplingModelFamily(super().select(capabilities))
@@ -0,0 +1,142 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Profile exact model-preparation collaborators without copying orchestration."""
from __future__ import annotations
from typing import ClassVar
import torch
from simple_syrup.domain.processed_regional_attention import (
ProcessedRegionalAttentionPlan,
)
from simple_syrup.domain.raw_regional_attention import RawRegionalAttentionPlan
from simple_syrup.runtime.attention_coupling.context_validation import (
RegionalContextValidator,
)
from simple_syrup.runtime.comfy_conditioning_model_loader import (
ComfyConditioningModelLoader,
)
from simple_syrup.runtime.comfy_conditioning_processing import (
ComfyRegionalConditioningProcessor,
)
from simple_syrup.runtime.regional_lora_conditioning_adapter import (
RegionalLoraConditioningAdapter,
)
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
from simple_syrup.runtime.regional_model_patch_interop import (
RegionalModelPatchInteropValidator,
)
from simple_syrup.services.attention_coupling_model_family_selector import (
AttentionCouplingModelFamilySelector,
)
from simple_syrup.services.attention_coupling_model_preparation_service import (
AttentionCouplingModelPreparationService,
)
from simple_syrup.services.attention_coupling_preparation_service import (
AttentionCouplingPreparation,
AttentionCouplingPreparationService,
)
from .interop_validation_profile import ProfiledRegionalModelPatchInteropValidator
from .model_family_profile import ProfiledAttentionCouplingModelFamilySelector
from .synchronized_phase_timing import measure_synchronized_phase, model_device
class ProfiledComfyConditioningModelLoader(ComfyConditioningModelLoader):
"""Time exact source-model loading through the production owner."""
def load(self, model: object) -> None:
"""Delegate source loading while recording its synchronized duration."""
with measure_synchronized_phase(
"source_model_load",
device=model_device(model),
):
return super().load(model)
class ProfiledRegionalLoraConditioningAdapter(RegionalLoraConditioningAdapter):
"""Time exact regional-LoRA collection and adaptation."""
def adapt(
self,
plan: RawRegionalAttentionPlan,
*,
model: object,
) -> RegionalLoraPlanAdaptation:
"""Delegate adaptation while recording its CPU-owned duration."""
with measure_synchronized_phase(
"regional_lora_adaptation",
device=model_device(model),
):
return super().adapt(plan, model=model)
class ProfiledAttentionCouplingPreparationService(AttentionCouplingPreparationService):
"""Time exact architecture-neutral regional-plan preparation."""
def prepare(self, plan: RawRegionalAttentionPlan) -> AttentionCouplingPreparation:
"""Delegate plan preparation while recording its CPU duration."""
with measure_synchronized_phase(
"regional_plan_preparation",
device=None,
):
return super().prepare(plan)
class ProfiledComfyRegionalConditioningProcessor(ComfyRegionalConditioningProcessor):
"""Time exact Comfy conditioning conversion and encoding."""
def process(
self,
preparation: AttentionCouplingPreparation,
*,
model: object,
noise: torch.Tensor,
device: torch.device,
context_validator: RegionalContextValidator,
) -> ProcessedRegionalAttentionPlan:
"""Delegate conditioning processing with synchronized device timing."""
with measure_synchronized_phase(
"conditioning_processing",
device=device,
):
return super().process(
preparation,
model=model,
noise=noise,
device=device,
context_validator=context_validator,
)
class ProfiledAttentionCouplingModelPreparationService(
AttentionCouplingModelPreparationService
):
"""Run production orchestration with exact-delegating timed collaborators."""
model_loader_class: ClassVar[type[ComfyConditioningModelLoader]] = (
ProfiledComfyConditioningModelLoader
)
lora_adapter_class: ClassVar[type[RegionalLoraConditioningAdapter]] = (
ProfiledRegionalLoraConditioningAdapter
)
preparation_service_class: ClassVar[type[AttentionCouplingPreparationService]] = (
ProfiledAttentionCouplingPreparationService
)
conditioning_processor_class: ClassVar[type[ComfyRegionalConditioningProcessor]] = (
ProfiledComfyRegionalConditioningProcessor
)
interop_validator_class: ClassVar[type[RegionalModelPatchInteropValidator]] = (
ProfiledRegionalModelPatchInteropValidator
)
model_family_selector_class: ClassVar[
type[AttentionCouplingModelFamilySelector]
] = ProfiledAttentionCouplingModelFamilySelector
@@ -0,0 +1,61 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Measure benchmark phases around exact delegates with device synchronization."""
from __future__ import annotations
import logging
import time
from collections.abc import Iterator
from contextlib import contextmanager
import torch
_LOGGER = logging.getLogger(
"simple_syrup.runtime.regional_lora.standard_unet_cold_path"
)
@contextmanager
def measure_synchronized_phase(
stage: str,
*,
device: torch.device | None,
) -> Iterator[None]:
"""Emit one synchronized phase only while cold capture enables DEBUG."""
if not _LOGGER.isEnabledFor(logging.DEBUG):
yield
return
synchronize_device(device)
started_at_ns = time.perf_counter_ns()
try:
yield
finally:
synchronize_device(device)
_LOGGER.debug(
"Measured benchmark Attention Coupling phase",
extra={
"cold_path_diagnostics": {
"stage": stage,
"elapsed_ms": (time.perf_counter_ns() - started_at_ns)
/ 1_000_000.0,
}
},
)
def model_device(model: object) -> torch.device | None:
"""Return a valid Comfy load device when the object exposes one."""
device = getattr(model, "load_device", None)
return device if isinstance(device, torch.device) else None
def synchronize_device(device: torch.device | None) -> None:
"""Synchronize only a live CUDA device explicitly owned by the phase."""
if device is not None and device.type == "cuda" and torch.cuda.is_available():
torch.cuda.synchronize(device)
@@ -0,0 +1,77 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Trace upstream Comfy nodes around one standard-UNet regional cold request."""
from __future__ import annotations
import argparse
import logging
from collections.abc import Sequence
from pathlib import Path
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.sdxl_attention_couple_parity.cases import load_parity_case
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
from tools.sdxl_attention_coupling_integration.visual_lora_baseline_cases import (
SdxlVisualPromptSet,
)
from tools.sdxl_cold_upstream_trace.runner import run_cold_upstream_trace
from tools.sdxl_full_strength_lora_fidelity.composition_cases import (
full_strength_composition_cases,
)
LOGGER = logging.getLogger(__name__)
DEFAULT_OUTPUT_ROOT = Path(
r"<COMFY_ROOT>\benchmark_artifacts\universal-regional-adapter\sdxl-cold-upstream-trace"
)
def main(argv: Sequence[str] | None = None) -> int:
"""Load external fixtures and run one managed image-free node trace."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--inventory", type=Path, required=True)
parser.add_argument("--prompt-case", type=Path, required=True)
parser.add_argument("--comfy-root", type=Path, default=Path(r"<COMFY_ROOT>"))
parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT)
args = parser.parse_args(argv)
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
artifacts = IntegrationArtifacts(args.output_root)
try:
inventory = SdxlVisualInventory.load(args.inventory)
prompts = load_parity_case(args.prompt_case)
prompt_set = SdxlVisualPromptSet(
base_positive_g=prompts.base_positive_g,
base_positive_l=prompts.base_positive_l,
base_negative_g=prompts.base_negative_g,
base_negative_l=prompts.base_negative_l,
left_positive_g=prompts.left_positive_g,
left_positive_l=prompts.left_positive_l,
right_positive_g=prompts.right_positive_g,
right_positive_l=prompts.right_positive_l,
left_negative_g=prompts.left_negative_g,
left_negative_l=prompts.left_negative_l,
right_negative_g=prompts.right_negative_g,
right_negative_l=prompts.right_negative_l,
)
case = full_strength_composition_cases(inventory, prompt_set)[2]
result = run_cold_upstream_trace(
artifacts,
inventory=inventory,
case=case,
comfy_root=args.comfy_root,
)
except BaseException as error:
artifacts.record_failure(error)
LOGGER.exception("Cold upstream trace failed at %s", artifacts.root)
return 1
LOGGER.info("Cold upstream trace completed: %s", result)
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,77 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Compare CPU and selected-device standard-UNet regional LoRA materialization."""
from __future__ import annotations
import argparse
import logging
from collections.abc import Sequence
from pathlib import Path
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.sdxl_attention_couple_parity.cases import load_parity_case
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
from tools.sdxl_attention_coupling_integration.visual_lora_baseline_cases import (
SdxlVisualPromptSet,
)
from tools.sdxl_full_strength_lora_fidelity.composition_cases import (
full_strength_composition_cases,
)
from tools.sdxl_materialization_parity.runner import run_materialization_parity
LOGGER = logging.getLogger(__name__)
DEFAULT_OUTPUT_ROOT = Path(
r"<COMFY_ROOT>\benchmark_artifacts\universal-regional-adapter\sdxl-materialization-parity"
)
def main(argv: Sequence[str] | None = None) -> int:
"""Load external fixtures and run one image-free managed comparison."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--inventory", type=Path, required=True)
parser.add_argument("--prompt-case", type=Path, required=True)
parser.add_argument("--comfy-root", type=Path, default=Path(r"<COMFY_ROOT>"))
parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT)
args = parser.parse_args(argv)
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
artifacts = IntegrationArtifacts(args.output_root)
try:
inventory = SdxlVisualInventory.load(args.inventory)
prompts = load_parity_case(args.prompt_case)
prompt_set = SdxlVisualPromptSet(
base_positive_g=prompts.base_positive_g,
base_positive_l=prompts.base_positive_l,
base_negative_g=prompts.base_negative_g,
base_negative_l=prompts.base_negative_l,
left_positive_g=prompts.left_positive_g,
left_positive_l=prompts.left_positive_l,
right_positive_g=prompts.right_positive_g,
right_positive_l=prompts.right_positive_l,
left_negative_g=prompts.left_negative_g,
left_negative_l=prompts.left_negative_l,
right_negative_g=prompts.right_negative_g,
right_negative_l=prompts.right_negative_l,
)
case = full_strength_composition_cases(inventory, prompt_set)[2]
result = run_materialization_parity(
artifacts,
inventory=inventory,
case=case,
comfy_root=args.comfy_root,
)
except BaseException as error:
artifacts.record_failure(error)
LOGGER.exception("Materialization parity failed at %s", artifacts.root)
return 1
LOGGER.info("Materialization parity completed: %s", result)
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,79 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Attribute one cold and one warmed two-regional-LoRA SDXL request."""
from __future__ import annotations
import argparse
import logging
from collections.abc import Sequence
from pathlib import Path
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.sdxl_attention_couple_parity.cases import load_parity_case
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
from tools.sdxl_attention_coupling_integration.visual_lora_baseline_cases import (
SdxlVisualPromptSet,
)
from tools.sdxl_full_strength_lora_fidelity.composition_cases import (
full_strength_composition_cases,
)
from tools.sdxl_regional_lora_performance.cold_path_runner import (
run_cold_path_attribution,
)
LOGGER = logging.getLogger(__name__)
DEFAULT_OUTPUT_ROOT = Path(
r"<COMFY_ROOT>\benchmark_artifacts\universal-regional-adapter\sdxl-cold-path"
)
def main(argv: Sequence[str] | None = None) -> int:
"""Load external fixtures and run the image-free attribution matrix."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--inventory", type=Path, required=True)
parser.add_argument("--prompt-case", type=Path, required=True)
parser.add_argument("--comfy-root", type=Path, default=Path(r"<COMFY_ROOT>"))
parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT)
args = parser.parse_args(argv)
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
artifacts = IntegrationArtifacts(args.output_root)
try:
inventory = SdxlVisualInventory.load(args.inventory)
prompts = load_parity_case(args.prompt_case)
prompt_set = SdxlVisualPromptSet(
base_positive_g=prompts.base_positive_g,
base_positive_l=prompts.base_positive_l,
base_negative_g=prompts.base_negative_g,
base_negative_l=prompts.base_negative_l,
left_positive_g=prompts.left_positive_g,
left_positive_l=prompts.left_positive_l,
right_positive_g=prompts.right_positive_g,
right_positive_l=prompts.right_positive_l,
left_negative_g=prompts.left_negative_g,
left_negative_l=prompts.left_negative_l,
right_negative_g=prompts.right_negative_g,
right_negative_l=prompts.right_negative_l,
)
case = full_strength_composition_cases(inventory, prompt_set)[2]
result = run_cold_path_attribution(
artifacts,
inventory=inventory,
case=case,
comfy_root=args.comfy_root,
)
except BaseException as error:
artifacts.record_failure(error)
LOGGER.exception("Cold-path SDXL attribution failed at %s", artifacts.root)
return 1
LOGGER.info("Cold-path SDXL attribution completed: %s", result)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+9
View File
@@ -57,6 +57,15 @@ class LoopbackComfyClient:
self._poll_interval = poll_interval
self._client_id = str(uuid.uuid4())
@property
def websocket_url(self) -> str:
"""Return the matching loopback WebSocket endpoint for this session."""
parsed = urllib.parse.urlparse(self._base_url)
host = parsed.netloc
query = urllib.parse.urlencode({"clientId": self._client_id})
return f"ws://{host}/ws?{query}"
def verify_server(self, required_node_ids: Collection[str]) -> JsonObject:
"""Verify readiness and the caller's required node contracts."""
+214
View File
@@ -0,0 +1,214 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Trace ordered Comfy node execution over one loopback WebSocket session."""
from __future__ import annotations
import json
import time
from dataclasses import dataclass
from importlib import import_module
from typing import Any
from tools.comfy_api import JsonObject, LoopbackComfyClient
@dataclass(frozen=True, slots=True)
class ComfyNodeExecutionTiming:
"""Retain one node's observed start and duration until the next event."""
node_id: str
class_type: str
started_at_ns: int
elapsed_ms: float
@dataclass(frozen=True, slots=True)
class ComfyExecutionTraceResult:
"""Retain one completed prompt history and its ordered node timings."""
prompt_id: str
submission_started_at_ns: int
completed_at_ns: int
cached_node_ids: tuple[str, ...]
nodes: tuple[ComfyNodeExecutionTiming, ...]
history: JsonObject
@property
def elapsed_ms(self) -> float:
"""Return observed submission-to-terminal wall time."""
return (self.completed_at_ns - self.submission_started_at_ns) / 1_000_000.0
def class_totals_ms(self) -> dict[str, float]:
"""Aggregate ordered node durations by exact runtime class type."""
totals: dict[str, float] = {}
for node in self.nodes:
totals[node.class_type] = totals.get(node.class_type, 0.0) + node.elapsed_ms
return totals
class ComfyExecutionTraceAccumulator:
"""Convert WebSocket lifecycle messages into ordered node intervals."""
def __init__(
self,
*,
prompt_id: str,
prompt: dict[str, JsonObject],
submission_started_at_ns: int,
) -> None:
"""Retain immutable prompt identity and initialize empty trace state."""
if not prompt_id:
raise ValueError("Execution trace prompt id must be non-empty.")
if submission_started_at_ns < 1:
raise ValueError("Execution trace submission timestamp must be positive.")
self._prompt_id = prompt_id
self._prompt = prompt
self._submission_started_at_ns = submission_started_at_ns
self._cached_node_ids: tuple[str, ...] = ()
self._current: tuple[str, int] | None = None
self._nodes: list[ComfyNodeExecutionTiming] = []
self._completed_at_ns: int | None = None
def observe(self, message: object, *, observed_at_ns: int) -> bool:
"""Consume one JSON message and report terminal completion."""
if observed_at_ns < self._submission_started_at_ns:
raise ValueError("Execution trace observation precedes submission.")
if not isinstance(message, dict):
return False
event_type = message.get("type")
data = message.get("data")
if not isinstance(event_type, str) or not isinstance(data, dict):
return False
message_prompt_id = data.get("prompt_id")
if message_prompt_id != self._prompt_id:
return False
if event_type == "execution_cached":
self._cached_node_ids = self._cached_nodes(data.get("nodes"))
return False
if event_type == "executing":
node_id = data.get("node")
if not isinstance(node_id, str) or not node_id:
raise ValueError("Execution trace node id is invalid.")
self._close_current(observed_at_ns)
self._current = (node_id, observed_at_ns)
return False
if event_type == "execution_error":
raise RuntimeError(
"Comfy execution trace observed an error: "
f"{data.get('exception_type')}: {data.get('exception_message')}"
)
if event_type == "execution_success":
self._close_current(observed_at_ns)
self._completed_at_ns = observed_at_ns
return True
return False
def finish(self, history: JsonObject) -> ComfyExecutionTraceResult:
"""Return a completed trace or fail on missing lifecycle evidence."""
if self._completed_at_ns is None:
raise TimeoutError("Comfy execution trace did not reach terminal success.")
if not self._nodes:
raise ValueError("Comfy execution trace observed no node starts.")
return ComfyExecutionTraceResult(
prompt_id=self._prompt_id,
submission_started_at_ns=self._submission_started_at_ns,
completed_at_ns=self._completed_at_ns,
cached_node_ids=self._cached_node_ids,
nodes=tuple(self._nodes),
history=history,
)
def _close_current(self, completed_at_ns: int) -> None:
"""Close the current node interval against one observed boundary."""
if self._current is None:
return
node_id, started_at_ns = self._current
if completed_at_ns < started_at_ns:
raise ValueError("Execution trace node completion precedes its start.")
node = self._prompt.get(node_id)
class_type = node.get("class_type") if isinstance(node, dict) else None
if not isinstance(class_type, str) or not class_type:
raise ValueError(f"Execution trace node {node_id!r} has no class type.")
self._nodes.append(
ComfyNodeExecutionTiming(
node_id=node_id,
class_type=class_type,
started_at_ns=started_at_ns,
elapsed_ms=(completed_at_ns - started_at_ns) / 1_000_000.0,
)
)
self._current = None
@staticmethod
def _cached_nodes(value: object) -> tuple[str, ...]:
"""Return canonical cached node ids from one lifecycle message."""
if not isinstance(value, list) or any(
not isinstance(item, str) for item in value
):
raise ValueError("Execution trace cached node ids are invalid.")
return tuple(value)
class LoopbackExecutionTrace:
"""Submit one prompt and capture its loopback WebSocket node lifecycle."""
def execute(
self,
client: LoopbackComfyClient,
*,
prompt: dict[str, JsonObject],
timeout: float,
) -> ComfyExecutionTraceResult:
"""Return one synchronized trace and completed history."""
if timeout <= 0.0:
raise ValueError("Execution trace timeout must be positive.")
websocket: Any = import_module("websocket")
connection: Any = websocket.create_connection(
client.websocket_url,
timeout=min(timeout, 5.0),
)
try:
submission_started_at_ns = time.perf_counter_ns()
prompt_id = client.submit(prompt)
accumulator = ComfyExecutionTraceAccumulator(
prompt_id=prompt_id,
prompt=prompt,
submission_started_at_ns=submission_started_at_ns,
)
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
connection.settimeout(min(1.0, max(deadline - time.monotonic(), 0.01)))
try:
raw_message = connection.recv()
except Exception as error:
if isinstance(error, websocket.WebSocketTimeoutException):
continue
raise
if not isinstance(raw_message, str):
continue
decoded: object = json.loads(raw_message)
if accumulator.observe(
decoded,
observed_at_ns=time.perf_counter_ns(),
):
history = client.wait_for_history(prompt_id, timeout=timeout)
return accumulator.finish(history)
raise TimeoutError(
f"Comfy execution trace exceeded {timeout}s: {prompt_id}."
)
finally:
connection.close()
LOOPBACK_EXECUTION_TRACE = LoopbackExecutionTrace()
@@ -12,6 +12,12 @@ from collections.abc import Sequence
from pathlib import Path
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.sdxl_attention_coupling_integration.sampling_controls import (
SDXL_VISUAL_SAMPLING,
)
from tools.sdxl_attention_coupling_integration.visual_case_selection import (
select_visual_cases,
)
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
@@ -32,13 +38,15 @@ DEFAULT_OUTPUT_ROOT = Path(
def main(argv: Sequence[str] | None = None) -> int:
"""Load locked fixtures and execute exactly two labeled artifacts."""
"""Load locked fixtures and execute the explicitly selected artifacts."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--inventory", type=Path, required=True)
parser.add_argument("--prompt-case", type=Path, required=True)
parser.add_argument("--comfy-root", type=Path, default=Path(r"<COMFY_ROOT>"))
parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT)
parser.add_argument("--case-id", action="append", default=[])
parser.add_argument("--seed", type=int, default=SDXL_VISUAL_SAMPLING.seed)
parser.add_argument("--readiness-timeout", type=float, default=240.0)
parser.add_argument("--prompt-timeout", type=float, default=1200.0)
args = parser.parse_args(argv)
@@ -47,13 +55,18 @@ def main(argv: Sequence[str] | None = None) -> int:
try:
inventory = SdxlVisualInventory.load(args.inventory)
prompt_set = load_visual_prompt_set(args.prompt_case)
cases = select_visual_cases(
post_optimization_visual_cases(inventory, prompt_set),
tuple(args.case_id),
)
result = execute_visual_cases(
artifacts,
inventory=inventory,
cases=post_optimization_visual_cases(inventory, prompt_set),
cases=cases,
comfy_root=args.comfy_root,
readiness_timeout=args.readiness_timeout,
prompt_timeout=args.prompt_timeout,
seed=args.seed,
)
except BaseException as error:
artifacts.record_failure(error)
@@ -42,6 +42,7 @@ def add_sampler_branch(
latent: NodeReference,
vae: NodeReference,
run_id: str,
seed: int,
steps: int,
denoise: float,
regional_prompt_weight: float,
@@ -66,7 +67,7 @@ def add_sampler_branch(
sampled = graph.add(
node_id,
model=[captured, 0],
seed=SDXL_VISUAL_SAMPLING.seed,
seed=seed,
steps=steps,
cfg=SDXL_VISUAL_SAMPLING.cfg,
sampler_name=SDXL_VISUAL_SAMPLING.sampler,
@@ -8,6 +8,8 @@ from __future__ import annotations
from dataclasses import dataclass
MAX_COMFY_SEED = 0xFFFFFFFFFFFFFFFF
@dataclass(frozen=True, slots=True)
class SdxlVisualSamplingControls:
@@ -27,3 +29,13 @@ SDXL_VISUAL_SAMPLING = SdxlVisualSamplingControls(
sampler="euler_ancestral",
scheduler="karras",
)
def validate_sdxl_visual_seed(seed: object) -> int:
"""Return one seed accepted by Comfy's sampler schema."""
if isinstance(seed, bool) or not isinstance(seed, int):
raise TypeError("SDXL visual seed must be an integer.")
if not 0 <= seed <= MAX_COMFY_SEED:
raise ValueError(f"SDXL visual seed must be between 0 and {MAX_COMFY_SEED}.")
return seed
@@ -20,6 +20,7 @@ from tools.comfy_integration.managed_server import (
from .comfy_model_root import resolve_active_comfy_model_root
from .managed_model_links import ManagedModelLink, ManagedSdxlVisualModelLinks
from .sampling_controls import SDXL_VISUAL_SAMPLING, validate_sdxl_visual_seed
from .visual_adapter_selections import (
CHECKPOINT_SELECTION,
LEFT_CHARACTER_SELECTION,
@@ -46,11 +47,13 @@ def execute_visual_cases(
comfy_root: Path,
readiness_timeout: float,
prompt_timeout: float,
seed: int = SDXL_VISUAL_SAMPLING.seed,
) -> Path:
"""Execute every explicit case in one managed server trajectory."""
if not cases:
raise ValueError("SDXL visual execution requires at least one case.")
validated_seed = validate_sdxl_visual_seed(seed)
recorder = SdxlVisualResultRecorder(artifacts.root, cases=cases)
model_links = build_sdxl_visual_model_links(comfy_root, inventory)
masks = ManagedSdxlVisualMasks(
@@ -70,6 +73,7 @@ def execute_visual_cases(
checkpoint_name=CHECKPOINT_SELECTION,
mask_names=masks.names(case.mask_profile),
case=case,
seed=validated_seed,
),
)
for case in cases
@@ -19,7 +19,6 @@ from .evidence_validation import (
validate_sdxl_metrics,
)
from .matrix import MODES, SdxlIntegrationMode
from .sampling_controls import SDXL_VISUAL_SAMPLING
from .visual_case_model import RegionalVisualAdapter, SdxlVisualCase, VisualMode
from .visual_history import SdxlVisualHistoryEvidence
from .visual_runtime_expectations import SDXL_VISUAL_RUNTIME_EXPECTATIONS
@@ -98,7 +97,7 @@ class SdxlVisualResultRecorder:
"case_id": case.case_id,
"label": case.label,
"mode": output.mode.value,
"seed": SDXL_VISUAL_SAMPLING.seed,
"seed": workflow.seed,
"mask_profile": case.mask_profile.value,
"regional_prompt_start_percent": (
case.regional_prompt_start_percent
@@ -24,7 +24,7 @@ from .matrix import (
TILE_SIZE,
)
from .sampler_branch import SdxlWorkflowOutputs, add_sampler_branch
from .sampling_controls import SDXL_VISUAL_SAMPLING
from .sampling_controls import SDXL_VISUAL_SAMPLING, validate_sdxl_visual_seed
from .visual_case_model import SdxlVisualCase, VisualMode
from .visual_conditioning import SdxlVisualConditioningBuilder
@@ -46,6 +46,7 @@ class BuiltSdxlVisualWorkflow:
prompt: dict[str, JsonObject]
outputs: tuple[SdxlVisualWorkflowOutput, ...]
seed: int
@property
def required_node_ids(self) -> frozenset[str]:
@@ -60,9 +61,11 @@ def build_sdxl_visual_workflow(
checkpoint_name: str,
mask_names: tuple[str, str],
case: SdxlVisualCase,
seed: int = SDXL_VISUAL_SAMPLING.seed,
) -> BuiltSdxlVisualWorkflow:
"""Build one full source and only the case's declared refinement modes."""
validated_seed = validate_sdxl_visual_seed(seed)
graph = SdxlWorkflowGraph()
loader = graph.add("CheckpointLoaderSimple", ckpt_name=checkpoint_name)
conditioning = SdxlVisualConditioningBuilder().build(
@@ -94,6 +97,7 @@ def build_sdxl_visual_workflow(
latent=[source_latent, 0],
vae=[loader, 2],
run_id=case_run_id,
seed=validated_seed,
steps=SDXL_VISUAL_SAMPLING.steps,
denoise=1.0,
regional_prompt_weight=case.regional_prompt_weight,
@@ -131,6 +135,7 @@ def build_sdxl_visual_workflow(
latent=[upscaled_latent, 0],
vae=[loader, 2],
run_id=case_run_id,
seed=validated_seed,
steps=REFINEMENT_STEPS,
denoise=REFINEMENT_DENOISE,
regional_prompt_weight=case.regional_prompt_weight,
@@ -141,7 +146,7 @@ def build_sdxl_visual_workflow(
output_records.append(_output_record(case, mode, branch.outputs))
if tuple(record.mode for record in output_records) != case.modes:
raise ValueError("U11 workflow output order must match the case declaration.")
return BuiltSdxlVisualWorkflow(graph.prompt, tuple(output_records))
return BuiltSdxlVisualWorkflow(graph.prompt, tuple(output_records), validated_seed)
def _output_record(
@@ -91,6 +91,7 @@ def build_sdxl_attention_coupling_workflow(
latent=[source_latent, 0],
vae=[loader, 2],
run_id=run_id,
seed=SDXL_VISUAL_SAMPLING.seed,
steps=SDXL_VISUAL_SAMPLING.steps,
denoise=1.0,
regional_prompt_weight=REGIONAL_PROMPT_WEIGHT,
@@ -122,6 +123,7 @@ def build_sdxl_attention_coupling_workflow(
latent=[upscaled_latent, 0],
vae=[loader, 2],
run_id=run_id,
seed=SDXL_VISUAL_SAMPLING.seed,
steps=REFINEMENT_STEPS,
denoise=REFINEMENT_DENOISE,
regional_prompt_weight=REGIONAL_PROMPT_WEIGHT,
@@ -146,6 +148,7 @@ def build_sdxl_attention_coupling_workflow(
latent=[upscaled_latent, 0],
vae=[loader, 2],
run_id=run_id,
seed=SDXL_VISUAL_SAMPLING.seed,
steps=REFINEMENT_STEPS,
denoise=REFINEMENT_DENOISE,
regional_prompt_weight=REGIONAL_PROMPT_WEIGHT,
@@ -0,0 +1,5 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Attribute upstream Comfy graph work around one cold regional request."""
@@ -0,0 +1,39 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Substitute only the dev-only profiled sampler in one cold trace graph."""
from __future__ import annotations
from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase
from tools.sdxl_regional_lora_performance.cold_path_workflow import (
BuiltSdxlColdPathWorkflow,
build_sdxl_cold_path_workflow,
)
_PRODUCTION_SAMPLER = "SimpleSyrup.KSamplerAttentionCoupling"
_PROFILED_SAMPLER = "SimpleSyrupBenchmark.ProfiledKSamplerAttentionCoupling"
def build_profiled_cold_path_workflow(
*,
checkpoint_name: str,
mask_names: tuple[str, str],
case: SdxlVisualCase,
) -> BuiltSdxlColdPathWorkflow:
"""Return the unchanged cold graph with one sampler identity substitution."""
built = build_sdxl_cold_path_workflow(
checkpoint_name=checkpoint_name,
mask_names=mask_names,
case=case,
)
sampler = built.prompt.get(built.sampler_node_id)
if (
not isinstance(sampler, dict)
or sampler.get("class_type") != _PRODUCTION_SAMPLER
):
raise ValueError("Cold trace lost its production Attention Coupling sampler.")
sampler["class_type"] = _PROFILED_SAMPLER
return built
+100
View File
@@ -0,0 +1,100 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Decompose cold wall time across traced Comfy nodes and sampler stages."""
from __future__ import annotations
from dataclasses import dataclass
from tools.comfy_integration.execution_trace import ComfyExecutionTraceResult
from tools.sdxl_regional_lora_performance.cold_path_results import (
SdxlColdPathTiming,
)
_TOP_LEVEL_STAGES = frozenset(
{
"admission_resolution",
"template_preparation",
"model_residency",
"sampling",
}
)
@dataclass(frozen=True, slots=True)
class SdxlColdUpstreamAttribution:
"""Retain one complete node/stage decomposition of cold wall time."""
submission_to_first_node_ms: float
pre_sampler_node_ms: float
sampler_node_ms: float
sampler_instrumented_ms: float
sampler_unattributed_ms: float
post_sampler_node_ms: float
trace_residual_ms: float
cold_unattributed_ms: float
named_upstream_fraction: float
upstream_class_totals_ms: dict[str, float]
all_class_totals_ms: dict[str, float]
def attribute_cold_upstream(
trace: ComfyExecutionTraceResult,
cold: SdxlColdPathTiming,
*,
sampler_node_id: str,
) -> SdxlColdUpstreamAttribution:
"""Decompose traced graph work around one instrumented sampler node."""
sampler_indices = tuple(
index
for index, timing in enumerate(trace.nodes)
if timing.node_id == sampler_node_id
)
if len(sampler_indices) != 1:
raise ValueError("Cold upstream trace requires exactly one sampler node.")
sampler_index = sampler_indices[0]
pre_sampler = trace.nodes[:sampler_index]
sampler = trace.nodes[sampler_index]
post_sampler = trace.nodes[sampler_index + 1 :]
first_started_at_ns = trace.nodes[0].started_at_ns
submission_to_first_ms = (
first_started_at_ns - trace.submission_started_at_ns
) / 1_000_000.0
stage_totals: dict[str, float] = {}
for stage in cold.stages:
stage_totals[stage.stage] = (
stage_totals.get(stage.stage, 0.0) + stage.elapsed_ms
)
sampler_instrumented_ms = sum(
stage_totals.get(stage, 0.0) for stage in _TOP_LEVEL_STAGES
)
cold_unattributed_ms = cold.runtime_ms - sampler_instrumented_ms
pre_sampler_ms = sum(node.elapsed_ms for node in pre_sampler)
post_sampler_ms = sum(node.elapsed_ms for node in post_sampler)
sampler_unattributed_ms = sampler.elapsed_ms - sampler_instrumented_ms
observed_components = (
submission_to_first_ms + pre_sampler_ms + sampler.elapsed_ms + post_sampler_ms
)
upstream_totals: dict[str, float] = {}
for node in pre_sampler:
upstream_totals[node.class_type] = (
upstream_totals.get(node.class_type, 0.0) + node.elapsed_ms
)
return SdxlColdUpstreamAttribution(
submission_to_first_node_ms=submission_to_first_ms,
pre_sampler_node_ms=pre_sampler_ms,
sampler_node_ms=sampler.elapsed_ms,
sampler_instrumented_ms=sampler_instrumented_ms,
sampler_unattributed_ms=sampler_unattributed_ms,
post_sampler_node_ms=post_sampler_ms,
trace_residual_ms=trace.elapsed_ms - observed_components,
cold_unattributed_ms=cold_unattributed_ms,
named_upstream_fraction=(
pre_sampler_ms / cold_unattributed_ms if cold_unattributed_ms > 0.0 else 0.0
),
upstream_class_totals_ms=upstream_totals,
all_class_totals_ms=trace.class_totals_ms(),
)
+162
View File
@@ -0,0 +1,162 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Run one managed node-level trace of the unchanged cold regional graph."""
from __future__ import annotations
import json
from dataclasses import asdict
from pathlib import Path
from tools.comfy_api import JsonObject
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.comfy_integration.execution_trace import LOOPBACK_EXECUTION_TRACE
from tools.comfy_integration.loopback_port import is_loopback_port_available
from tools.comfy_integration.managed_server import ManagedComfyServer
from tools.sdxl_attention_coupling_integration.sampling_controls import (
SDXL_VISUAL_SAMPLING,
)
from tools.sdxl_attention_coupling_integration.visual_adapter_selections import (
CHECKPOINT_SELECTION,
)
from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
from tools.sdxl_attention_coupling_integration.visual_launch import (
sdxl_visual_sampling_launch_arguments,
)
from tools.sdxl_attention_coupling_integration.visual_masks import (
ManagedSdxlVisualMasks,
)
from tools.sdxl_attention_coupling_integration.visual_matrix_execution import (
build_sdxl_visual_model_links,
)
from tools.sdxl_regional_lora_performance.cold_path_priming import (
build_cold_path_primers,
execute_cold_path_primers,
)
from tools.sdxl_regional_lora_performance.cold_path_results import (
decode_cold_path_timing,
)
from tools.sdxl_regional_lora_performance.cold_path_workflow import (
BuiltSdxlColdPathWorkflow,
)
from .profile_workflow import build_profiled_cold_path_workflow
from .results import attribute_cold_upstream
def run_cold_upstream_trace(
artifacts: IntegrationArtifacts,
*,
inventory: SdxlVisualInventory,
case: SdxlVisualCase,
comfy_root: Path,
readiness_timeout: float = 240.0,
prompt_timeout: float = 1200.0,
) -> Path:
"""Persist one node-level trace plus the existing cold-stage diagnostics."""
links = build_sdxl_visual_model_links(comfy_root, inventory)
masks = ManagedSdxlVisualMasks(
input_root=comfy_root / "input",
run_id=artifacts.run_id,
)
system_stats: JsonObject = {}
server_cleanup = False
with links:
with masks:
workflow: BuiltSdxlColdPathWorkflow = build_profiled_cold_path_workflow(
checkpoint_name=CHECKPOINT_SELECTION,
mask_names=masks.names(case.mask_profile),
case=case,
)
primers = build_cold_path_primers(
checkpoint_name=CHECKPOINT_SELECTION,
mask_names=masks.names(case.mask_profile),
case=case,
)
required_node_ids = frozenset(
{
*workflow.required_node_ids,
*(
node_id
for primer in primers
for node_id in primer.workflow.required_node_ids
),
}
)
with ManagedComfyServer(
comfy_root=comfy_root,
artifacts=artifacts,
required_node_ids=required_node_ids,
readiness_timeout=readiness_timeout,
launch_arguments=sdxl_visual_sampling_launch_arguments(),
) as running:
system_stats = running.system_stats
primer_timings = execute_cold_path_primers(
running.client,
primers=primers,
seed=SDXL_VISUAL_SAMPLING.seed,
prompt_timeout=prompt_timeout,
)
execution_id = f"{artifacts.run_id}:cold-upstream"
prompt = workflow.prompt_for_execution(
seed=SDXL_VISUAL_SAMPLING.seed,
run_id=execution_id,
)
trace = LOOPBACK_EXECUTION_TRACE.execute(
running.client,
prompt=prompt,
timeout=prompt_timeout,
)
cold = decode_cold_path_timing(
trace.history,
terminal_node_id=workflow.terminal_node_id,
started_at_ns=trace.submission_started_at_ns,
seed=SDXL_VISUAL_SAMPLING.seed,
)
attribution = attribute_cold_upstream(
trace,
cold,
sampler_node_id=workflow.sampler_node_id,
)
port = running.port
process = running.process
server_cleanup = not process.is_running and is_loopback_port_available(port)
artifacts.record_cleanup(
process_running=process.is_running,
port_available=is_loopback_port_available(port),
)
result = artifacts.root / "sdxl-cold-upstream-trace.json"
result.write_text(
json.dumps(
{
"status": "completed",
"cold": asdict(cold),
"trace": {
"elapsed_ms": trace.elapsed_ms,
"cached_node_ids": trace.cached_node_ids,
"nodes": [asdict(node) for node in trace.nodes],
},
"attribution": asdict(attribution),
"primers": {
label: asdict(timing) for label, timing in primer_timings.items()
},
"system_stats": system_stats,
"cleanup": {
"server": server_cleanup,
"model_links": links.cleaned,
"masks": masks.cleaned,
},
},
indent=2,
sort_keys=True,
)
+ "\n",
encoding="utf-8",
)
return result
@@ -0,0 +1,5 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Run focused standard-UNet materialization parity evidence."""
@@ -0,0 +1,55 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Decode one managed materialization parity terminal fail-closed."""
from __future__ import annotations
from tools.comfy_api import JsonObject
def decode_materialization_parity(
history: JsonObject,
*,
terminal_node_id: str,
expected_run_id: str,
) -> JsonObject:
"""Return the one complete validated UI comparison object."""
outputs = history.get("outputs")
if not isinstance(outputs, dict):
raise ValueError("Materialization parity history is missing outputs.")
output = outputs.get(terminal_node_id)
if not isinstance(output, dict):
raise ValueError("Materialization parity history is missing its terminal.")
values = output.get("materialization_parity")
if not isinstance(values, list) or len(values) != 1:
raise ValueError("Materialization parity terminal requires one result.")
payload = values[0]
if not isinstance(payload, dict):
raise ValueError("Materialization parity result must be an object.")
if payload.get("run_id") != expected_run_id:
raise ValueError("Materialization parity run identity does not match.")
_require_non_negative_int(payload, "variant_count")
_require_non_negative_int(payload, "element_count")
_require_non_negative_int(payload, "differing_element_count")
_require_non_negative_int(payload, "peak_vram_bytes")
if payload["variant_count"] < 1 or payload["element_count"] < 1:
raise ValueError("Materialization parity result must compare nonempty banks.")
if payload["differing_element_count"] > payload["element_count"]:
raise ValueError("Materialization parity differing count exceeds its total.")
if not isinstance(payload.get("exact"), bool):
raise ValueError("Materialization parity exactness flag is invalid.")
variants = payload.get("variants")
if not isinstance(variants, list) or len(variants) != payload["variant_count"]:
raise ValueError("Materialization parity variant evidence is incomplete.")
return {str(key): value for key, value in payload.items()}
def _require_non_negative_int(payload: dict[object, object], key: str) -> None:
"""Validate one required non-negative integer result field."""
value = payload.get(key)
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValueError(f"Materialization parity field {key!r} is invalid.")
+113
View File
@@ -0,0 +1,113 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Run one managed sampler-free materialization parity comparison."""
from __future__ import annotations
import json
import time
from pathlib import Path
from tools.comfy_api import JsonObject
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.comfy_integration.loopback_port import is_loopback_port_available
from tools.comfy_integration.managed_server import ManagedComfyServer
from tools.sdxl_attention_coupling_integration.visual_adapter_selections import (
CHECKPOINT_SELECTION,
)
from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
from tools.sdxl_attention_coupling_integration.visual_launch import (
sdxl_visual_sampling_launch_arguments,
)
from tools.sdxl_attention_coupling_integration.visual_masks import (
ManagedSdxlVisualMasks,
)
from tools.sdxl_attention_coupling_integration.visual_matrix_execution import (
build_sdxl_visual_model_links,
)
from .results import decode_materialization_parity
from .workflow import build_materialization_parity_workflow
def run_materialization_parity(
artifacts: IntegrationArtifacts,
*,
inventory: SdxlVisualInventory,
case: SdxlVisualCase,
comfy_root: Path,
readiness_timeout: float = 240.0,
prompt_timeout: float = 1200.0,
) -> Path:
"""Persist one complete comparison and managed lifecycle evidence."""
links = build_sdxl_visual_model_links(comfy_root, inventory)
masks = ManagedSdxlVisualMasks(
input_root=comfy_root / "input",
run_id=artifacts.run_id,
)
system_stats: JsonObject = {}
server_cleanup = False
started_at_ns = 0
completed_at_ns = 0
with links:
with masks:
workflow = build_materialization_parity_workflow(
run_id=artifacts.run_id,
checkpoint_name=CHECKPOINT_SELECTION,
mask_names=masks.names(case.mask_profile),
case=case,
)
with ManagedComfyServer(
comfy_root=comfy_root,
artifacts=artifacts,
required_node_ids=workflow.required_node_ids,
readiness_timeout=readiness_timeout,
launch_arguments=sdxl_visual_sampling_launch_arguments(),
) as running:
system_stats = running.system_stats
started_at_ns = time.perf_counter_ns()
prompt_id = running.client.submit(workflow.prompt)
history = running.client.wait_for_history(
prompt_id,
timeout=prompt_timeout,
)
completed_at_ns = time.perf_counter_ns()
comparison = decode_materialization_parity(
history,
terminal_node_id=workflow.terminal_node_id,
expected_run_id=artifacts.run_id,
)
port = running.port
process = running.process
server_cleanup = not process.is_running and is_loopback_port_available(port)
artifacts.record_cleanup(
process_running=process.is_running,
port_available=is_loopback_port_available(port),
)
result = artifacts.root / "materialization-parity.json"
result.write_text(
json.dumps(
{
"status": "completed",
"workflow_elapsed_ms": (completed_at_ns - started_at_ns) / 1_000_000.0,
"comparison": comparison,
"system_stats": system_stats,
"cleanup": {
"server": server_cleanup,
"model_links": links.cleaned,
"masks": masks.cleaned,
},
},
indent=2,
sort_keys=True,
)
+ "\n",
encoding="utf-8",
)
return result
@@ -0,0 +1,79 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Build one sampler-free regional materialization parity workflow."""
from __future__ import annotations
from dataclasses import dataclass
from tools.comfy_api import JsonObject
from tools.sdxl_attention_coupling_integration.graph import NodeReference
from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase
from tools.sdxl_regional_lora_performance.two_adapter_workflow import (
TwoAdapterPerformanceMode,
prepare_two_adapter_sampling,
)
_PARITY_NODE = "SimpleSyrupBenchmark.CompareMaterializationParity"
@dataclass(frozen=True, slots=True)
class BuiltMaterializationParityWorkflow:
"""Expose one immutable image-free graph and its terminal identity."""
prompt: dict[str, JsonObject]
terminal_node_id: str
@property
def required_node_ids(self) -> frozenset[str]:
"""Return every Comfy node required by the graph."""
return frozenset(str(node["class_type"]) for node in self.prompt.values())
def build_materialization_parity_workflow(
*,
run_id: str,
checkpoint_name: str,
mask_names: tuple[str, str],
case: SdxlVisualCase,
) -> BuiltMaterializationParityWorkflow:
"""Build the accepted two-regional-LoRA request without a sampler."""
if not isinstance(run_id, str) or not run_id:
raise ValueError("Materialization parity run id must be non-empty.")
prepared = prepare_two_adapter_sampling(
checkpoint_name=checkpoint_name,
mask_names=mask_names,
case=case,
mode=TwoAdapterPerformanceMode.REGIONAL,
)
region_masks = _region_masks(prepared.sampler_extras)
terminal = prepared.graph.add(
_PARITY_NODE,
model=prepared.model,
positive=prepared.positive,
negative=prepared.negative,
region_masks=region_masks,
latent_image=prepared.latent,
region_mask_feather=case.region_mask_feather,
run_id=run_id,
)
return BuiltMaterializationParityWorkflow(prepared.graph.prompt, terminal)
def _region_masks(extras: dict[str, object]) -> NodeReference:
"""Return the exact mask reference retained by regional preparation."""
value = extras.get("region_masks")
if (
not isinstance(value, list)
or len(value) != 2
or not isinstance(value[0], str)
or isinstance(value[1], bool)
or not isinstance(value[1], int)
):
raise ValueError("Materialization parity graph lost its region masks.")
return value
@@ -0,0 +1,85 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Execute one cold and one warmed SDXL attribution request."""
from __future__ import annotations
import time
from dataclasses import dataclass
from typing import Protocol
from tools.comfy_api import JsonObject
from .cold_path_results import SdxlColdPathTiming, decode_cold_path_timing
from .cold_path_workflow import BuiltSdxlColdPathWorkflow
class ColdPathClient(Protocol):
"""Submit Comfy workflows and return their completed histories."""
def submit(self, prompt: dict[str, JsonObject]) -> str:
"""Submit one API-format graph."""
def wait_for_history(self, prompt_id: str, *, timeout: float) -> JsonObject:
"""Wait for one completed prompt history."""
@dataclass(frozen=True, slots=True)
class SdxlColdPathObservation:
"""Retain the predeclared adjacent cold and warm executions."""
cold: SdxlColdPathTiming
warm: SdxlColdPathTiming
def measure_cold_path(
client: ColdPathClient,
*,
workflow: BuiltSdxlColdPathWorkflow,
first_seed: int,
run_id: str,
prompt_timeout: float,
) -> SdxlColdPathObservation:
"""Measure exactly one cache miss followed by one exact warmed reuse."""
if not isinstance(run_id, str) or not run_id:
raise ValueError("Cold-path measurement run id must be non-empty.")
cold = _execute(
client,
workflow=workflow,
seed=first_seed,
execution_id=f"{run_id}:cold",
prompt_timeout=prompt_timeout,
)
warm = _execute(
client,
workflow=workflow,
seed=first_seed + 1,
execution_id=f"{run_id}:warm",
prompt_timeout=prompt_timeout,
)
return SdxlColdPathObservation(cold, warm)
def _execute(
client: ColdPathClient,
*,
workflow: BuiltSdxlColdPathWorkflow,
seed: int,
execution_id: str,
prompt_timeout: float,
) -> SdxlColdPathTiming:
"""Submit and decode one synchronized attribution execution."""
prompt = workflow.prompt_for_execution(seed=seed, run_id=execution_id)
started_at_ns = time.perf_counter_ns()
prompt_id = client.submit(prompt)
history = client.wait_for_history(prompt_id, timeout=prompt_timeout)
return decode_cold_path_timing(
history,
terminal_node_id=workflow.terminal_node_id,
started_at_ns=started_at_ns,
seed=seed,
)
@@ -0,0 +1,115 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Prime the exact historical dependencies before regional cold attribution."""
from __future__ import annotations
import time
from dataclasses import dataclass
from tools.comfy_api import JsonObject
from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase
from .cold_path_measurement import ColdPathClient
from .conventional_variant_workflow import (
ConventionalVariantBranch,
build_conventional_variant_steady_state_workflow,
)
from .steady_state_measurement import (
SdxlSteadyStateTiming,
decode_steady_state_timing,
)
from .steady_state_workflow import BuiltSdxlSteadyStateWorkflow
from .two_adapter_workflow import (
TwoAdapterPerformanceMode,
build_two_adapter_steady_state_workflow,
)
@dataclass(frozen=True, slots=True)
class SdxlColdPathPrimer:
"""Bind one historical primer label to its stable workflow."""
label: str
workflow: BuiltSdxlSteadyStateWorkflow
def build_cold_path_primers(
*,
checkpoint_name: str,
mask_names: tuple[str, str],
case: SdxlVisualCase,
) -> tuple[SdxlColdPathPrimer, ...]:
"""Build the global, left, then right historical preparation order."""
return (
SdxlColdPathPrimer(
TwoAdapterPerformanceMode.GLOBAL_REFERENCE.value,
build_two_adapter_steady_state_workflow(
checkpoint_name=checkpoint_name,
mask_names=mask_names,
case=case,
mode=TwoAdapterPerformanceMode.GLOBAL_REFERENCE,
),
),
*(
SdxlColdPathPrimer(
branch.value,
build_conventional_variant_steady_state_workflow(
checkpoint_name=checkpoint_name,
case=case,
branch=branch,
),
)
for branch in ConventionalVariantBranch
),
)
def execute_cold_path_primers(
client: ColdPathClient,
*,
primers: tuple[SdxlColdPathPrimer, ...],
seed: int,
prompt_timeout: float,
) -> dict[str, SdxlSteadyStateTiming]:
"""Execute each primer once and return synchronized timing evidence."""
if tuple(primer.label for primer in primers) != (
TwoAdapterPerformanceMode.GLOBAL_REFERENCE.value,
ConventionalVariantBranch.LEFT.value,
ConventionalVariantBranch.RIGHT.value,
):
raise ValueError("Cold-path primers must preserve global, left, right order.")
return {
primer.label: _execute_primer(
client,
workflow=primer.workflow,
seed=seed + index,
prompt_timeout=prompt_timeout,
)
for index, primer in enumerate(primers)
}
def _execute_primer(
client: ColdPathClient,
*,
workflow: BuiltSdxlSteadyStateWorkflow,
seed: int,
prompt_timeout: float,
) -> SdxlSteadyStateTiming:
"""Submit one image-free primer and decode its synchronized terminal."""
prompt: dict[str, JsonObject] = workflow.prompt_for_seed(seed)
started_at_ns = time.perf_counter_ns()
prompt_id = client.submit(prompt)
history = client.wait_for_history(prompt_id, timeout=prompt_timeout)
return decode_steady_state_timing(
history,
completion_node_id=workflow.completion_node_id,
started_at_ns=started_at_ns,
seed=seed,
)
@@ -0,0 +1,184 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Decode and validate standard-UNet cold-attribution evidence."""
from __future__ import annotations
from dataclasses import dataclass
from tools.comfy_api import JsonObject
_REQUIRED_COLD_STAGES = frozenset(
{
"admission_resolution",
"variant_materialization",
"variant_shell",
"template_preparation",
"model_residency",
"sampling",
}
)
_TOP_LEVEL_STAGES = (
"admission_resolution",
"template_preparation",
"model_residency",
"sampling",
)
@dataclass(frozen=True, slots=True)
class ColdPathStageTiming:
"""Retain one ordered structured stage observation."""
stage: str
elapsed_ms: float
metadata: JsonObject
@dataclass(frozen=True, slots=True)
class SdxlColdPathTiming:
"""Retain one synchronized cold or warmed workflow execution."""
run_id: str
seed: int
runtime_ms: float
stages: tuple[ColdPathStageTiming, ...]
model_call_count: int
peak_vram_bytes: int
@dataclass(frozen=True, slots=True)
class SdxlColdPathSummary:
"""Attribute the cold interval without double-counting nested stages."""
cold_runtime_ms: float
warm_runtime_ms: float
stage_totals_ms: dict[str, float]
top_level_accounted_ms: float
unattributed_ms: float
materialized_parameter_count: int
materialized_parameter_bytes: int
def decode_cold_path_timing(
history: JsonObject,
*,
terminal_node_id: str,
started_at_ns: int,
seed: int,
) -> SdxlColdPathTiming:
"""Decode one synchronized benchmark terminal and its ordered stages."""
if isinstance(started_at_ns, bool) or not isinstance(started_at_ns, int):
raise TypeError("Cold-path start timestamp must be an integer.")
if started_at_ns < 1:
raise ValueError("Cold-path start timestamp must be positive.")
outputs = history.get("outputs")
if not isinstance(outputs, dict):
raise ValueError("Cold-path history is missing outputs.")
output = outputs.get(terminal_node_id)
if not isinstance(output, dict):
raise ValueError("Cold-path history is missing its terminal node.")
values = output.get("cold_path_diagnostics")
if not isinstance(values, list) or len(values) != 1:
raise ValueError("Cold-path terminal must contain one diagnostic object.")
payload = values[0]
if not isinstance(payload, dict):
raise ValueError("Cold-path terminal diagnostic must be an object.")
run_id = payload.get("run_id")
completed_at_ns = payload.get("completed_at_ns")
records = payload.get("records")
model_call_count = payload.get("model_call_count")
peak_vram_bytes = payload.get("peak_vram_bytes")
if not isinstance(run_id, str) or not run_id:
raise ValueError("Cold-path terminal run id is invalid.")
if isinstance(completed_at_ns, bool) or not isinstance(completed_at_ns, int):
raise ValueError("Cold-path completion timestamp is invalid.")
if completed_at_ns <= started_at_ns:
raise ValueError("Cold-path completion must follow submission start.")
if not isinstance(records, list) or not records:
raise ValueError("Cold-path terminal contains no stage records.")
if isinstance(model_call_count, bool) or not isinstance(model_call_count, int):
raise ValueError("Cold-path model-call count is invalid.")
if isinstance(peak_vram_bytes, bool) or not isinstance(peak_vram_bytes, int):
raise ValueError("Cold-path peak allocation is invalid.")
stages = tuple(_decode_stage(record) for record in records)
return SdxlColdPathTiming(
run_id=run_id,
seed=seed,
runtime_ms=(completed_at_ns - started_at_ns) / 1_000_000.0,
stages=stages,
model_call_count=model_call_count,
peak_vram_bytes=peak_vram_bytes,
)
def summarize_cold_path(
cold: SdxlColdPathTiming,
warm: SdxlColdPathTiming,
) -> SdxlColdPathSummary:
"""Require the complete matrix and calculate non-overlapping attribution."""
cold_stages = {stage.stage for stage in cold.stages}
missing = _REQUIRED_COLD_STAGES - cold_stages
if missing:
raise ValueError(
f"Cold-path attribution is missing stages: {sorted(missing)!r}."
)
warm_stages = {stage.stage for stage in warm.stages}
if not {"model_residency", "sampling"}.issubset(warm_stages):
raise ValueError("Warmed attribution requires residency and sampling stages.")
if cold.model_call_count != 30 or warm.model_call_count != 30:
raise ValueError("Cold-path attribution requires exactly 30 model calls.")
totals: dict[str, float] = {}
for stage in cold.stages:
totals[stage.stage] = totals.get(stage.stage, 0.0) + stage.elapsed_ms
accounted = sum(totals[stage] for stage in _TOP_LEVEL_STAGES)
materialized = tuple(
stage for stage in cold.stages if stage.stage == "variant_materialization"
)
return SdxlColdPathSummary(
cold_runtime_ms=cold.runtime_ms,
warm_runtime_ms=warm.runtime_ms,
stage_totals_ms=totals,
top_level_accounted_ms=accounted,
unattributed_ms=cold.runtime_ms - accounted,
materialized_parameter_count=sum(
_metadata_int(stage, "parameter_count") for stage in materialized
),
materialized_parameter_bytes=sum(
_metadata_int(stage, "parameter_bytes") for stage in materialized
),
)
def _decode_stage(value: object) -> ColdPathStageTiming:
"""Narrow one JSON stage without discarding its bounded metadata."""
if not isinstance(value, dict):
raise ValueError("Cold-path stage must be an object.")
stage = value.get("stage")
elapsed_ms = value.get("elapsed_ms")
if not isinstance(stage, str) or not stage:
raise ValueError("Cold-path stage name is invalid.")
if isinstance(elapsed_ms, bool) or not isinstance(elapsed_ms, int | float):
raise ValueError("Cold-path stage elapsed time is invalid.")
if float(elapsed_ms) < 0.0:
raise ValueError("Cold-path stage elapsed time cannot be negative.")
metadata = {
str(key): item
for key, item in value.items()
if key not in {"stage", "elapsed_ms"}
}
return ColdPathStageTiming(stage, float(elapsed_ms), metadata)
def _metadata_int(stage: ColdPathStageTiming, key: str) -> int:
"""Return one required non-negative integer stage metric."""
value = stage.metadata.get(key)
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValueError(f"Cold-path stage metric {key!r} is invalid.")
return value
@@ -0,0 +1,138 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Run one managed SDXL cold-path attribution matrix."""
from __future__ import annotations
import json
from dataclasses import asdict
from pathlib import Path
from tools.comfy_api import JsonObject
from tools.comfy_integration.artifacts import IntegrationArtifacts
from tools.comfy_integration.loopback_port import is_loopback_port_available
from tools.comfy_integration.managed_server import ManagedComfyServer
from tools.sdxl_attention_coupling_integration.sampling_controls import (
SDXL_VISUAL_SAMPLING,
)
from tools.sdxl_attention_coupling_integration.visual_adapter_selections import (
CHECKPOINT_SELECTION,
)
from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase
from tools.sdxl_attention_coupling_integration.visual_inventory import (
SdxlVisualInventory,
)
from tools.sdxl_attention_coupling_integration.visual_launch import (
sdxl_visual_sampling_launch_arguments,
)
from tools.sdxl_attention_coupling_integration.visual_masks import (
ManagedSdxlVisualMasks,
)
from tools.sdxl_attention_coupling_integration.visual_matrix_execution import (
build_sdxl_visual_model_links,
)
from .cold_path_measurement import measure_cold_path
from .cold_path_priming import (
build_cold_path_primers,
execute_cold_path_primers,
)
from .cold_path_results import summarize_cold_path
from .cold_path_workflow import build_sdxl_cold_path_workflow
def run_cold_path_attribution(
artifacts: IntegrationArtifacts,
*,
inventory: SdxlVisualInventory,
case: SdxlVisualCase,
comfy_root: Path,
readiness_timeout: float = 240.0,
prompt_timeout: float = 1200.0,
) -> Path:
"""Record exactly one cold request and one warmed exact reuse."""
links = build_sdxl_visual_model_links(comfy_root, inventory)
masks = ManagedSdxlVisualMasks(
input_root=comfy_root / "input",
run_id=artifacts.run_id,
)
system_stats: JsonObject = {}
server_cleanup = False
with links:
with masks:
workflow = build_sdxl_cold_path_workflow(
checkpoint_name=CHECKPOINT_SELECTION,
mask_names=masks.names(case.mask_profile),
case=case,
)
primers = build_cold_path_primers(
checkpoint_name=CHECKPOINT_SELECTION,
mask_names=masks.names(case.mask_profile),
case=case,
)
required_node_ids = frozenset(
{
*workflow.required_node_ids,
*(
node_id
for primer in primers
for node_id in primer.workflow.required_node_ids
),
}
)
with ManagedComfyServer(
comfy_root=comfy_root,
artifacts=artifacts,
required_node_ids=required_node_ids,
readiness_timeout=readiness_timeout,
launch_arguments=sdxl_visual_sampling_launch_arguments(),
) as running:
system_stats = running.system_stats
primer_timings = execute_cold_path_primers(
running.client,
primers=primers,
seed=SDXL_VISUAL_SAMPLING.seed,
prompt_timeout=prompt_timeout,
)
observation = measure_cold_path(
running.client,
workflow=workflow,
first_seed=SDXL_VISUAL_SAMPLING.seed,
run_id=artifacts.run_id,
prompt_timeout=prompt_timeout,
)
port = running.port
process = running.process
server_cleanup = not process.is_running and is_loopback_port_available(port)
artifacts.record_cleanup(
process_running=process.is_running,
port_available=is_loopback_port_available(port),
)
summary = summarize_cold_path(observation.cold, observation.warm)
result = artifacts.root / "sdxl-cold-path-attribution.json"
result.write_text(
json.dumps(
{
"status": "completed",
"observation": asdict(observation),
"primers": {
label: asdict(timing) for label, timing in primer_timings.items()
},
"summary": asdict(summary),
"system_stats": system_stats,
"cleanup": {
"server": server_cleanup,
"model_links": links.cleaned,
"masks": masks.cleaned,
},
},
indent=2,
sort_keys=True,
)
+ "\n",
encoding="utf-8",
)
return result
@@ -0,0 +1,101 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Build a seed- and capture-revisable SDXL cold-attribution workflow."""
from __future__ import annotations
import copy
from dataclasses import dataclass
from tools.comfy_api import JsonObject
from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase
from .sampling_workflow import add_sdxl_sampler
from .two_adapter_workflow import (
TwoAdapterPerformanceMode,
prepare_two_adapter_sampling,
)
_CAPTURE_NODE = "SimpleSyrupBenchmark.CaptureColdPathDiagnostics"
_READ_NODE = "SimpleSyrupBenchmark.ReadColdPathDiagnostics"
@dataclass(frozen=True, slots=True)
class BuiltSdxlColdPathWorkflow:
"""Expose one stable graph with explicit execution identity inputs."""
prompt: dict[str, JsonObject]
sampler_node_id: str
capture_node_id: str
terminal_node_id: str
@property
def required_node_ids(self) -> frozenset[str]:
"""Return every Comfy node required by the graph."""
return frozenset(str(node["class_type"]) for node in self.prompt.values())
def prompt_for_execution(
self,
*,
seed: int,
run_id: str,
) -> dict[str, JsonObject]:
"""Change only sampler seed and the paired capture identity."""
if isinstance(seed, bool) or not isinstance(seed, int) or seed < 0:
raise ValueError("Cold-path sampler seed must be a non-negative integer.")
if not isinstance(run_id, str) or not run_id:
raise ValueError("Cold-path execution run id must be non-empty.")
prompt = copy.deepcopy(self.prompt)
_inputs(prompt, self.sampler_node_id)["seed"] = seed
_inputs(prompt, self.capture_node_id)["run_id"] = run_id
_inputs(prompt, self.terminal_node_id)["run_id"] = run_id
return prompt
def build_sdxl_cold_path_workflow(
*,
checkpoint_name: str,
mask_names: tuple[str, str],
case: SdxlVisualCase,
) -> BuiltSdxlColdPathWorkflow:
"""Build one image-free two-regional-LoRA attribution graph."""
prepared = prepare_two_adapter_sampling(
checkpoint_name=checkpoint_name,
mask_names=mask_names,
case=case,
mode=TwoAdapterPerformanceMode.REGIONAL,
)
capture = prepared.graph.add(
_CAPTURE_NODE,
model=prepared.model,
run_id="cold-path-template",
)
sampler = add_sdxl_sampler(prepared, model=[capture, 0])
terminal = prepared.graph.add(
_READ_NODE,
latent=[sampler, 0],
run_id="cold-path-template",
)
return BuiltSdxlColdPathWorkflow(
prompt=prepared.graph.prompt,
sampler_node_id=sampler,
capture_node_id=capture,
terminal_node_id=terminal,
)
def _inputs(prompt: dict[str, JsonObject], node_id: str) -> dict[str, object]:
"""Return one validated mutable API node input mapping."""
node = prompt.get(node_id)
if not isinstance(node, dict):
raise ValueError(f"Cold-path workflow is missing node {node_id!r}.")
inputs = node.get("inputs")
if not isinstance(inputs, dict):
raise ValueError(f"Cold-path workflow node {node_id!r} has invalid inputs.")
return inputs
@@ -67,7 +67,7 @@ def build_two_adapter_performance_workflow(
) -> BuiltTwoAdapterPerformanceWorkflow:
"""Build one matched 30-step sampler without decode or artifact work."""
prepared = _prepare_two_adapter_sampling(
prepared = prepare_two_adapter_sampling(
checkpoint_name=checkpoint_name,
mask_names=mask_names,
case=case,
@@ -98,7 +98,7 @@ def build_two_adapter_steady_state_workflow(
) -> BuiltSdxlSteadyStateWorkflow:
"""Build one cache-realistic graph with no MODEL-cloning timing probe."""
prepared = _prepare_two_adapter_sampling(
prepared = prepare_two_adapter_sampling(
checkpoint_name=checkpoint_name,
mask_names=mask_names,
case=case,
@@ -116,7 +116,7 @@ def build_two_adapter_steady_state_workflow(
)
def _prepare_two_adapter_sampling(
def prepare_two_adapter_sampling(
*,
checkpoint_name: str,
mask_names: tuple[str, str],