* Support for adding more input vars from outside the prompt.
* Added input variables `_input_prev_pos_prompt` and `_input_prev_neg_prompt` for hiresfix phase of A1111 compatible hosts.
This commit is contained in:
+73
-52
@@ -79,6 +79,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
try:
|
||||
# Support for SD.Next
|
||||
import installer # type: ignore
|
||||
|
||||
if hasattr(installer, "control_extensions"):
|
||||
if self.title() not in installer.control_extensions:
|
||||
installer.control_extensions.append(self.title()) # We add the extension to the whitelist.
|
||||
@@ -86,7 +87,12 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
# log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.WARNING, "Could not import control_extensions from installer, SD.Next support will not work.")
|
||||
pass
|
||||
except Exception as e: # pylint: disable=broad-except
|
||||
log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.ERROR, f"Error while adding to the SD.Next extension whitelist: {e}")
|
||||
log(
|
||||
self.ppp_logger,
|
||||
DEBUG_LEVEL.minimal,
|
||||
logging.ERROR,
|
||||
f"Error while adding to the SD.Next extension whitelist: {e}",
|
||||
)
|
||||
|
||||
def title(self):
|
||||
"""
|
||||
@@ -388,8 +394,8 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
else:
|
||||
calculated_seeds = seeds
|
||||
|
||||
# (prompt type, typeindex) -> (new positive prompt, new negative prompt)
|
||||
prompts_list: dict[tuple[str, int], tuple[str, str]] = {}
|
||||
# [index][indextype] -> (new positive prompt, new negative prompt)
|
||||
prompts_list: list[list[tuple[str, str]]] = []
|
||||
extra_params = {}
|
||||
|
||||
# adds prompts
|
||||
@@ -402,10 +408,10 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
rnh: list[str] = getattr(p, "all_hr_negative_prompts", None)
|
||||
hiresfix_exists = bool(rph) and bool(rnh)
|
||||
for i in range(len(calculated_seeds)):
|
||||
if regular_exists:
|
||||
prompts_list[(regular_type, i)] = None
|
||||
prompts_list.append([])
|
||||
prompts_list[i].append(None)
|
||||
if hiresfix_exists:
|
||||
prompts_list[(hiresfix_type, i)] = None
|
||||
prompts_list[i].append(None)
|
||||
|
||||
ppp.process_prompts_group_start()
|
||||
if input_combinatorial:
|
||||
@@ -429,7 +435,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
num_comb = len(comb_results)
|
||||
for i in range(len(rpr)): # pylint: disable=consider-using-enumerate
|
||||
posp, negp, _ = comb_results[i % num_comb]
|
||||
prompts_list[(regular_type, i)] = (posp, negp)
|
||||
prompts_list[i][0] = (posp, negp)
|
||||
extra_params["PPP combination"] = [str(1 + (i % num_comb)) for i in range(len(rpr))]
|
||||
if hiresfix_exists:
|
||||
hiresfix_equal = regular_exists and rph == rpr and rnh == rnr
|
||||
@@ -441,7 +447,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
"hiresfix prompts are the same as regular prompts, skipping combinatorial processing for hiresfix",
|
||||
)
|
||||
for i in range(len(rph)): # pylint: disable=consider-using-enumerate
|
||||
prompts_list[(hiresfix_type, i)] = prompts_list.get((regular_type, i))
|
||||
prompts_list[i][1] = prompts_list[i][0]
|
||||
else:
|
||||
log(
|
||||
self.ppp_logger,
|
||||
@@ -462,43 +468,54 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
num_comb_hr = len(comb_results_hr)
|
||||
for i in range(len(rph)): # pylint: disable=consider-using-enumerate
|
||||
posp, negp, _ = comb_results_hr[i % num_comb_hr]
|
||||
prompts_list[(hiresfix_type, i)] = (posp, negp)
|
||||
prompts_list[i][1] = (posp, negp)
|
||||
extra_params["PPP HR combination"] = [str(1 + (i % num_comb_hr)) for i in range(len(rph))]
|
||||
else:
|
||||
# processes prompts
|
||||
for prompttype, typeindex in prompts_list.keys():
|
||||
log(
|
||||
self.ppp_logger,
|
||||
self.ppp_debug_level,
|
||||
logging.INFO,
|
||||
f"processing prompts ({prompttype}[{typeindex+1}])",
|
||||
)
|
||||
key = (
|
||||
(hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex])
|
||||
if prompttype == regular_type
|
||||
else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex])
|
||||
)
|
||||
cached = self.lru_cache.get(key)
|
||||
if cached is None:
|
||||
hsh, seed, prompt, negative_prompt = key
|
||||
results = ppp.process_prompt(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
seed,
|
||||
jobinfo={
|
||||
"job_timestamp": shared.state.job_timestamp,
|
||||
"job": shared.state.job,
|
||||
"detail": f"{prompttype} prompt",
|
||||
},
|
||||
for index, grouplist in enumerate(prompts_list):
|
||||
for typeindex in range(len(grouplist)):
|
||||
typeprompt = [regular_type, hiresfix_type][typeindex]
|
||||
log(
|
||||
self.ppp_logger,
|
||||
self.ppp_debug_level,
|
||||
logging.INFO,
|
||||
f"processing prompts ({typeprompt}[{index+1}])",
|
||||
)
|
||||
posp, negp, _ = results[0]
|
||||
cached = (posp, negp)
|
||||
self.lru_cache.put(key, cached)
|
||||
# adds also the result so i2i doesn't process it unnecessarily
|
||||
self.lru_cache.put((hsh, seed, posp, negp), cached)
|
||||
else:
|
||||
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache")
|
||||
prompts_list[(prompttype, typeindex)] = cached
|
||||
key = (
|
||||
(hash_fullenv, calculated_seeds[index], rpr[index], rnr[index])
|
||||
if typeindex == 0
|
||||
else (hash_fullenv, calculated_seeds[index], rph[index], rnh[index])
|
||||
)
|
||||
cached = self.lru_cache.get(key)
|
||||
if cached is None:
|
||||
hsh, seed, prompt, negative_prompt = key
|
||||
if typeindex > 0:
|
||||
prev_prompts = prompts_list[index][typeindex - 1]
|
||||
input_vars = {
|
||||
"prev_pos_prompt": prev_prompts[0],
|
||||
"prev_neg_prompt": prev_prompts[1],
|
||||
}
|
||||
else:
|
||||
input_vars = None
|
||||
results = ppp.process_prompt(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
seed,
|
||||
jobinfo={
|
||||
"job_timestamp": shared.state.job_timestamp,
|
||||
"job": shared.state.job,
|
||||
"detail": f"{typeprompt} prompt",
|
||||
},
|
||||
input_vars=input_vars,
|
||||
)
|
||||
posp, negp, _ = results[0]
|
||||
cached = (posp, negp)
|
||||
self.lru_cache.put(key, cached)
|
||||
# adds also the result so i2i doesn't process it unnecessarily
|
||||
self.lru_cache.put((hsh, seed, posp, negp), cached)
|
||||
else:
|
||||
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache")
|
||||
prompts_list[index][typeindex] = cached
|
||||
ppp.process_prompts_group_end()
|
||||
|
||||
# updates the prompts
|
||||
@@ -506,17 +523,21 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None)
|
||||
regular_changes = False
|
||||
hiresfix_changes = False
|
||||
for (prompttype, typeindex), (posp, negp) in prompts_list.items():
|
||||
if prompttype == regular_type:
|
||||
if rpr[typeindex].strip() != posp.strip() or rnr[typeindex].strip() != negp.strip():
|
||||
regular_changes = True
|
||||
rpr[typeindex] = posp
|
||||
rnr[typeindex] = negp
|
||||
elif prompttype == hiresfix_type:
|
||||
if rph[typeindex].strip() != posp.strip() or rnh[typeindex].strip() != negp.strip():
|
||||
hiresfix_changes = True
|
||||
rph[typeindex] = posp
|
||||
rnh[typeindex] = negp
|
||||
for index, grouplist in enumerate(prompts_list):
|
||||
for typeindex, groupprompts in enumerate(grouplist):
|
||||
if groupprompts is None:
|
||||
continue
|
||||
(posp, negp) = groupprompts
|
||||
if typeindex == 0:
|
||||
if rpr[index].strip() != posp.strip() or rnr[index].strip() != negp.strip():
|
||||
regular_changes = True
|
||||
rpr[index] = posp
|
||||
rnr[index] = negp
|
||||
elif typeindex == 1:
|
||||
if rph[index].strip() != posp.strip() or rnh[index].strip() != negp.strip():
|
||||
hiresfix_changes = True
|
||||
rph[index] = posp
|
||||
rnh[index] = negp
|
||||
|
||||
# initialize extra generation parameters
|
||||
if add_prompts:
|
||||
|
||||
Reference in New Issue
Block a user