fix(metadata): record the actual checkpoint on Set/Get and stale-cache workflows

trace_model_path dead-ended on KJNodes GetNode virtual links and fell back
to the first entry in MODELS, which could be a stale checkpoint name
resurrected from the process-lifetime node_cache — images ended up with a
wrong or missing Hashes.model.

- SetNodeExtractor now records MODEL passthrough as a model_variable entry
  (variable name -> source node id), cached per node like other metadata
- trace_model_path resolves GetNode variable references through those
  entries instead of giving up (last SetNode wins, per KJNodes semantics)
- _fill_missing_metadata validates cached checkpoint names against the
  node's current widget values and rebuilds stale entries from the current
  inputs instead of resurrecting the previous model
This commit is contained in:
Will Miao
2026-10-09 20:44:39 +08:00
parent 1a6b0f78a3
commit 406096c0e1
4 changed files with 370 additions and 7 deletions
+64 -3
View File
@@ -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):
"""
+90 -4
View File
@@ -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"""
+28
View File
@@ -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")
@@ -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()