Files
Artificial-Sweetener-Simple…/simple_syrup/runtime/regional_lora/linear_execution.py
T

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()