refactor node setup

This commit is contained in:
mozman
2023-12-18 07:00:55 +01:00
parent 489e27e419
commit a13a4c5a7f
3 changed files with 18 additions and 14 deletions
+9 -1
View File
@@ -1,3 +1,11 @@
from .apply_sdxl_style import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
# Copyright (c) 2023, Manfred Moitzi
# License: MIT License
from __future__ import annotations
from . import apply_sdxl_style
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
NODE_CLASS_MAPPINGS: dict[str, type] = {}
NODE_DISPLAY_NAME_MAPPINGS: dict[str, str] = {}
apply_sdxl_style.setup_nodes(NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS)
+6 -8
View File
@@ -5,6 +5,7 @@ import pathlib
import dataclasses
import json
__all__ = ["setup_nodes"]
BYPASS = "bypass"
BYPASS_PROMPT = "{prompt}"
@@ -118,8 +119,6 @@ class ApplyStyle:
return output_positive, output_negative
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_DEFINITIONS = [
("sdxl_styles_sai.json", "Apply SDXL Style SAI"),
("sdxl_styles_twri.json", "Apply SDXL Style TWRI"),
@@ -131,17 +130,16 @@ NODE_DEFINITIONS = [
]
def _setup_classes():
def setup_nodes(
class_mapping: dict[str, type], display_name_mapping: dict[str, str]
) -> None:
cwd = pathlib.Path(__file__).parent
for file_name, display_name in NODE_DEFINITIONS:
templates = load_style_templates(cwd / "styles" / file_name)
class_name = display_name.replace(" ", "")
# create classes dynamically:
NODE_CLASS_MAPPINGS[class_name] = type(
class_mapping[class_name] = type(
class_name, (ApplyStyle,), {"TEMPLATES": templates}
)
NODE_DISPLAY_NAME_MAPPINGS[class_name] = display_name
_setup_classes()
display_name_mapping[class_name] = display_name
+3 -5
View File
@@ -3,11 +3,8 @@
import pytest
import re
from .apply_sdxl_style import (
NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS,
NODE_DEFINITIONS,
)
from . import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
from .apply_sdxl_style import NODE_DEFINITIONS
def test_all_classes_created():
@@ -25,6 +22,7 @@ def test_classes_have_templates():
for cls in NODE_CLASS_MAPPINGS.values():
assert len(cls.TEMPLATES) > 1, "no styles loaded - correct filename?"
def test_templates_have_bypass_style():
"""The style "bypass" will be added automatically to every style file."""
for cls in NODE_CLASS_MAPPINGS.values():