* 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:
Antonio Cordero Balcazar
2026-06-12 22:53:21 +02:00
parent 5e6460c099
commit 3a2571c008
5 changed files with 137 additions and 72 deletions
+73 -52
View File
@@ -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: