mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-14 17:53:21 -03:00
Fix prompt tracking through conditioning transforms
This commit is contained in:
@@ -40,7 +40,7 @@ class GenericNodeExtractor(NodeMetadataExtractor):
|
|||||||
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
|
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
|
||||||
are checked for a model file name and stored as checkpoint metadata.
|
are checked for a model file name and stored as checkpoint metadata.
|
||||||
* ``CONDITIONING`` output: common text input fields are checked for
|
* ``CONDITIONING`` output: common text input fields are checked for
|
||||||
prompt text and stored as prompt metadata.
|
prompt text, and conditioning inputs are tracked through transforms.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Input field names that carry a model path in loader-style nodes.
|
# Input field names that carry a model path in loader-style nodes.
|
||||||
@@ -73,7 +73,7 @@ class GenericNodeExtractor(NodeMetadataExtractor):
|
|||||||
_store_checkpoint_metadata(metadata, node_id, name)
|
_store_checkpoint_metadata(metadata, node_id, name)
|
||||||
return
|
return
|
||||||
|
|
||||||
# — CONDITIONING encoder detection (CLIPTextEncode, Flux, custom) —
|
# — CONDITIONING encoder / transform detection —
|
||||||
if "CONDITIONING" in return_types or any("CONDITIONING" in str(t) for t in return_types):
|
if "CONDITIONING" in return_types or any("CONDITIONING" in str(t) for t in return_types):
|
||||||
text = None
|
text = None
|
||||||
for field in GenericNodeExtractor._TEXT_FIELDS:
|
for field in GenericNodeExtractor._TEXT_FIELDS:
|
||||||
@@ -81,12 +81,18 @@ class GenericNodeExtractor(NodeMetadataExtractor):
|
|||||||
if val and isinstance(val, str) and val.strip():
|
if val and isinstance(val, str) and val.strip():
|
||||||
text = val.strip()
|
text = val.strip()
|
||||||
break
|
break
|
||||||
if text:
|
|
||||||
prompt_data = metadata.setdefault(PROMPTS, {})
|
input_conditionings = [
|
||||||
prompt_data[node_id] = {
|
value
|
||||||
"text": text,
|
for input_name, value in inputs.items()
|
||||||
"node_id": node_id,
|
if input_name.startswith("conditioning") and value is not None
|
||||||
}
|
]
|
||||||
|
if text or input_conditionings:
|
||||||
|
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
||||||
|
if text:
|
||||||
|
prompt_metadata["text"] = text
|
||||||
|
if input_conditionings:
|
||||||
|
prompt_metadata["orig_conditionings"] = input_conditionings
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def update(node_id, outputs, metadata, return_types=None):
|
def update(node_id, outputs, metadata, return_types=None):
|
||||||
@@ -98,11 +104,22 @@ class GenericNodeExtractor(NodeMetadataExtractor):
|
|||||||
return
|
return
|
||||||
if node_id not in metadata.get(PROMPTS, {}):
|
if node_id not in metadata.get(PROMPTS, {}):
|
||||||
return
|
return
|
||||||
if outputs and isinstance(outputs, list) and len(outputs) > 0:
|
output_tuple = _first_output_tuple(outputs)
|
||||||
if isinstance(outputs[0], tuple) and len(outputs[0]) > 0:
|
if not output_tuple or len(output_tuple) < 1:
|
||||||
cond = outputs[0][0]
|
return
|
||||||
if cond is not None:
|
|
||||||
metadata[PROMPTS][node_id]["conditioning"] = cond
|
output_conditioning = output_tuple[0]
|
||||||
|
if output_conditioning is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
prompt_metadata = metadata[PROMPTS][node_id]
|
||||||
|
prompt_metadata["conditioning"] = output_conditioning
|
||||||
|
_record_conditioning_source(
|
||||||
|
metadata,
|
||||||
|
node_id,
|
||||||
|
output_conditioning,
|
||||||
|
prompt_metadata.get("orig_conditionings", []),
|
||||||
|
)
|
||||||
|
|
||||||
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -471,6 +471,172 @@ def test_conditioning_provenance_recovers_combined_controlnet_prompts(
|
|||||||
assert params["negative_prompt"] == "low quality"
|
assert params["negative_prompt"] == "low quality"
|
||||||
|
|
||||||
|
|
||||||
|
def test_conditioning_provenance_recovers_transformed_switched_prompts(
|
||||||
|
metadata_registry, monkeypatch
|
||||||
|
):
|
||||||
|
prompt_graph = {
|
||||||
|
"encode_pos": {
|
||||||
|
"class_type": "CLIPTextEncode",
|
||||||
|
"inputs": {"text": "expected positive", "clip": ["clip", 0]},
|
||||||
|
},
|
||||||
|
"encode_other_pos": {
|
||||||
|
"class_type": "CLIPTextEncode",
|
||||||
|
"inputs": {"text": "wrong positive", "clip": ["clip", 0]},
|
||||||
|
},
|
||||||
|
"encode_neg": {
|
||||||
|
"class_type": "CLIPTextEncode",
|
||||||
|
"inputs": {"text": "expected negative", "clip": ["clip", 0]},
|
||||||
|
},
|
||||||
|
"encode_other_neg": {
|
||||||
|
"class_type": "CLIPTextEncode",
|
||||||
|
"inputs": {"text": "wrong negative", "clip": ["clip", 0]},
|
||||||
|
},
|
||||||
|
"enhancer": {
|
||||||
|
"class_type": "KreaSeedVarianceEnhancer",
|
||||||
|
"inputs": {"conditioning": ["encode_pos", 0]},
|
||||||
|
},
|
||||||
|
"zero_out": {
|
||||||
|
"class_type": "ConditioningZeroOut",
|
||||||
|
"inputs": {"conditioning": ["encode_neg", 0]},
|
||||||
|
},
|
||||||
|
"positive_switch": {
|
||||||
|
"class_type": "ComfySwitchNode",
|
||||||
|
"inputs": {
|
||||||
|
"switch": True,
|
||||||
|
"on_false": ["encode_other_pos", 0],
|
||||||
|
"on_true": ["enhancer", 0],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"negative_switch": {
|
||||||
|
"class_type": "ComfySwitchNode",
|
||||||
|
"inputs": {
|
||||||
|
"switch": True,
|
||||||
|
"on_false": ["encode_other_neg", 0],
|
||||||
|
"on_true": ["zero_out", 0],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"sampler": {
|
||||||
|
"class_type": "ClownsharKSampler_Beta",
|
||||||
|
"inputs": {
|
||||||
|
"seed": 123,
|
||||||
|
"steps": 8,
|
||||||
|
"cfg": 1.0,
|
||||||
|
"sampler_name": "linear/euler",
|
||||||
|
"scheduler": "beta57",
|
||||||
|
"denoise": 1.0,
|
||||||
|
"positive": ["positive_switch", 0],
|
||||||
|
"negative": ["negative_switch", 0],
|
||||||
|
"latent_image": {
|
||||||
|
"samples": types.SimpleNamespace(shape=(1, 4, 16, 16))
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
prompt = SimpleNamespace(original_prompt=prompt_graph)
|
||||||
|
|
||||||
|
positive_conditioning = object()
|
||||||
|
other_positive_conditioning = object()
|
||||||
|
negative_conditioning = object()
|
||||||
|
other_negative_conditioning = object()
|
||||||
|
enhanced_conditioning = object()
|
||||||
|
zeroed_conditioning = object()
|
||||||
|
|
||||||
|
monkeypatch.setattr(metadata_processor, "standalone_mode", False)
|
||||||
|
|
||||||
|
metadata_registry.start_collection("prompt-transformed-switch")
|
||||||
|
metadata_registry.set_current_prompt(prompt)
|
||||||
|
|
||||||
|
for node_id, text, conditioning in (
|
||||||
|
("encode_pos", "expected positive", positive_conditioning),
|
||||||
|
("encode_other_pos", "wrong positive", other_positive_conditioning),
|
||||||
|
("encode_neg", "expected negative", negative_conditioning),
|
||||||
|
("encode_other_neg", "wrong negative", other_negative_conditioning),
|
||||||
|
):
|
||||||
|
metadata_registry.record_node_execution(
|
||||||
|
node_id, "CLIPTextEncode", {"text": text}, None
|
||||||
|
)
|
||||||
|
metadata_registry.update_node_execution(
|
||||||
|
node_id, "CLIPTextEncode", [(conditioning,)]
|
||||||
|
)
|
||||||
|
|
||||||
|
metadata_registry.record_node_execution(
|
||||||
|
"enhancer",
|
||||||
|
"KreaSeedVarianceEnhancer",
|
||||||
|
{"conditioning": positive_conditioning},
|
||||||
|
None,
|
||||||
|
return_types=("CONDITIONING", "STRING"),
|
||||||
|
)
|
||||||
|
metadata_registry.update_node_execution(
|
||||||
|
"enhancer",
|
||||||
|
"KreaSeedVarianceEnhancer",
|
||||||
|
[(enhanced_conditioning, "diagnostics")],
|
||||||
|
return_types=("CONDITIONING", "STRING"),
|
||||||
|
)
|
||||||
|
metadata_registry.record_node_execution(
|
||||||
|
"zero_out",
|
||||||
|
"ConditioningZeroOut",
|
||||||
|
{"conditioning": negative_conditioning},
|
||||||
|
None,
|
||||||
|
return_types=("CONDITIONING",),
|
||||||
|
)
|
||||||
|
metadata_registry.update_node_execution(
|
||||||
|
"zero_out",
|
||||||
|
"ConditioningZeroOut",
|
||||||
|
[(zeroed_conditioning,)],
|
||||||
|
return_types=("CONDITIONING",),
|
||||||
|
)
|
||||||
|
metadata_registry.record_node_execution(
|
||||||
|
"positive_switch",
|
||||||
|
"ComfySwitchNode",
|
||||||
|
{
|
||||||
|
"switch": True,
|
||||||
|
"on_false": other_positive_conditioning,
|
||||||
|
"on_true": enhanced_conditioning,
|
||||||
|
},
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
metadata_registry.update_node_execution(
|
||||||
|
"positive_switch", "ComfySwitchNode", [(enhanced_conditioning,)]
|
||||||
|
)
|
||||||
|
metadata_registry.record_node_execution(
|
||||||
|
"negative_switch",
|
||||||
|
"ComfySwitchNode",
|
||||||
|
{
|
||||||
|
"switch": True,
|
||||||
|
"on_false": other_negative_conditioning,
|
||||||
|
"on_true": zeroed_conditioning,
|
||||||
|
},
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
metadata_registry.update_node_execution(
|
||||||
|
"negative_switch", "ComfySwitchNode", [(zeroed_conditioning,)]
|
||||||
|
)
|
||||||
|
metadata_registry.record_node_execution(
|
||||||
|
"sampler",
|
||||||
|
"ClownsharKSampler_Beta",
|
||||||
|
{
|
||||||
|
"seed": 123,
|
||||||
|
"steps": 8,
|
||||||
|
"cfg": 1.0,
|
||||||
|
"sampler_name": "linear/euler",
|
||||||
|
"scheduler": "beta57",
|
||||||
|
"denoise": 1.0,
|
||||||
|
"positive": enhanced_conditioning,
|
||||||
|
"negative": zeroed_conditioning,
|
||||||
|
"latent_image": {
|
||||||
|
"samples": types.SimpleNamespace(shape=(1, 4, 16, 16))
|
||||||
|
},
|
||||||
|
},
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
|
metadata = metadata_registry.get_metadata("prompt-transformed-switch")
|
||||||
|
params = MetadataProcessor.extract_generation_params(metadata)
|
||||||
|
|
||||||
|
assert params["prompt"] == "expected positive"
|
||||||
|
assert params["negative_prompt"] == "expected negative"
|
||||||
|
|
||||||
|
|
||||||
def test_conditioning_provenance_recovers_kj_set_get_prompts(
|
def test_conditioning_provenance_recovers_kj_set_get_prompts(
|
||||||
metadata_registry, monkeypatch
|
metadata_registry, monkeypatch
|
||||||
):
|
):
|
||||||
|
|||||||
Reference in New Issue
Block a user