mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-10-09 11:02:12 -03:00
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:
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user