338 lines
13 KiB
Python
338 lines
13 KiB
Python
import ast
|
|
import csv
|
|
from functools import reduce
|
|
import logging
|
|
from pathlib import Path
|
|
import re
|
|
import textwrap
|
|
import time
|
|
import lark
|
|
from ruamel.yaml import YAML as _YAML
|
|
|
|
from ppp_logging import log
|
|
from ppp_classes import ONWARNING_CHOICES, PPPException, PPPInterrupt, PPPState
|
|
from ppp_utils import escape_single_quotes, format_output
|
|
|
|
|
|
def parse_prompt(
|
|
state: PPPState,
|
|
prompt_description: str,
|
|
prompt: str,
|
|
parser: lark.Lark,
|
|
raise_parsing_error: bool = False,
|
|
):
|
|
"""
|
|
Parses a prompt using the specified parser.
|
|
|
|
Args:
|
|
prompt_description (str): The description of the prompt.
|
|
prompt (str): The prompt to be parsed.
|
|
parser (lark.Lark): The parser to be used.
|
|
raise_parsing_error (bool): Whether to raise a parsing error.
|
|
|
|
Returns:
|
|
Tree: The parsed prompt.
|
|
"""
|
|
t1 = time.monotonic_ns()
|
|
parsed_prompt = None
|
|
try:
|
|
log(
|
|
state.logger,
|
|
state.options.debug_level,
|
|
logging.DEBUG,
|
|
f"Parsing {prompt_description}: '{escape_single_quotes(prompt)}'",
|
|
)
|
|
parsed_prompt = parser.parse(prompt)
|
|
# we store the contents so we can use them later even if the meta position is not valid anymore
|
|
if isinstance(parsed_prompt, lark.Tree):
|
|
for n in parsed_prompt.iter_subtrees():
|
|
if isinstance(n, lark.Tree):
|
|
if n.meta.empty:
|
|
n.meta.content = ""
|
|
else:
|
|
n.meta.content = prompt[n.meta.start_pos : n.meta.end_pos]
|
|
except lark.exceptions.UnexpectedInput:
|
|
if raise_parsing_error:
|
|
raise
|
|
log(
|
|
state.logger,
|
|
state.options.debug_level,
|
|
logging.ERROR,
|
|
f"Parsing failed on prompt!: {escape_single_quotes(prompt)}",
|
|
)
|
|
t2 = time.monotonic_ns()
|
|
log(
|
|
state.logger,
|
|
state.options.debug_level,
|
|
logging.DEBUG,
|
|
f"Parse {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds",
|
|
)
|
|
if parsed_prompt:
|
|
log(
|
|
state.logger,
|
|
state.options.debug_level,
|
|
logging.DEBUG,
|
|
"Tree:\n"
|
|
+ textwrap.indent(
|
|
re.sub(r"\n$", "", (parsed_prompt.pretty() if isinstance(parsed_prompt, lark.Tree) else parsed_prompt)),
|
|
" ",
|
|
),
|
|
formatted=False,
|
|
)
|
|
return parsed_prompt
|
|
|
|
|
|
def warn_or_stop(state: PPPState, is_negative: bool, message: str, e: Exception = None):
|
|
INVALID_CONTENT_STOP = "INVALID CONTENT! {0}\nBREAK "
|
|
if state.options.on_warning == ONWARNING_CHOICES.stop:
|
|
raise PPPInterrupt(
|
|
message,
|
|
INVALID_CONTENT_STOP.format(message) if not is_negative else "",
|
|
INVALID_CONTENT_STOP.format(message) if is_negative else "",
|
|
) from e
|
|
log(state.logger, state.options.debug_level, logging.WARNING, format_output(message))
|
|
|
|
|
|
def load_grammar() -> str:
|
|
# Process with lark (debug with https://www.lark-parser.org/ide/)
|
|
grammar_filename = Path(__file__).resolve().parent / "grammar.lark"
|
|
with open(grammar_filename, "r", encoding="utf-8") as file:
|
|
grammar_content = file.read()
|
|
return grammar_content
|
|
|
|
|
|
def preprocess_grammar(grammar_content: str, options: dict[str, bool], logger: logging.Logger, debug_level: int) -> str:
|
|
"""
|
|
Preprocesses the grammar content to handle conditional compilation directives.
|
|
|
|
Args:
|
|
grammar_content (str): The raw grammar content.
|
|
options (dict[str,bool]): Options for preprocessing.
|
|
logger (logging.Logger): The logger object.
|
|
debug_level (int): The debug level for logging.
|
|
|
|
Returns:
|
|
str: The preprocessed grammar content.
|
|
"""
|
|
lines = grammar_content.split("\n")
|
|
result_lines = []
|
|
skip_current_block = []
|
|
all_blocks_skipped = []
|
|
|
|
def eval_bool_expr(expr: str, constants: dict[str, bool]) -> bool:
|
|
"""
|
|
Evaluates a boolean expression using known constants.
|
|
Supports: and, or, not, parentheses, and named constants.
|
|
|
|
Args:
|
|
expr (str): The boolean expression to evaluate.
|
|
constants (dict[str, bool]): A dictionary of constant values.
|
|
Returns:
|
|
bool: The result of the evaluated expression.
|
|
"""
|
|
tree = ast.parse(expr, mode="eval")
|
|
|
|
def _eval(node) -> bool:
|
|
if isinstance(node, ast.Expression):
|
|
return _eval(node.body)
|
|
if isinstance(node, ast.BoolOp):
|
|
if isinstance(node.op, ast.And):
|
|
return all(_eval(v) for v in node.values)
|
|
if isinstance(node.op, ast.Or):
|
|
return any(_eval(v) for v in node.values)
|
|
if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not):
|
|
return not _eval(node.operand)
|
|
if isinstance(node, ast.Name):
|
|
return bool(constants[node.id]) # raises KeyError for unknown names
|
|
if isinstance(node, ast.Constant) and isinstance(node.value, bool):
|
|
return node.value
|
|
raise ValueError(f"Unsupported construct: {ast.dump(node)}")
|
|
|
|
return _eval(tree)
|
|
|
|
for line in lines:
|
|
stripped_line = line.strip()
|
|
if stripped_line.startswith("//#if"):
|
|
# Extract condition from the #if directive
|
|
conditions = stripped_line[5:].strip()
|
|
# Evaluate the conditions
|
|
skip_current_block.append(not eval_bool_expr(conditions, options))
|
|
all_blocks_skipped.append(skip_current_block[-1])
|
|
elif stripped_line.startswith("//#elif"):
|
|
if not skip_current_block:
|
|
log(logger, debug_level, logging.WARNING, "Unmatched //#elif directive found in grammar content.")
|
|
elif all_blocks_skipped[-1]:
|
|
# No earlier branch in this if/elif chain was taken, so evaluate this one.
|
|
# all_blocks_skipped[-1] tracks whether any branch has matched yet;
|
|
# skip_current_block[-1] tracks whether the current branch should be emitted.
|
|
conditions = stripped_line[7:].strip()
|
|
skip_current_block[-1] = not eval_bool_expr(conditions, options)
|
|
if not skip_current_block[-1]:
|
|
all_blocks_skipped[-1] = False # mark that a branch was taken
|
|
else:
|
|
# A previous branch already matched - skip all remaining elif/else branches.
|
|
skip_current_block[-1] = True
|
|
elif stripped_line.startswith("//#else"):
|
|
if not skip_current_block:
|
|
log(logger, debug_level, logging.WARNING, "Unmatched //#else directive found in grammar content.")
|
|
elif all_blocks_skipped[-1]:
|
|
skip_current_block[-1] = False
|
|
else:
|
|
skip_current_block[-1] = True
|
|
elif stripped_line.startswith("//#endif"):
|
|
if not skip_current_block:
|
|
log(logger, debug_level, logging.WARNING, "Unmatched //#endif directive found in grammar content.")
|
|
else:
|
|
skip_current_block.pop()
|
|
all_blocks_skipped.pop()
|
|
elif stripped_line.startswith("//#"):
|
|
log(
|
|
logger,
|
|
debug_level,
|
|
logging.WARNING,
|
|
f"Unrecognized directive found in grammar content: {stripped_line}",
|
|
)
|
|
elif not any(skip_current_block):
|
|
# Include the line if we're not skipping any current block
|
|
result_lines.append(stripped_line)
|
|
# Check for unclosed blocks at the end
|
|
if skip_current_block:
|
|
raise PPPInterrupt(
|
|
f"Found {len(skip_current_block)} unclosed conditional directive(s) at the end of the grammar file"
|
|
)
|
|
return "\n".join(result_lines)
|
|
|
|
|
|
def get_model_config_from_filename(filename: Path) -> object | None:
|
|
"""
|
|
Attempts to detect the model class from the given filename by inspecting the file header.
|
|
Currently only supports ComfyUI models in .safetensors format.
|
|
The path must be relative to a model folder.
|
|
"""
|
|
try:
|
|
import folder_paths # type: ignore
|
|
import comfy.utils # type: ignore
|
|
import comfy.model_detection as model_detection # type: ignore
|
|
except ImportError as e:
|
|
raise PPPException(f"Error detecting class from '{filename}': {e}") from e
|
|
import json
|
|
|
|
if not filename:
|
|
return None
|
|
|
|
path_keys = ["diffusion_models", "checkpoints", "unet"]
|
|
full_path: Path | None = None
|
|
if filename.is_absolute():
|
|
full_path = filename
|
|
else:
|
|
base_folders = []
|
|
for key in path_keys:
|
|
base_folders.extend(folder_paths.get_folder_paths(key))
|
|
for base in base_folders:
|
|
fname: Path = base / filename
|
|
if fname.exists():
|
|
full_path = fname
|
|
break
|
|
if not full_path or not full_path.suffix.lower() in (".safetensors", ".sft"):
|
|
return None
|
|
try:
|
|
header_bytes = comfy.utils.safetensors_header(full_path)
|
|
if header_bytes is None:
|
|
raise PPPException(f"Error detecting class from '{full_path}': no header")
|
|
header = json.loads(header_bytes)
|
|
|
|
# model_config_from_unet only inspects tensor shapes, not actual data.
|
|
# A lightweight proxy that exposes .shape and index access lets us avoid
|
|
# loading the full model into memory.
|
|
class _ShapeProxy:
|
|
def __init__(self, shape):
|
|
self.shape = shape
|
|
|
|
def __getitem__(self, i):
|
|
return self.shape[i]
|
|
|
|
mock_sd = {k: _ShapeProxy(v["shape"]) for k, v in header.items() if k != "__metadata__" and "shape" in v}
|
|
|
|
prefix = model_detection.unet_prefix_from_state_dict(mock_sd)
|
|
config = model_detection.model_config_from_unet(mock_sd, prefix, True)
|
|
if not config:
|
|
raise PPPException(f"Error detecting class from '{full_path}': no config found")
|
|
return config
|
|
except Exception as e: # pylint: disable=broad-except
|
|
raise PPPException(f"Error detecting class from file '{full_path}': {e}") from e
|
|
|
|
|
|
def get_model_class_from_filename(filename: Path) -> str:
|
|
config = get_model_config_from_filename(filename)
|
|
if not config:
|
|
return ""
|
|
c = config.__class__.__name__
|
|
if not c:
|
|
raise PPPException(f"Error detecting class from '{filename}': config has no class name {config}")
|
|
return c
|
|
|
|
|
|
def sanitize_wc_name(name: str) -> str:
|
|
# Remove invalid characters
|
|
return re.sub(r"[^a-zA-Z0-9-_]+", "", re.sub(r"_{2,}", "_", name.replace(" ", "_")))
|
|
|
|
|
|
def convert_a1111_styles_to_wildcard(inp: Path, out: Path):
|
|
"""
|
|
Converts styles from A1111 format to wildcard format in a YAML file.
|
|
Args:
|
|
inp (Path): The input CSV file path.
|
|
out (Path): The output YAML file path.
|
|
"""
|
|
wildcards = {}
|
|
with open(inp, "r", encoding="utf-8-sig") as f:
|
|
for reg in csv.reader(f):
|
|
name = reg[0].strip()
|
|
if name.lower() == "name":
|
|
continue
|
|
positive = reg[1].strip()
|
|
negative = reg[2].strip()
|
|
if negative:
|
|
positive = f"{positive}<ppp:stn>{negative}<ppp:/stn>"
|
|
name = sanitize_wc_name(name)
|
|
wildcards[name] = positive
|
|
if not wildcards:
|
|
raise RuntimeError(f"No styles found in {inp} to convert to wildcards.")
|
|
with open(out, "w", encoding="utf-8-sig") as f:
|
|
f.write(f"# Original names may contain characters that are replaced in the output.\n# Converted from {inp}\n")
|
|
_YAML().dump(wildcards, f)
|
|
|
|
|
|
def convert_sdnext_styles_to_wildcard(inp: Path, out: Path):
|
|
"""
|
|
Converts styles from SD.Next format to wildcard format in a YAML file.
|
|
Args:
|
|
inp (Path): The input folder path containing style json files or a single json file.
|
|
out (Path): The output YAML file path.
|
|
"""
|
|
wildcards = {}
|
|
files = inp.glob("*.json") if inp.is_dir() else [inp]
|
|
for file in files:
|
|
with open(file, "r", encoding="utf-8-sig") as f:
|
|
data = _YAML(typ="safe").load(f)
|
|
if not isinstance(data, list):
|
|
continue
|
|
wildcards[file] = {}
|
|
for style in data:
|
|
name = style.get("name", "").strip()
|
|
positive = style.get("prompt", "").strip()
|
|
negative = style.get("negative", "").strip()
|
|
# extra = style.get("extra", "").strip()
|
|
if negative:
|
|
positive = f"{positive}<ppp:stn>{negative}<ppp:/stn>"
|
|
name = sanitize_wc_name(name)
|
|
wildcards[file][name] = positive
|
|
if not reduce(lambda acc, d: acc or bool(d), wildcards.values(), False):
|
|
raise RuntimeError(f"No styles found in {inp} to convert to wildcards.")
|
|
with open(out, "w", encoding="utf-8-sig") as f:
|
|
f.write("# Original names may contain characters that are replaced in the output.\n")
|
|
for name, wcs in wildcards.items():
|
|
f.write(f"# Converted from {name}\n")
|
|
yaml_writer = _YAML()
|
|
yaml_writer.dump(wcs, f)
|