refactor node setup
This commit is contained in:
+9
-1
@@ -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
@@ -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,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():
|
||||
|
||||
Reference in New Issue
Block a user