refactor class definition

This commit is contained in:
mozman
2023-12-17 09:45:33 +01:00
parent 355654d00a
commit bf02220e88
2 changed files with 42 additions and 44 deletions
+5
View File
@@ -0,0 +1,5 @@
{
"cSpell.words": [
"twri"
]
}
+37 -44
View File
@@ -28,7 +28,6 @@ class Template:
BYPASS_TEMPLATE = Template(BYPASS, BYPASS_PROMPT)
TEMPLATE_FILES: dict[str, dict[str, Template]] = dict()
def load_json_data(style_path: pathlib.Path) -> list:
@@ -55,15 +54,15 @@ def load_style_templates(style_path: pathlib.Path) -> dict[str, Template]:
class ApplyStyle:
FILE: str = "styles.json"
TEMPLATES: dict[str, Template] = {}
FILE = ""
def __str__(self):
return f"{self.__class__}, filename={self.FILE}, #styles={len(self.TEMPLATES)}"
@classmethod
def INPUT_TYPES(self):
cwd = pathlib.Path(__file__).parent
templates = load_style_templates(cwd / self.FILE)
templates[BYPASS] = BYPASS_TEMPLATE
TEMPLATE_FILES[self.FILE] = templates
style_names = sorted(templates.keys())
style_names = sorted(self.TEMPLATES.keys())
style_names.remove(BYPASS)
style_names.insert(0, BYPASS)
@@ -102,61 +101,55 @@ class ApplyStyle:
apply_negative_style: bool,
log_prompt: bool,
):
templates = TEMPLATE_FILES.get(self.FILE, {})
template = templates.get(style, BYPASS_TEMPLATE)
template = self.TEMPLATES.get(style, BYPASS_TEMPLATE)
output_positive, output_negative = template.apply(
positive_prompt, negative_prompt, apply_negative_style
)
if log_prompt:
print("-" * 79)
line = "-" * 79
print(line)
print(f"Style: {style} from '{self.FILE}' file")
print(f"Input Positive: {positive_prompt}")
print(f"Input Negative: {negative_prompt}")
print(f"Output Positive: {output_positive}")
print(f"Output Negative: {output_negative}")
print("-" * 79)
print(line)
return output_positive, output_negative
class ApplySDXLStyleSAI(ApplyStyle):
FILE = "sdxl_styles_sai.json"
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"),
("sdxl_styles_mre.json", "Apply SDXL Style MRE"),
("clipdrop_styles.json", "Apply ClipDrop Style"),
("fooocus.json", "Apply Fooocus Style"),
("art_styles_expansion.json", "Apply Art Style Expansion"),
]
class ApplySDXLStyleTWRI(ApplyStyle):
FILE = "sdxl_styles_twri.json"
def _setup_classes():
cwd = pathlib.Path(__file__).parent
for file_name, display_name in NODE_DEFINITIONS:
templates = load_style_templates(cwd / file_name)
class_name = display_name.replace(" ", "")
# create classes dynamically:
NODE_CLASS_MAPPINGS[class_name] = type(
class_name, (ApplyStyle,), {"TEMPLATES": templates, "FILE": file_name}
)
NODE_DISPLAY_NAME_MAPPINGS[class_name] = display_name
class ApplySDXLStyleMRE(ApplyStyle):
FILE = "sdxl_styles_mre.json"
_setup_classes()
class ApplyFooocusStyle(ApplyStyle):
FILE = "fooocus.json"
def print_loaded_classes():
for name, cls in NODE_CLASS_MAPPINGS.items():
print(f"{name}: {cls()}")
class ApplyClipDropStyle(ApplyStyle):
FILE = "clipdrop_styles.json"
class ApplyArtStyleExpansion(ApplyStyle):
FILE = "art_styles_expansion.json"
NODE_CLASS_MAPPINGS = {
"ApplySDXLStyleSAI": ApplySDXLStyleSAI,
"ApplySDXLStyleTWRI": ApplySDXLStyleTWRI,
"ApplySDXLStyleMRE": ApplySDXLStyleMRE,
"ApplyClipDropStyle": ApplyClipDropStyle,
"ApplyFooocusStyle": ApplyFooocusStyle,
"ApplyArtStyleExpansion": ApplyArtStyleExpansion,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ApplySDXLStyleSAI": "Apply SDXL Style SAI",
"ApplySDXLStyleTWRI": "Apply SDXL Style TWRI",
"ApplySDXLStyleMRE": "Apply SDXL Style MRE",
"ApplyClipDropStyle": "Apply ClipDrop Style",
"ApplyFooocusStyle": "Apply Fooocus Style",
"ApplyArtStyleExpansion": "Apply Art Style Expansion",
}
if __name__ == "__main__":
print_loaded_classes()