203 lines
8.3 KiB
Python
203 lines
8.3 KiB
Python
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
|
# Copyright (C) 2026 Artificial Sweetener and contributors
|
|
# SPDX-License-Identifier: AGPL-3.0-or-later
|
|
|
|
"""Execute generic regional ordinary-LoRA deltas around one original Linear call."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import math
|
|
from collections.abc import Callable
|
|
|
|
import torch
|
|
|
|
from ...domain.regional_activation_geometry import RegionalActivationLayout
|
|
from .delta_execution import REGIONAL_LORA_DELTA_EXECUTOR, RegionalLoraDeltaExecutor
|
|
from .host_linear_parameters import RegionalHostLinearParameterProvider
|
|
from .linear_execution_plan import RegionalLinearExecutionPlan
|
|
from .operation_mask_resolution import RegionalOperationMaskBatch
|
|
from .partitioned_linear_execution import (
|
|
REGIONAL_PARTITIONED_LINEAR_EXECUTOR,
|
|
RegionalPartitionedLinearExecutor,
|
|
)
|
|
|
|
|
|
class RegionalLinearExecutor:
|
|
"""Call the original once and add exact active adapter groups in order."""
|
|
|
|
def __init__(
|
|
self,
|
|
delta_executor: RegionalLoraDeltaExecutor = REGIONAL_LORA_DELTA_EXECUTOR,
|
|
partitioned_executor: RegionalPartitionedLinearExecutor = (
|
|
REGIONAL_PARTITIONED_LINEAR_EXECUTOR
|
|
),
|
|
) -> None:
|
|
"""Retain the existing model-neutral projection and accumulation owner."""
|
|
|
|
if not isinstance(delta_executor, RegionalLoraDeltaExecutor):
|
|
raise TypeError("Regional Linear execution requires a delta executor.")
|
|
if not isinstance(partitioned_executor, RegionalPartitionedLinearExecutor):
|
|
raise TypeError(
|
|
"Regional Linear execution requires a partitioned executor."
|
|
)
|
|
self._delta_executor = delta_executor
|
|
self._partitioned_executor = partitioned_executor
|
|
|
|
def execute(
|
|
self,
|
|
original: Callable[..., object],
|
|
inputs: torch.Tensor,
|
|
*args: object,
|
|
plan: RegionalLinearExecutionPlan,
|
|
masks: RegionalOperationMaskBatch,
|
|
schedule_strengths: tuple[float, ...],
|
|
parameter_provider: RegionalHostLinearParameterProvider | None = None,
|
|
**kwargs: object,
|
|
) -> torch.Tensor:
|
|
"""Execute one original operation and ordered regional low-rank additions."""
|
|
|
|
self._validate_call(inputs, plan, masks, schedule_strengths)
|
|
invocation = masks.linear_invocations.resolve(
|
|
plan,
|
|
mask_multipliers=masks.multipliers,
|
|
composition_indices=masks.composition_indices,
|
|
schedule_strengths=schedule_strengths,
|
|
leading_shape=tuple(int(size) for size in inputs.shape[:-1]),
|
|
)
|
|
if not args and not kwargs:
|
|
partitioned = self._partitioned_executor.execute(
|
|
inputs,
|
|
plan=plan,
|
|
multipliers=invocation.multipliers,
|
|
parameter_provider=parameter_provider,
|
|
)
|
|
if partitioned is not None:
|
|
return partitioned
|
|
original_output = original(inputs, *args, **kwargs)
|
|
if not isinstance(original_output, torch.Tensor):
|
|
raise TypeError("Regional Linear original operation must return a tensor.")
|
|
expected_output_shape = (
|
|
*inputs.shape[:-1],
|
|
plan.groups[0].preparation.target.output_features,
|
|
)
|
|
if (
|
|
tuple(original_output.shape) != expected_output_shape
|
|
or original_output.device != inputs.device
|
|
or original_output.dtype != inputs.dtype
|
|
):
|
|
raise ValueError(
|
|
"Regional Linear original output must match its target shape and "
|
|
"input execution type."
|
|
)
|
|
active = tuple(
|
|
index
|
|
for index, multiplier in enumerate(invocation.multipliers)
|
|
if multiplier is not None
|
|
)
|
|
if not active:
|
|
return original_output
|
|
if len(active) == 1:
|
|
index = active[0]
|
|
multiplier = invocation.multipliers[index]
|
|
if multiplier is None:
|
|
raise AssertionError("Active Regional Linear multiplier disappeared.")
|
|
return self._delta_executor.add_masked_delta(
|
|
original_output,
|
|
inputs,
|
|
preparation=plan.groups[index].preparation,
|
|
multiplier=multiplier,
|
|
)
|
|
deltas: list[torch.Tensor | None] = [None] * len(plan.groups)
|
|
for indices, preparation in plan.rank_batches:
|
|
selected = tuple(
|
|
index for index in indices if invocation.multipliers[index] is not None
|
|
)
|
|
if not selected:
|
|
continue
|
|
if len(selected) == 1:
|
|
index = selected[0]
|
|
multiplier = invocation.multipliers[index]
|
|
if multiplier is None:
|
|
raise AssertionError(
|
|
"Active Regional Linear multiplier disappeared."
|
|
)
|
|
deltas[index] = self._delta_executor.masked_delta(
|
|
inputs,
|
|
preparation=plan.groups[index].preparation,
|
|
multiplier=multiplier,
|
|
)
|
|
continue
|
|
selected_preparation = plan.preparation_for(
|
|
selected_indices=selected,
|
|
complete_indices=indices,
|
|
complete=preparation,
|
|
)
|
|
compatible_deltas = self._delta_executor.compatible_deltas(
|
|
inputs,
|
|
preparation=selected_preparation,
|
|
multipliers=tuple(
|
|
multiplier
|
|
for index in selected
|
|
if (multiplier := invocation.multipliers[index]) is not None
|
|
),
|
|
)
|
|
for local_index, group_index in enumerate(selected):
|
|
deltas[group_index] = compatible_deltas[local_index]
|
|
return self._delta_executor.accumulate_ordered(
|
|
original_output,
|
|
tuple(delta for delta in deltas if delta is not None),
|
|
)
|
|
|
|
@staticmethod
|
|
def _validate_call(
|
|
inputs: object,
|
|
plan: object,
|
|
masks: object,
|
|
schedule_strengths: object,
|
|
) -> None:
|
|
"""Validate the explicit call contract before invoking the original."""
|
|
|
|
if not isinstance(inputs, torch.Tensor) or not inputs.is_floating_point():
|
|
raise TypeError("Regional Linear inputs must be a floating tensor.")
|
|
if not isinstance(plan, RegionalLinearExecutionPlan):
|
|
raise TypeError("Regional Linear execution requires a typed plan.")
|
|
if not isinstance(masks, RegionalOperationMaskBatch):
|
|
raise TypeError("Regional Linear execution requires operation masks.")
|
|
if masks.geometry.layout not in (
|
|
RegionalActivationLayout.FLATTENED_SPATIAL_TOKENS,
|
|
RegionalActivationLayout.CONSUMER_SPATIALIZED,
|
|
RegionalActivationLayout.BRANCH_TOKENS,
|
|
):
|
|
raise ValueError(
|
|
"Regional Linear execution requires spatial or branch-token geometry."
|
|
)
|
|
if tuple(inputs.shape) != masks.geometry.invocation_shape:
|
|
raise ValueError("Regional Linear input must match activation geometry.")
|
|
if (
|
|
masks.multipliers.device != inputs.device
|
|
or masks.multipliers.dtype != inputs.dtype
|
|
):
|
|
raise ValueError("Regional Linear masks must match input device and dtype.")
|
|
if not isinstance(schedule_strengths, tuple) or len(schedule_strengths) != len(
|
|
plan.uses
|
|
):
|
|
raise ValueError("Regional Linear schedule strengths must align to uses.")
|
|
if any(
|
|
isinstance(strength, bool)
|
|
or not isinstance(strength, int | float)
|
|
or not math.isfinite(float(strength))
|
|
for strength in schedule_strengths
|
|
):
|
|
raise TypeError("Regional Linear schedule strengths must be finite.")
|
|
if int(inputs.shape[-1]) != plan.groups[0].preparation.target.input_features:
|
|
raise ValueError("Regional Linear input feature count is invalid.")
|
|
if masks.composition_indices != tuple(
|
|
use.composition_index for use in plan.uses
|
|
):
|
|
raise ValueError(
|
|
"Regional Linear operation masks must align to target uses."
|
|
)
|
|
|
|
|
|
REGIONAL_LINEAR_EXECUTOR = RegionalLinearExecutor()
|