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")