70 lines
2.2 KiB
Python
70 lines
2.2 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
|
|
|
|
"""Domain model for per-segment conditioning batches."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from typing import Any, TypeAlias
|
|
|
|
Conditioning: TypeAlias = Any
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ConditioningBatch:
|
|
"""Store ordered conditioning entries for per-SEG selection."""
|
|
|
|
entries: tuple[Conditioning, ...]
|
|
|
|
def __post_init__(self) -> None:
|
|
"""Reject batches that cannot select a conditioning."""
|
|
|
|
if not self.entries:
|
|
raise ValueError(
|
|
"conditioning batch must contain at least one conditioning."
|
|
)
|
|
|
|
def select(self, index: int) -> Conditioning:
|
|
"""Return the entry for an index, reusing the last entry as fallback."""
|
|
|
|
if index < 0:
|
|
raise ValueError("conditioning batch index must be non-negative.")
|
|
return self.entries[min(index, len(self.entries) - 1)]
|
|
|
|
def append(self, conditioning: Conditioning) -> ConditioningBatch:
|
|
"""Return a new batch with one conditioning appended."""
|
|
|
|
return ConditioningBatch((*self.entries, conditioning))
|
|
|
|
|
|
def batch_conditioning(
|
|
values: tuple[Conditioning | ConditioningBatch, ...],
|
|
) -> ConditioningBatch:
|
|
"""Flatten conditioning values and batches into one ordered batch."""
|
|
|
|
if not values:
|
|
raise ValueError("Batch Region Conditioning requires one or more inputs.")
|
|
|
|
entries: list[Conditioning] = []
|
|
for value in values:
|
|
if isinstance(value, ConditioningBatch):
|
|
entries.extend(value.entries)
|
|
else:
|
|
entries.append(value)
|
|
return ConditioningBatch(tuple(entries))
|
|
|
|
|
|
def select_conditioning(
|
|
conditioning: Conditioning | ConditioningBatch,
|
|
index: int,
|
|
) -> Conditioning:
|
|
"""Select per-index conditioning or broadcast normal conditioning unchanged."""
|
|
|
|
if isinstance(conditioning, ConditioningBatch):
|
|
return conditioning.select(index)
|
|
if index < 0:
|
|
raise ValueError("conditioning batch index must be non-negative.")
|
|
return conditioning
|