diff --git a/py/metadata_collector/metadata_processor.py b/py/metadata_collector/metadata_processor.py index 8b9a541f..e42b614d 100644 --- a/py/metadata_collector/metadata_processor.py +++ b/py/metadata_collector/metadata_processor.py @@ -7,7 +7,7 @@ from .constants import IMAGES 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, OVERWRITE -from .node_extractors import NODE_EXTRACTORS +from .node_extractors import NODE_EXTRACTORS, _get_variable_name logger = logging.getLogger(__name__) @@ -356,8 +356,18 @@ class MetadataProcessor: # Handle pipe nodes like FromBasicPipe by following the pipeline next_input_name = "basic_pipe" else: - # Dead end - no model input to follow - return None + # No direct model input to follow. GetNode (KJNodes) links + # are virtual and absent from the API prompt — resolve the + # variable reference back to the SetNode's MODEL source + # before giving up. + resolved_id = MetadataProcessor._resolve_model_variable_source( + metadata, prompt, current_node_id + ) + if resolved_id is None: + return None + current_node_id = resolved_id + depth += 1 + continue # Get connected node input_val = inputs[next_input_name] @@ -370,6 +380,57 @@ class MetadataProcessor: return None + @staticmethod + def _resolve_model_variable_source(metadata, prompt, node_id): + """Resolve a GetNode-style virtual reference to its MODEL source node. + + SetNodeExtractor records ``model_variable`` entries (variable name → + source node id) when a SetNode carrying a MODEL link executes, and + GetNodeExtractor records the variable name each GetNode reads. When + both are available, the virtual Set/Get link can be followed just like + a real connection. + """ + if not prompt or not getattr(prompt, "original_prompt", None): + return None + if node_id not in prompt.original_prompt: + return None + + # Variable name read by this node: prefer the runtime record, + # fall back to the prompt inputs. + variable_name = None + prompt_info = metadata.get(PROMPTS, {}).get(node_id) + if isinstance(prompt_info, dict): + variable_name = prompt_info.get("variable_name") + if not variable_name: + node_inputs = prompt.original_prompt[node_id].get("inputs", {}) + variable_name = _get_variable_name(node_inputs) + if not variable_name: + return None + + candidates = [ + info + for info in metadata.get(MODELS, {}).values() + if isinstance(info, dict) + and info.get("type") == "model_variable" + and info.get("variable_name") == variable_name + and info.get("source_node_id") + ] + if not candidates: + return None + + # KJNodes semantics: when several SetNodes share a variable name, + # the last one to execute wins. + execution_order = metadata.get("execution_order") or [] + + def _order(info): + try: + return execution_order.index(info.get("node_id")) + except ValueError: + return -1 + + best = max(candidates, key=_order) + return best.get("source_node_id") + @staticmethod def find_primary_checkpoint(metadata, downstream_id=None, primary_sampler_id=None): """ diff --git a/py/metadata_collector/metadata_registry.py b/py/metadata_collector/metadata_registry.py index 15f4d8e4..5b216cb5 100644 --- a/py/metadata_collector/metadata_registry.py +++ b/py/metadata_collector/metadata_registry.py @@ -1,8 +1,47 @@ +import os +import re import time from typing import Any from nodes import NODE_CLASS_MAPPINGS # pyright: ignore[reportMissingImports, reportAttributeAccessIssue] from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor -from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE +from .constants import METADATA_CATEGORIES, IMAGES, MODELS, OVERWRITE + + +# Input fields that carry a model filename in loader-style nodes. Mirrors +# GenericNodeExtractor._MODEL_NAME_FIELDS, plus the TensorRT engine extension. +_MODEL_NAME_FIELDS = ("ckpt_name", "unet_name", "model_path", "model_name", "gguf_name") +_MODEL_FILE_EXTENSIONS = ( + ".ckpt", ".pt", ".pt2", ".bin", ".pth", + ".safetensors", ".pkl", ".sft", ".gguf", ".engine", +) + + +def _checkpoint_name_candidates(inputs): + """Derive the checkpoint names a loader node's current inputs could record. + + Used to keep stale ``node_cache`` entries (from before the user switched + models) out of prompts they no longer belong to. Returns an empty set when + the inputs carry no recognizable model field. + """ + candidates = set() + for field in _MODEL_NAME_FIELDS: + value = inputs.get(field) + if not isinstance(value, str) or not value.strip(): + continue + name = value.strip() + if not name.lower().endswith(_MODEL_FILE_EXTENSIONS): + continue + candidates.add(name) + base = os.path.splitext(os.path.basename(name))[0] + candidates.add(base) + # TensorRTLoaderExtractor derivation: drop the "_$profile" part and + # any trailing save counter (e.g. "_00001_"). + derived = base + if "_$" in derived: + derived = derived[: derived.index("_$")] + derived = re.sub(r"_\d+_?$", "", derived) + candidates.add(derived) + return candidates class MetadataRegistry: @@ -155,9 +194,56 @@ class MetadataRegistry: 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][ - node_id - ] + cached_entry = cached_data[category][node_id] + if category == MODELS: + cached_entry = self._validated_checkpoint_fill( + cached_entry, node_data + ) + if cached_entry is None: + continue + metadata[category][node_id] = cached_entry + + @staticmethod + def _validated_checkpoint_fill(cached_entry, node_data): + """Guard a MODELS cache fill against stale checkpoint names. + + A loader that has not executed since the user switched models still + holds the previous model in its cache entry. When the node's current + inputs name a different file, the cached name is replaced with the one + derived from the current inputs; when no current name can be derived, + the stale entry is dropped entirely. + + The cache is trusted as-is when: + * the entry is not a checkpoint record, + * the cached name carries no model-file extension (e.g. names set + from runtime attachments rather than a widget value), or + * the node inputs carry no recognizable model field (unknown loader). + """ + if not isinstance(cached_entry, dict): + return cached_entry + if cached_entry.get("type") != "checkpoint": + return cached_entry + cached_name = cached_entry.get("name") + if not cached_name or not cached_name.lower().endswith(_MODEL_FILE_EXTENSIONS): + return cached_entry + + inputs = node_data.get("inputs", {}) + candidates = _checkpoint_name_candidates(inputs) + if not candidates: + return cached_entry + if cached_name in candidates: + return cached_entry + + # Stale — rebuild from the current inputs when possible. + for field in _MODEL_NAME_FIELDS: + value = inputs.get(field) + if isinstance(value, str) and value.strip().lower().endswith( + _MODEL_FILE_EXTENSIONS + ): + rebuilt = dict(cached_entry) + rebuilt["name"] = value.strip() + return rebuilt + return None def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None): """Record information about a node's execution""" diff --git a/py/metadata_collector/node_extractors.py b/py/metadata_collector/node_extractors.py index 2eeb31bd..457ef2b1 100644 --- a/py/metadata_collector/node_extractors.py +++ b/py/metadata_collector/node_extractors.py @@ -584,6 +584,20 @@ class ConditioningCombineExtractor(NodeMetadataExtractor): ) +def _get_model_link_source(metadata, node_id): + """Return the upstream node id feeding this node's MODEL input link.""" + prompt = metadata.get("current_prompt") + original_prompt = getattr(prompt, "original_prompt", None) + if not original_prompt or node_id not in original_prompt: + return None + node_inputs = original_prompt[node_id].get("inputs", {}) + for key in ("MODEL", "model"): + link = node_inputs.get(key) + if isinstance(link, list) and link: + return str(link[0]) + return None + + class SetNodeExtractor(NodeMetadataExtractor): @staticmethod def extract(node_id, inputs, outputs, metadata): @@ -591,6 +605,20 @@ class SetNodeExtractor(NodeMetadataExtractor): return variable_name = _get_node_variable_name(metadata, node_id, inputs) + + # Record MODEL passthrough as a variable reference so trace_model_path + # can follow the matching GetNode back to the real loader — Set/Get + # links are virtual and absent from the API prompt. + if variable_name: + source_node_id = _get_model_link_source(metadata, node_id) + if source_node_id is not None: + metadata.setdefault(MODELS, {})[node_id] = { + "type": "model_variable", + "variable_name": variable_name, + "source_node_id": source_node_id, + "node_id": node_id, + } + conditioning = inputs.get("CONDITIONING") if conditioning is None: conditioning = inputs.get("conditioning") diff --git a/tests/metadata_collector/test_model_variable_trace.py b/tests/metadata_collector/test_model_variable_trace.py new file mode 100644 index 00000000..84b89d6e --- /dev/null +++ b/tests/metadata_collector/test_model_variable_trace.py @@ -0,0 +1,188 @@ +"""Tests for MODEL variable tracing (KJNodes Set/Get) and stale checkpoint +cache-fill validation in the metadata collector.""" + +import types +from types import SimpleNamespace + +from py.metadata_collector import metadata_processor +from py.metadata_collector.constants import MODELS +from py.metadata_collector.metadata_processor import MetadataProcessor +from py.metadata_collector.metadata_registry import ( + MetadataRegistry, + _checkpoint_name_candidates, +) + + +def _ksampler_inputs(): + return { + "seed": 123, + "steps": 8, + "cfg": 1.0, + "sampler_name": "er_sde", + "scheduler": "simple", + "denoise": 1.0, + "latent_image": {"samples": types.SimpleNamespace(shape=(1, 4, 16, 16))}, + } + + +def test_primary_checkpoint_follows_kj_set_get_model_chain( + metadata_registry, monkeypatch +): + """Sampler <- GetNode <- (virtual) <- SetNode <- UNETLoader must resolve to + the UNETLoader even though the Set/Get links do not exist in the API prompt.""" + monkeypatch.setattr(metadata_processor, "standalone_mode", False) + + unet_name = "Krea 2/base model/krea2_turbo_int8_convrot.safetensors" + prompt_graph = { + "54": { + "class_type": "UNETLoader", + "inputs": {"unet_name": unet_name, "weight_dtype": "default"}, + }, + "32": { + "class_type": "SetNode", + "inputs": {"MODEL": ["54", 0], "name": "MODEL"}, + }, + "57": {"class_type": "GetNode", "inputs": {"name": "MODEL"}}, + "9": { + "class_type": "KSampler", + "inputs": {**_ksampler_inputs(), "model": ["57", 0]}, + }, + } + prompt = SimpleNamespace(original_prompt=prompt_graph) + + metadata_registry.start_collection("prompt-set-get-model") + metadata_registry.set_current_prompt(prompt) + metadata_registry.record_node_execution( + "54", "UNETLoader", {"unet_name": unet_name, "weight_dtype": "default"}, None + ) + metadata_registry.record_node_execution( + "32", "SetNode", {"MODEL": object(), "name": "MODEL"}, None + ) + metadata_registry.record_node_execution("57", "GetNode", {"name": "MODEL"}, None) + metadata_registry.record_node_execution("9", "KSampler", _ksampler_inputs(), None) + + metadata = metadata_registry.get_metadata("prompt-set-get-model") + + # The SetNode recorded a variable reference, not a checkpoint entry. + set_entry = metadata[MODELS]["32"] + assert set_entry["type"] == "model_variable" + assert set_entry["variable_name"] == "MODEL" + assert set_entry["source_node_id"] == "54" + + params = MetadataProcessor.extract_generation_params(metadata) + assert params["checkpoint"] == unet_name + + +def test_stale_checkpoint_cache_fill_rebuilt_from_current_inputs( + metadata_registry, monkeypatch +): + """A loader cached with model A, then switched to model B but not + re-executed (ComfyUI served its output from cache), must not resurrect + model A's name into the new prompt's metadata.""" + monkeypatch.setattr(metadata_processor, "standalone_mode", False) + + # Run A: loader executes with the old model. + prompt_a = SimpleNamespace( + original_prompt={ + "54": { + "class_type": "UNETLoader", + "inputs": {"unet_name": "old/myKrea2.safetensors"}, + }, + } + ) + metadata_registry.start_collection("prompt-old") + metadata_registry.set_current_prompt(prompt_a) + metadata_registry.record_node_execution( + "54", "UNETLoader", {"unet_name": "old/myKrea2.safetensors"}, None + ) + metadata_registry.get_metadata("prompt-old") + + # Run B: widget switched to a new model; the loader itself does not + # execute (its output comes from ComfyUI's execution cache). + prompt_b = SimpleNamespace( + original_prompt={ + "54": { + "class_type": "UNETLoader", + "inputs": {"unet_name": "new/krea2_turbo_int8_convrot.safetensors"}, + }, + "9": {"class_type": "KSampler", "inputs": {"model": ["54", 0]}}, + } + ) + metadata_registry.start_collection("prompt-new") + metadata_registry.set_current_prompt(prompt_b) + metadata_registry.record_node_execution("9", "KSampler", _ksampler_inputs(), None) + + metadata = metadata_registry.get_metadata("prompt-new") + assert ( + metadata[MODELS]["54"]["name"] == "new/krea2_turbo_int8_convrot.safetensors" + ) + + params = MetadataProcessor.extract_generation_params(metadata) + assert params["checkpoint"] == "new/krea2_turbo_int8_convrot.safetensors" + + +def test_checkpoint_cache_fill_trusted_when_inputs_match(metadata_registry): + """Unchanged inputs: the cached entry is filled unchanged.""" + prompt = SimpleNamespace( + original_prompt={ + "54": { + "class_type": "UNETLoader", + "inputs": {"unet_name": "models/flux.safetensors"}, + }, + } + ) + metadata_registry.start_collection("prompt-one") + metadata_registry.set_current_prompt(prompt) + metadata_registry.record_node_execution( + "54", "UNETLoader", {"unet_name": "models/flux.safetensors"}, None + ) + metadata_registry.get_metadata("prompt-one") + + # Identical re-run: nothing executes, everything fills from cache. + metadata_registry.start_collection("prompt-two") + metadata_registry.set_current_prompt(prompt) + metadata = metadata_registry.get_metadata("prompt-two") + assert metadata[MODELS]["54"]["name"] == "models/flux.safetensors" + + +def test_validated_checkpoint_fill_trust_cases(): + """Direct checks of the trust/rebuild/drop rules.""" + node_data = {"inputs": {"unet_name": "a/model_b.safetensors"}} + + # Not a checkpoint entry — passed through. + assert MetadataRegistry._validated_checkpoint_fill( + {"type": "model_variable"}, node_data + ) == {"type": "model_variable"} + + # Extension-less cached name (e.g. set from runtime attachments) — trusted. + entry = {"type": "checkpoint", "name": "flux1-dev"} + assert MetadataRegistry._validated_checkpoint_fill(entry, node_data) is entry + + # Matching name — trusted unchanged. + entry = {"type": "checkpoint", "name": "a/model_b.safetensors"} + assert MetadataRegistry._validated_checkpoint_fill(entry, node_data) is entry + + # Stale name — rebuilt from current inputs. + entry = {"type": "checkpoint", "name": "a/model_a.safetensors"} + rebuilt = MetadataRegistry._validated_checkpoint_fill(entry, node_data) + assert rebuilt["name"] == "a/model_b.safetensors" + assert rebuilt["type"] == "checkpoint" + + # Unparseable model field (no recognizable value) — cache trusted as-is. + entry = {"type": "checkpoint", "name": "a/model_a.safetensors"} + assert ( + MetadataRegistry._validated_checkpoint_fill(entry, {"inputs": {"unet_name": 123}}) + is entry + ) + + +def test_checkpoint_name_candidates_derivations(): + candidates = _checkpoint_name_candidates( + {"unet_name": "dir/sub/model_$fp16_00001_.engine"} + ) + assert "dir/sub/model_$fp16_00001_.engine" in candidates + assert "model_$fp16_00001_" in candidates + assert "model" in candidates # TensorRT-style derivation + + assert _checkpoint_name_candidates({"unet_name": "not-a-model"}) == set() + assert _checkpoint_name_candidates({}) == set()