add style as en provider

This commit is contained in:
Vladimir Mandic
2023-10-07 08:07:42 -04:00
parent 109b1d6907
commit 8e22fbbdb6
4 changed files with 42 additions and 4 deletions
+4 -2
View File
@@ -17,6 +17,8 @@ def register_extra_network(extra_network):
def register_default_extra_networks():
from modules.extra_networks_hypernet import ExtraNetworkHypernet
register_extra_network(ExtraNetworkHypernet())
from modules.ui_extra_networks_styles import ExtraNetworkStyles
register_extra_network(ExtraNetworkStyles())
class ExtraNetworkParams:
@@ -70,7 +72,7 @@ def activate(p, extra_network_data):
try:
extra_network.activate(p, extra_network_args)
except Exception as e:
errors.display(e, f"activating extra network {extra_network_name} with arguments {extra_network_args}")
errors.display(e, f"activating extra network: name={extra_network_name} args:{extra_network_args}")
for extra_network_name, extra_network in extra_network_registry.items():
args = extra_network_data.get(extra_network_name, None)
@@ -79,7 +81,7 @@ def activate(p, extra_network_data):
try:
extra_network.activate(p, [])
except Exception as e:
errors.display(e, f"activating extra network {extra_network_name}")
errors.display(e, f"activating extra network: name={extra_network_name}")
def deactivate(p, extra_network_data):
+31 -1
View File
@@ -1,7 +1,7 @@
import os
import html
import json
from modules import shared, ui_extra_networks
from modules import shared, script_callbacks, extra_networks, ui_extra_networks, styles
class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage):
@@ -92,3 +92,33 @@ class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage):
def allowed_directories_for_previews(self):
return [v for v in [shared.opts.styles_dir] if v is not None] + ['html']
class ExtraNetworkStyles(extra_networks.ExtraNetwork):
def __init__(self):
super().__init__('style')
self.indexes = {}
def activate(self, p, params_list):
for param in params_list:
if len(param.items) > 0:
style = None
search = param.items[0]
# style = shared.prompt_styles.find_style(param.items[0])
match = [s for s in shared.prompt_styles.styles.values() if s.name == search]
if len(match) > 0:
style = match[0]
else:
match = [s for s in shared.prompt_styles.styles.values() if s.name.startswith(search)]
if len(match) > 0:
i = self.indexes.get(search, 0)
self.indexes[search] = (i + 1) % len(match)
style = match[self.indexes[search]]
if style is not None:
p.styles.append(style.name)
p.prompts = [styles.merge_prompts(style.prompt, prompt) for prompt in p.prompts]
p.negative_prompts = [styles.merge_prompts(style.negative_prompt, prompt) for prompt in p.negative_prompts]
def deactivate(self, p):
pass