diff --git a/__init__.py b/__init__.py index 6f07f90a..a5b691bc 100644 --- a/__init__.py +++ b/__init__.py @@ -18,6 +18,7 @@ try: # pragma: no cover - import fallback for pytest collection from .py.nodes.lora_info import LoraInfoLM from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath from .py.nodes.create_hook_lora import CreateHookLoraLM + from .py.nodes.metadata_overwrite import MetadataOverwriteLM from .py.metadata_collector import init as init_metadata_collector except ( ImportError @@ -66,6 +67,9 @@ except ( CreateHookLoraLM = importlib.import_module( "py.nodes.create_hook_lora" ).CreateHookLoraLM + MetadataOverwriteLM = importlib.import_module( + "py.nodes.metadata_overwrite" + ).MetadataOverwriteLM init_metadata_collector = importlib.import_module("py.metadata_collector").init NODE_CLASS_MAPPINGS = { @@ -88,6 +92,7 @@ NODE_CLASS_MAPPINGS = { LoraInfoLM.NAME: LoraInfoLM, LoraSyntaxToPath.NAME: LoraSyntaxToPath, CreateHookLoraLM.NAME: CreateHookLoraLM, + MetadataOverwriteLM.NAME: MetadataOverwriteLM, } WEB_DIRECTORY = "./web/comfyui" diff --git a/py/metadata_collector/constants.py b/py/metadata_collector/constants.py index 85072c05..36a59c28 100644 --- a/py/metadata_collector/constants.py +++ b/py/metadata_collector/constants.py @@ -9,6 +9,14 @@ EMBEDDINGS = "embeddings" SIZE = "size" IMAGES = "images" IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes +OVERWRITE = "overwrite" # Manual metadata overwrite from MetadataOverwriteLM node + +# Field names that the MetadataOverwriteLM node and its extractor share +METADATA_OVERWRITE_FIELDS = ( + "prompt", "negative_prompt", "seed", "steps", "cfg_scale", + "sampler", "scheduler", "checkpoint", "loras", "size", + "clip_skip", "additional_data", +) # Complete list of categories to track -METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES] +METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES, OVERWRITE] diff --git a/py/metadata_collector/metadata_processor.py b/py/metadata_collector/metadata_processor.py index 0d883d71..b5616c43 100644 --- a/py/metadata_collector/metadata_processor.py +++ b/py/metadata_collector/metadata_processor.py @@ -6,7 +6,7 @@ from .constants import IMAGES # Check if running in standalone mode standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0" -from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER +from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER, OVERWRITE from .node_extractors import NODE_EXTRACTORS logger = logging.getLogger(__name__) @@ -524,7 +524,8 @@ class MetadataProcessor: "checkpoint": None, "loras": "", "size": None, - "clip_skip": None + "clip_skip": None, + "additional_data": "", } # Get the prompt object for node relationship tracing @@ -672,7 +673,14 @@ class MetadataProcessor: break if params["clip_skip"] is None: params["clip_skip"] = "1" - + + # ---- Apply manual metadata overwrites ---- + for overwrite_info in metadata.get(OVERWRITE, {}).values(): + overwrite_params = overwrite_info.get("parameters", {}) + for key, value in overwrite_params.items(): + if value: # truthy check — only overwrite when user provided a real value + params[key] = value + return params @staticmethod diff --git a/py/metadata_collector/metadata_registry.py b/py/metadata_collector/metadata_registry.py index ef754a06..e8caa4e2 100644 --- a/py/metadata_collector/metadata_registry.py +++ b/py/metadata_collector/metadata_registry.py @@ -1,7 +1,7 @@ import time from nodes import NODE_CLASS_MAPPINGS # type: ignore from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor -from .constants import METADATA_CATEGORIES, IMAGES +from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE class MetadataRegistry: @@ -133,8 +133,16 @@ class MetadataRegistry: if cache_key in self.node_cache: cached_data = self.node_cache[cache_key] + # Detect bypass (mode=4) / mute (mode=2) — these nodes + # were intentionally disabled and should not contribute + # overwrite values from a previous execution's cache. + node_mode = node_data.get("mode", 0) + node_is_disabled = node_mode in (2, 4) + # Apply cached metadata to the current metadata for category in self.metadata_categories: + if category == OVERWRITE and node_is_disabled: + continue if category in cached_data and node_id in cached_data[category]: if node_id not in metadata[category]: metadata[category][node_id] = cached_data[category][ diff --git a/py/metadata_collector/node_extractors.py b/py/metadata_collector/node_extractors.py index 27f67e54..d75e0d90 100644 --- a/py/metadata_collector/node_extractors.py +++ b/py/metadata_collector/node_extractors.py @@ -2,7 +2,7 @@ import json import os import re -from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER +from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE, METADATA_OVERWRITE_FIELDS def _store_checkpoint_metadata(metadata, node_id, model_name): @@ -1221,6 +1221,32 @@ class CR_ApplyControlNetStackExtractor(NodeMetadataExtractor): metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative +class MetadataOverwriteExtractor(NodeMetadataExtractor): + """Extract manually specified metadata from MetadataOverwriteLM node. + + Stores truthy input values under the OVERWRITE category so that + extract_generation_params can merge them over the inferred params. + """ + + @staticmethod + def extract(node_id, inputs, outputs, metadata): + if not inputs: + return + + overwrite_params = {} + for key in METADATA_OVERWRITE_FIELDS: + value = inputs.get(key) + if value: # truthy — only overwrite when user provided a real value + overwrite_params[key] = value + + if overwrite_params: + metadata.setdefault(OVERWRITE, {}) + metadata[OVERWRITE][node_id] = { + "parameters": overwrite_params, + "node_id": node_id, + } + + # Registry of node-specific extractors # Keys are node class names NODE_EXTRACTORS = { @@ -1288,5 +1314,7 @@ NODE_EXTRACTORS = { "CFGGuider": CFGGuiderExtractor, # Add CFGGuider # Image "VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor + # Metadata overwrite + "MetadataOverwriteLM": MetadataOverwriteExtractor, # Add other nodes as needed } diff --git a/py/nodes/metadata_overwrite.py b/py/nodes/metadata_overwrite.py new file mode 100644 index 00000000..f8476e55 --- /dev/null +++ b/py/nodes/metadata_overwrite.py @@ -0,0 +1,154 @@ +"""Metadata Overwrite node — allows users to manually specify generation parameters +that override the automatically collected/inferred metadata. + +All inputs have falsy defaults: only truthy (non-empty / non-zero) values +will overwrite the corresponding field in the final metadata. +""" + +from typing import Any + +from ..metadata_collector.constants import METADATA_OVERWRITE_FIELDS + + +class MetadataOverwriteLM: + NAME = "Metadata Overwrite (LoraManager)" + CATEGORY = "Lora Manager/utils" + DESCRIPTION = ( + "Manually specify generation parameters to override automatically collected " + "metadata. Only filled/connected inputs will take effect — empty defaults " + "are ignored." + ) + + @classmethod + def INPUT_TYPES(cls) -> dict[str, Any]: + return { + "optional": { + "prompt": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": "Positive prompt. Only overwrites when non-empty.", + }, + ), + "negative_prompt": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": "Negative prompt. Only overwrites when non-empty.", + }, + ), + "seed": ( + "INT", + { + "default": 0, + "min": 0, + "max": 0xFFFFFFFFFFFFFFFF, + "control_after_generate": False, + "tooltip": "Seed value. Only overwrites when > 0.", + }, + ), + "steps": ( + "INT", + { + "default": 0, + "min": 0, + "max": 10000, + "tooltip": "Number of steps. Only overwrites when > 0.", + }, + ), + "cfg_scale": ( + "FLOAT", + { + "default": 0.0, + "min": 0.0, + "max": 100.0, + "tooltip": "CFG scale. Only overwrites when > 0.", + }, + ), + "sampler": ( + "STRING", + { + "default": "", + "tooltip": "Sampler name. Only overwrites when non-empty.", + }, + ), + "scheduler": ( + "STRING", + { + "default": "", + "tooltip": "Scheduler name. Only overwrites when non-empty.", + }, + ), + "checkpoint": ( + "STRING", + { + "default": "", + "tooltip": "Checkpoint / model name. Only overwrites when non-empty.", + }, + ), + "loras": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": ( + "LoRA syntax, e.g. " + "or , " + "separated by spaces. Only overwrites when non-empty." + ), + }, + ), + "size": ( + "STRING", + { + "default": "", + "tooltip": ( + "Image size in WIDTHxHEIGHT format (e.g. 512x768). " + "Only overwrites when non-empty." + ), + }, + ), + "clip_skip": ( + "INT", + { + "default": 0, + "min": -24, + "max": 24, + "tooltip": "Clip skip. Only overwrites when non-zero.", + }, + ), + "additional_data": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": ( + "Additional data to embed in the image metadata. " + "Inserted between Clip skip and Model hash in the " + "A1111-compatible parameters string. " + 'Example: "Copyright": "Some license info"' + ), + }, + ), + }, + } + + RETURN_TYPES = ("METADATA",) + RETURN_NAMES = ("metadata",) + FUNCTION = "collect_metadata" + OUTPUT_NODE = True + + def collect_metadata(self, **kwargs: Any) -> tuple[dict[str, Any]]: + """Collect non-falsy input values into a metadata dict. + + Only values that are truthy (non-empty string, non-zero number) + are included — matching the overwrite logic in the metadata pipeline. + """ + result: dict[str, Any] = {} + for key in METADATA_OVERWRITE_FIELDS: + value = kwargs.get(key) + if value: + result[key] = value + return (result,) diff --git a/py/nodes/save_image.py b/py/nodes/save_image.py index b6620dc9..6d3ff2aa 100644 --- a/py/nodes/save_image.py +++ b/py/nodes/save_image.py @@ -471,6 +471,9 @@ class SaveImageLM: params.append(f"Clip skip: {abs(cs)}") except (ValueError, TypeError): pass + additional_data = metadata_dict.get("additional_data", "") + if additional_data: + params.append(additional_data) if ckpt_hash: params.append(f"Model hash: {ckpt_hash[:10].upper()}") if ckpt_display_name: diff --git a/tests/metadata_collector/test_metadata_collector.py b/tests/metadata_collector/test_metadata_collector.py index e7dc35c4..df3b7b6e 100644 --- a/tests/metadata_collector/test_metadata_collector.py +++ b/tests/metadata_collector/test_metadata_collector.py @@ -820,3 +820,220 @@ def test_lora_manager_checkpoint_and_unet_loaders_extract_models(metadata_regist "type": "checkpoint", "node_id": "unet_node", } + + +# --------------------------------------------------------------------------- +# MetadataOverwriteExtractor & overwrite merge tests +# --------------------------------------------------------------------------- + +from py.metadata_collector.constants import OVERWRITE, METADATA_OVERWRITE_FIELDS +from py.metadata_collector.node_extractors import MetadataOverwriteExtractor + + +def test_metadata_overwrite_extractor_stores_truthy_values(metadata_registry): + """Extractor should store truthy inputs under the OVERWRITE category.""" + metadata_registry.start_collection("prompt-ow") + metadata = metadata_registry.prompt_metadata["prompt-ow"] + + inputs = { + "prompt": "a beautiful landscape", + "negative_prompt": "", + "seed": 42, + "steps": 0, + "cfg_scale": 7.5, + "sampler": "", + "scheduler": "", + "checkpoint": "myModel.safetensors", + "loras": "", + "size": "1024x768", + "clip_skip": 0, + "additional_data": '{"Copyright": "CC0"}', + } + + MetadataOverwriteExtractor.extract("ow-1", inputs, None, metadata) + + assert OVERWRITE in metadata + assert "ow-1" in metadata[OVERWRITE] + params = metadata[OVERWRITE]["ow-1"]["parameters"] + + # Truthy values stored + assert params["prompt"] == "a beautiful landscape" + assert params["seed"] == 42 + assert params["cfg_scale"] == 7.5 + assert params["checkpoint"] == "myModel.safetensors" + assert params["loras"] == "" + assert params["size"] == "1024x768" + assert params["additional_data"] == '{"Copyright": "CC0"}' + + # Falsy values NOT stored + assert "negative_prompt" not in params + assert "steps" not in params + assert "sampler" not in params + assert "scheduler" not in params + assert "clip_skip" not in params + + metadata_registry.clear_metadata() + + +def test_metadata_overwrite_extractor_empty_inputs(metadata_registry): + """Extractor with all-falsy inputs should NOT create OVERWRITE category.""" + metadata_registry.start_collection("prompt-ow2") + metadata = metadata_registry.prompt_metadata["prompt-ow2"] + + inputs = {key: "" for key in METADATA_OVERWRITE_FIELDS} + inputs.update({"seed": 0, "steps": 0, "cfg_scale": 0.0, "clip_skip": 0}) + + MetadataOverwriteExtractor.extract("ow-2", inputs, None, metadata) + + # start_collection pre-creates empty dicts for all categories, + # but no node should have populated OVERWRITE with any data + assert not metadata[OVERWRITE] + + metadata_registry.clear_metadata() + + +def test_extract_generation_params_applies_overwrite(metadata_registry, populated_registry, monkeypatch): + """overwrite values should replace inferred params in extract_generation_params.""" + import py.metadata_collector.metadata_processor as mp + + monkeypatch.setattr(mp, "standalone_mode", False) + + metadata = populated_registry["metadata"] + registry_obj = populated_registry["registry"] + + # Simulate the MetadataOverwriteLM node having been executed with overwrite values + registry_obj.start_collection("promptA") + # Re-populate with the same data (start_collection resets) + registry_obj.set_current_prompt(populated_registry["prompt"]) + metadata2 = registry_obj.prompt_metadata["promptA"] + + # Inject overwrite data into metadata + metadata2[OVERWRITE] = { + "ow-1": { + "parameters": { + "seed": 777, + "additional_data": '{"AuthorURL": "https://civitai.com/user/foo"}', + }, + "node_id": "ow-1", + } + } + # Copy other categories from original populated metadata + for cat in ("models", "prompts", "sampling", "loras", "size", "images"): + if cat in metadata: + metadata2[cat] = metadata[cat] + metadata2["execution_order"] = metadata["execution_order"] + + params = MetadataProcessor.extract_generation_params(metadata2, id="vae") + + # Overwritten values + assert params["seed"] == 777 + assert params["additional_data"] == '{"AuthorURL": "https://civitai.com/user/foo"}' + + # Inferred values still present (not overwritten) + assert params["prompt"] == "A castle on a hill" + assert params["cfg_scale"] == 7.5 + assert params["checkpoint"] == "model.safetensors" + + registry_obj.clear_metadata() + + +def test_extract_generation_params_overwrite_falsy_skipped(metadata_registry, populated_registry, monkeypatch): + """Overwrite entries with falsy values should NOT replace inferred params.""" + import py.metadata_collector.metadata_processor as mp + + monkeypatch.setattr(mp, "standalone_mode", False) + + metadata = populated_registry["metadata"] + registry_obj = populated_registry["registry"] + + registry_obj.start_collection("promptA") + registry_obj.set_current_prompt(populated_registry["prompt"]) + metadata2 = registry_obj.prompt_metadata["promptA"] + + # Inject overwrite with falsy values + metadata2[OVERWRITE] = { + "ow-1": { + "parameters": { + "seed": 0, + "steps": 0, + "cfg_scale": 0.0, + "prompt": "", + "clip_skip": 0, + }, + "node_id": "ow-1", + } + } + for cat in ("models", "prompts", "sampling", "loras", "size", "images"): + if cat in metadata: + metadata2[cat] = metadata[cat] + metadata2["execution_order"] = metadata["execution_order"] + + params = MetadataProcessor.extract_generation_params(metadata2, id="vae") + + # Falsy overwrites should NOT have replaced inferred values + assert params["prompt"] == "A castle on a hill" + assert params["cfg_scale"] == 7.5 + + registry_obj.clear_metadata() + + +def test_fill_missing_metadata_skips_overwrite_for_bypassed_node(metadata_registry): + """Bypassed (mode=4) node should not have OVERWRITE filled from cache.""" + metadata_registry.start_collection("prompt-bypass") + + # Simulate a previous execution that cached overwrite data + metadata_registry.record_node_execution( + "ow-1", + "MetadataOverwriteLM", + {"seed": 99, "prompt": "test", "steps": 0, "cfg_scale": 0.0, + "negative_prompt": "", "sampler": "", "scheduler": "", "checkpoint": "", + "loras": "", "size": "", "clip_skip": 0, "additional_data": ""}, + None, + ) + + # Now start a new prompt where the node is bypassed (mode=4) + metadata_registry.start_collection("prompt-bypass-2") + original_prompt = { + "ow-1": {"class_type": "MetadataOverwriteLM", "inputs": {}, "mode": 4}, + } + metadata_registry.set_current_prompt( + SimpleNamespace(original_prompt=original_prompt) + ) + + metadata = metadata_registry.get_metadata("prompt-bypass-2") + + # The overwrite data should NOT be present (node was bypassed, not + # a cache hit — it should not inherit previous execution's overwrite) + assert "ow-1" not in metadata.get(OVERWRITE, {}) + + metadata_registry.clear_metadata() + + +def test_fill_missing_metadata_fills_overwrite_for_muted_node(metadata_registry): + """Muted (mode=2) node should also not have OVERWRITE filled from cache.""" + metadata_registry.start_collection("prompt-mute") + + # Simulate a previous execution that cached overwrite data + metadata_registry.record_node_execution( + "ow-1", + "MetadataOverwriteLM", + {"seed": 88, "prompt": "test2", "steps": 0, "cfg_scale": 0.0, + "negative_prompt": "", "sampler": "", "scheduler": "", "checkpoint": "", + "loras": "", "size": "", "clip_skip": 0, "additional_data": ""}, + None, + ) + + # Start a new prompt where the node is muted (mode=2) + metadata_registry.start_collection("prompt-mute-2") + original_prompt = { + "ow-1": {"class_type": "MetadataOverwriteLM", "inputs": {}, "mode": 2}, + } + metadata_registry.set_current_prompt( + SimpleNamespace(original_prompt=original_prompt) + ) + + metadata = metadata_registry.get_metadata("prompt-mute-2") + + assert "ow-1" not in metadata.get(OVERWRITE, {}) + + metadata_registry.clear_metadata()