343 lines
12 KiB
Python
343 lines
12 KiB
Python
"""Pydantic models for the PPP configuration file structure (ppp_config.yaml)."""
|
|
|
|
from dataclasses import dataclass, field
|
|
from logging import Logger
|
|
import re
|
|
from enum import Enum
|
|
from typing import Any, Literal, Optional
|
|
from lark import Lark
|
|
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
|
|
|
from ppp_logging import DEBUG_LEVEL
|
|
from ppp_wildcards import PPPWildcards
|
|
from ppp_enmappings import PPPExtraNetworkMappings
|
|
from ppp_variables import VariableRepository
|
|
|
|
|
|
class SUPPORTED_APPS(Enum):
|
|
comfyui = "comfyui"
|
|
a1111 = "a1111"
|
|
forge = "forge"
|
|
forgeneo = "forgeneo"
|
|
reforge = "reforge"
|
|
sdnext = "sdnext"
|
|
tests = "tests" # for testing purposes only, not a real app
|
|
|
|
|
|
SUPPORTED_APPS_NAMES = {
|
|
SUPPORTED_APPS.comfyui: "ComfyUI",
|
|
SUPPORTED_APPS.sdnext: "SD.Next",
|
|
SUPPORTED_APPS.forge: "Forge Classic",
|
|
SUPPORTED_APPS.forgeneo: "Forge Neo",
|
|
SUPPORTED_APPS.reforge: "reForge",
|
|
SUPPORTED_APPS.a1111: "A1111 (or compatible)",
|
|
SUPPORTED_APPS.tests: "Tests",
|
|
}
|
|
|
|
|
|
class IFWILDCARDS_CHOICES(Enum):
|
|
ignore = "ignore"
|
|
remove = "remove"
|
|
warn = "warn"
|
|
stop = "stop"
|
|
|
|
|
|
class ONWARNING_CHOICES(Enum):
|
|
warn = "warn"
|
|
stop = "stop"
|
|
|
|
|
|
class RUN_MODE(Enum):
|
|
single = "single"
|
|
multiple = "multiple"
|
|
combinatorial = "combinatorial"
|
|
|
|
|
|
class DEFAULT_SAMPLER(Enum):
|
|
random = "random"
|
|
cyclical = "cyclical"
|
|
|
|
|
|
# ------------------- Host configuration -------------------
|
|
|
|
AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"]
|
|
SchedulingOption = Literal["ok", "before", "after", "first", "remove", "error"]
|
|
AlternationOption = Literal["ok", "first", "remove", "error"]
|
|
AndOption = Literal["ok", "eol", "comma", "remove", "error"]
|
|
BreakOption = Literal["ok", "eol", "comma", "remove", "error"]
|
|
|
|
|
|
class HostConfig(BaseModel):
|
|
"""Configuration for a specific host application."""
|
|
|
|
model_config = ConfigDict(populate_by_name=True, extra="forbid")
|
|
|
|
attention: AttentionOption = "ok"
|
|
scheduling: SchedulingOption = "ok"
|
|
alternation: AlternationOption = "ok"
|
|
and_: AndOption = Field("ok", alias="and")
|
|
break_: BreakOption = Field("ok", alias="break")
|
|
seed_bits: int = Field(
|
|
64, validator=lambda v: v if v in (8, 16, 32, 64) else ValueError("seed_bits must be one of 8, 16, 32, or 64")
|
|
)
|
|
|
|
|
|
# ------------------- Model detection -------------------
|
|
|
|
|
|
class ModelDetectConfig(BaseModel):
|
|
"""Detection configuration for a specific host when loading a model."""
|
|
|
|
model_config = ConfigDict(populate_by_name=True)
|
|
|
|
class_: Optional[list[str]] = Field(None, alias="class")
|
|
property: Optional[str] = None
|
|
|
|
@model_validator(mode="after")
|
|
def check_class_or_property(self) -> "ModelDetectConfig":
|
|
if self.class_ is None and self.property is None:
|
|
raise ValueError("Either 'class' or 'property' must be specified")
|
|
return self
|
|
|
|
|
|
# ------------------- Variant find_in_filename -------------------
|
|
|
|
|
|
class FindInFilenamePattern(BaseModel):
|
|
"""A regex pattern with optional flags used to identify a model variant in the filename."""
|
|
|
|
regex: str
|
|
flags: int = 0
|
|
|
|
@field_validator("flags", mode="before")
|
|
@classmethod
|
|
def parse_flags(cls, v: object) -> int:
|
|
if isinstance(v, int):
|
|
return v
|
|
if isinstance(v, list):
|
|
flag_value = 0
|
|
for flag in v:
|
|
if not isinstance(flag, str) or not hasattr(re, flag):
|
|
raise ValueError(f"Invalid regex flag '{flag}'")
|
|
flag_value |= getattr(re, flag)
|
|
return flag_value
|
|
raise ValueError(f"Expected int or list of flag-name strings, got {type(v).__name__}")
|
|
|
|
@model_validator(mode="after")
|
|
def validate_regex(self) -> "FindInFilenamePattern":
|
|
try:
|
|
re.compile(self.regex, self.flags)
|
|
except re.error as exc:
|
|
raise ValueError(f"Invalid regex pattern '{self.regex}': {exc}") from exc
|
|
return self
|
|
|
|
|
|
class VariantConfig(BaseModel):
|
|
"""Configuration for a specific model variant."""
|
|
|
|
find_in_filename: list[FindInFilenamePattern]
|
|
|
|
@field_validator("find_in_filename", mode="before")
|
|
@classmethod
|
|
def normalize_find_in_filename(cls, v: object) -> list:
|
|
"""Normalize str / dict / list input to always be a list of FindInFilenamePattern-compatible dicts."""
|
|
if isinstance(v, str):
|
|
return [{"regex": v, "flags": re.IGNORECASE}]
|
|
if isinstance(v, dict):
|
|
return [v]
|
|
if isinstance(v, list):
|
|
normalized = []
|
|
for item in v:
|
|
if isinstance(item, str):
|
|
normalized.append({"regex": item, "flags": re.IGNORECASE})
|
|
elif isinstance(item, dict):
|
|
normalized.append(item)
|
|
else:
|
|
raise ValueError(f"Expected str or dict in 'find_in_filename' list, got {type(item).__name__}")
|
|
return normalized
|
|
raise ValueError(f"Expected str, dict, or list for 'find_in_filename', got {type(v).__name__}")
|
|
|
|
|
|
# ------------------- Model configuration -------------------
|
|
|
|
|
|
class ModelConfig(BaseModel):
|
|
"""Configuration for a supported base model."""
|
|
|
|
detect: Optional[dict[str, Optional[ModelDetectConfig]]] = None
|
|
variants: Optional[dict[str, VariantConfig]] = None
|
|
|
|
@model_validator(mode="after")
|
|
def check_detect_or_variants(self) -> "ModelConfig":
|
|
if self.detect is None and self.variants is None:
|
|
raise ValueError("At least one of 'detect' or 'variants' must be specified")
|
|
return self
|
|
|
|
|
|
# ------------------- Top-level configuration -------------------
|
|
|
|
|
|
class PPPConfig(BaseModel):
|
|
"""Top-level PPP configuration structure matching ppp_config.yaml."""
|
|
|
|
hosts: Optional[dict[str, Optional[HostConfig]]] = None
|
|
models: Optional[dict[str, Optional[ModelConfig | None]]] = None
|
|
|
|
@model_validator(mode="after")
|
|
def check_hosts_or_models(self) -> "PPPConfig":
|
|
if self.hosts is None and self.models is None:
|
|
raise ValueError("At least one of 'hosts' or 'models' must be specified")
|
|
return self
|
|
|
|
|
|
# ------------------- State object -------------------
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class PPPStateOptions:
|
|
"""Options that can be set for prompt processing."""
|
|
|
|
debug_level: DEBUG_LEVEL = DEBUG_LEVEL.minimal
|
|
on_warning: ONWARNING_CHOICES = ONWARNING_CHOICES.warn
|
|
process_wildcards: bool = True
|
|
keep_choices_order: bool = True
|
|
choice_separator: str = ", "
|
|
if_wildcards: IFWILDCARDS_CHOICES = IFWILDCARDS_CHOICES.stop
|
|
stn_ignore_repeats: bool = True
|
|
stn_separator: str = ", "
|
|
cup_do_cleanup: bool = True # whether to do cleanup at all (if False, all other cleanup options are ignored)
|
|
cup_cleanup_variables: bool = True
|
|
cup_extra_spaces: bool = True
|
|
cup_empty_constructs: bool = True
|
|
cup_extra_separators: bool = True
|
|
cup_extra_separators2: bool = True
|
|
cup_extra_separators_include_eol: bool = False
|
|
cup_breaks: bool = False
|
|
cup_breaks_eol: bool = False
|
|
cup_ands: bool = False
|
|
cup_ands_eol: bool = False
|
|
cup_extranetwork_tags: bool = False
|
|
cup_merge_attention: bool = True
|
|
cup_remove_extranetwork_tags: bool = False
|
|
strict_operators: bool = True
|
|
results_file: str = "" # empty = disabled; supports %datetime%, %date%, %time%, %host% tokens
|
|
run_mode: RUN_MODE = RUN_MODE.single
|
|
results_limit: int = 100 # 0 = no limit
|
|
results_shuffle: bool = False
|
|
comb_random_fixed: bool = True # if True, the random sampler will be fixed across all DFS runs
|
|
default_sampler: DEFAULT_SAMPLER = DEFAULT_SAMPLER.random
|
|
|
|
def __post_init__(self):
|
|
if not self.cup_do_cleanup:
|
|
object.__setattr__(self, "cup_cleanup_variables", False)
|
|
object.__setattr__(self, "cup_extra_spaces", False)
|
|
object.__setattr__(self, "cup_empty_constructs", False)
|
|
object.__setattr__(self, "cup_extra_separators", False)
|
|
object.__setattr__(self, "cup_extra_separators2", False)
|
|
object.__setattr__(self, "cup_extra_separators_include_eol", False)
|
|
object.__setattr__(self, "cup_breaks", False)
|
|
object.__setattr__(self, "cup_breaks_eol", False)
|
|
object.__setattr__(self, "cup_ands", False)
|
|
object.__setattr__(self, "cup_ands_eol", False)
|
|
object.__setattr__(self, "cup_extranetwork_tags", False)
|
|
object.__setattr__(self, "cup_merge_attention", False)
|
|
object.__setattr__(self, "cup_remove_extranetwork_tags", False)
|
|
|
|
|
|
@dataclass
|
|
class PPPStateInputs:
|
|
"""Structured inputs for a single prompt processing call."""
|
|
|
|
seed: int = -1
|
|
pos_prompt: str = ""
|
|
neg_prompt: str = ""
|
|
jobinfo: Any = None
|
|
|
|
|
|
class CyclicalSamplerState:
|
|
"""Maintains the cycling position for '@' choice samplers across process_prompt calls."""
|
|
|
|
def __init__(self):
|
|
self.current_path: list[int] = []
|
|
self.last_trace: list[int] = []
|
|
self.last_prompt_pair: tuple[str, str] | None = None
|
|
|
|
def advance(self):
|
|
"""Advance to the next combination, cycling back to the start when all are exhausted."""
|
|
if not self.last_trace:
|
|
self.current_path = []
|
|
return
|
|
path = list(self.current_path)
|
|
while len(path) < len(self.last_trace):
|
|
path.append(0)
|
|
# Mixed-radix increment: least significant position is last.
|
|
for i in range(len(path) - 1, -1, -1):
|
|
path[i] += 1
|
|
if path[i] < self.last_trace[i]:
|
|
break
|
|
path[i] = 0
|
|
self.current_path = path
|
|
|
|
def reset(self):
|
|
"""Reset the cyclical state to the beginning."""
|
|
self.current_path = []
|
|
self.last_trace = []
|
|
self.last_prompt_pair = None
|
|
|
|
|
|
@dataclass
|
|
class PPPEnvInfo:
|
|
"""Environment and model information passed to PPP at construction time."""
|
|
|
|
app: SUPPORTED_APPS = SUPPORTED_APPS.a1111
|
|
ppp_config: str | dict | None = None
|
|
model_class: str = ""
|
|
model_filename: str = ""
|
|
property_base: Any = None
|
|
models_path: str = ""
|
|
_is_flags: dict[str, bool] = field(default_factory=dict, init=False, repr=False)
|
|
|
|
@property
|
|
def is_flags(self) -> dict[str, bool]:
|
|
"""Boolean model-detection flags keyed by model name (e.g. 'sdxl' -> True)."""
|
|
return self._is_flags
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class PPPState:
|
|
"""State object passed to various PPP components during prompt processing."""
|
|
|
|
logger: Logger
|
|
env_info: PPPEnvInfo = field(default_factory=PPPEnvInfo)
|
|
host_config: HostConfig = field(default_factory=HostConfig)
|
|
options: PPPStateOptions = field(default_factory=PPPStateOptions)
|
|
inputs: PPPStateInputs = field(default_factory=PPPStateInputs)
|
|
variables: VariableRepository = field(default_factory=VariableRepository)
|
|
wildcards_obj: PPPWildcards = field(default_factory=PPPWildcards)
|
|
extranetwork_mappings_obj: PPPExtraNetworkMappings = field(default_factory=PPPExtraNetworkMappings)
|
|
parsers: dict[str, Lark] = field(default_factory=dict)
|
|
cyclical_state: CyclicalSamplerState = field(default_factory=CyclicalSamplerState)
|
|
|
|
|
|
class PPPException(Exception):
|
|
"""
|
|
Custom exception to handle exceptions in the PromptPostProcessor.
|
|
"""
|
|
|
|
def __init__(self, message: str = "An error occurred during prompt processing."):
|
|
super().__init__(message)
|
|
self.message = message
|
|
|
|
|
|
class PPPInterrupt(PPPException):
|
|
"""
|
|
Custom exception to handle interruptions in the PromptPostProcessor.
|
|
This exception can be raised to stop the processing of prompts.
|
|
"""
|
|
|
|
def __init__(self, message: str = "Processing interrupted.", pos_prefix: str = "", neg_prefix: str = ""):
|
|
super().__init__(message)
|
|
self.message = message
|
|
self.pos_prefix = pos_prefix
|
|
self.neg_prefix = neg_prefix
|