mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-21 13:01:27 -03:00
Compare commits
5 Commits
f34c02756d
...
f49b4ba4db
| Author | SHA1 | Date | |
|---|---|---|---|
| f49b4ba4db | |||
| 84e708328b | |||
| 125bed3f09 | |||
| 077e70169d | |||
| e6dc169a05 |
@@ -18,6 +18,7 @@ try: # pragma: no cover - import fallback for pytest collection
|
|||||||
from .py.nodes.lora_info import LoraInfoLM
|
from .py.nodes.lora_info import LoraInfoLM
|
||||||
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
|
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
|
||||||
from .py.nodes.create_hook_lora import CreateHookLoraLM
|
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
|
from .py.metadata_collector import init as init_metadata_collector
|
||||||
except (
|
except (
|
||||||
ImportError
|
ImportError
|
||||||
@@ -66,6 +67,9 @@ except (
|
|||||||
CreateHookLoraLM = importlib.import_module(
|
CreateHookLoraLM = importlib.import_module(
|
||||||
"py.nodes.create_hook_lora"
|
"py.nodes.create_hook_lora"
|
||||||
).CreateHookLoraLM
|
).CreateHookLoraLM
|
||||||
|
MetadataOverwriteLM = importlib.import_module(
|
||||||
|
"py.nodes.metadata_overwrite"
|
||||||
|
).MetadataOverwriteLM
|
||||||
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
@@ -88,6 +92,7 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
LoraInfoLM.NAME: LoraInfoLM,
|
LoraInfoLM.NAME: LoraInfoLM,
|
||||||
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||||
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||||
|
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
|
||||||
}
|
}
|
||||||
|
|
||||||
WEB_DIRECTORY = "./web/comfyui"
|
WEB_DIRECTORY = "./web/comfyui"
|
||||||
|
|||||||
@@ -9,6 +9,14 @@ EMBEDDINGS = "embeddings"
|
|||||||
SIZE = "size"
|
SIZE = "size"
|
||||||
IMAGES = "images"
|
IMAGES = "images"
|
||||||
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
|
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", "model", "loras", "size",
|
||||||
|
"clip_skip", "additional_data",
|
||||||
|
)
|
||||||
|
|
||||||
# Complete list of categories to track
|
# 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]
|
||||||
|
|||||||
@@ -83,7 +83,8 @@ class MetadataHook:
|
|||||||
|
|
||||||
# Record inputs before execution
|
# Record inputs before execution
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -114,7 +115,8 @@ class MetadataHook:
|
|||||||
|
|
||||||
# Record outputs after execution
|
# Record outputs after execution
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.update_node_execution(node_id, class_type, results)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -135,10 +137,13 @@ class MetadataHook:
|
|||||||
# Store the dynprompt reference for node lookups
|
# Store the dynprompt reference for node lookups
|
||||||
if hasattr(prompt, 'original_prompt'):
|
if hasattr(prompt, 'original_prompt'):
|
||||||
registry.set_current_prompt(prompt)
|
registry.set_current_prompt(prompt)
|
||||||
|
|
||||||
|
# Store extra_data for accessing full workflow node properties
|
||||||
|
registry.set_extra_data(extra_data)
|
||||||
|
|
||||||
# Execute the original function
|
# Execute the original function
|
||||||
return original_execute(*args, **kwargs)
|
return original_execute(*args, **kwargs)
|
||||||
|
|
||||||
# Replace the functions
|
# Replace the functions
|
||||||
execution._map_node_over_list = map_node_over_list_with_metadata
|
execution._map_node_over_list = map_node_over_list_with_metadata
|
||||||
execution.execute = execute_with_prompt_tracking
|
execution.execute = execute_with_prompt_tracking
|
||||||
@@ -163,7 +168,8 @@ class MetadataHook:
|
|||||||
class_type = obj.__class__.__name__
|
class_type = obj.__class__.__name__
|
||||||
node_id = unique_id
|
node_id = unique_id
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -180,7 +186,8 @@ class MetadataHook:
|
|||||||
class_type = obj.__class__.__name__
|
class_type = obj.__class__.__name__
|
||||||
node_id = unique_id
|
node_id = unique_id
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.update_node_execution(node_id, class_type, results)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -202,6 +209,9 @@ class MetadataHook:
|
|||||||
if hasattr(prompt, 'original_prompt'):
|
if hasattr(prompt, 'original_prompt'):
|
||||||
registry.set_current_prompt(prompt)
|
registry.set_current_prompt(prompt)
|
||||||
|
|
||||||
|
# Store extra_data for accessing full workflow node properties
|
||||||
|
registry.set_extra_data(extra_data)
|
||||||
|
|
||||||
# Execute the original function
|
# Execute the original function
|
||||||
return await original_execute(*args, **kwargs)
|
return await original_execute(*args, **kwargs)
|
||||||
|
|
||||||
|
|||||||
@@ -1,15 +1,68 @@
|
|||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
from .constants import IMAGES
|
from .constants import IMAGES
|
||||||
|
|
||||||
# Check if running in standalone mode
|
# 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"
|
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__)
|
||||||
|
|
||||||
|
# Keys that identify metadata hint marks stored in node.properties.lm_marker_role
|
||||||
|
_META_MARK_PREFIX = "meta_"
|
||||||
|
_MARK_PRIMARY_MODEL = "primary_model"
|
||||||
|
_MARK_PRIMARY_SAMPLER = "primary_sampler"
|
||||||
|
_MARK_POSITIVE_PROMPT = "positive_prompt"
|
||||||
|
_MARK_NEGATIVE_PROMPT = "negative_prompt"
|
||||||
|
|
||||||
class MetadataProcessor:
|
class MetadataProcessor:
|
||||||
"""Process and format collected metadata"""
|
"""Process and format collected metadata"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_user_marks(metadata):
|
||||||
|
"""Scan workflow nodes (from extra_data.extra_pnginfo.workflow) for user-assigned
|
||||||
|
metadata hint marks stored in node.properties.lm_marker_role.
|
||||||
|
|
||||||
|
Returns a dict mapping mark type keys to node IDs.
|
||||||
|
Example: {'primary_model': '42', 'primary_sampler': '17'}
|
||||||
|
"""
|
||||||
|
marks: dict[str, str] = {}
|
||||||
|
|
||||||
|
# Primary source: extra_data.extra_pnginfo.workflow.nodes (has full properties)
|
||||||
|
extra_data = metadata.get("extra_data")
|
||||||
|
if extra_data and isinstance(extra_data, dict):
|
||||||
|
extra_pnginfo = extra_data.get("extra_pnginfo", {})
|
||||||
|
if isinstance(extra_pnginfo, dict):
|
||||||
|
workflow = extra_pnginfo.get("workflow", {})
|
||||||
|
nodes = workflow.get("nodes", [])
|
||||||
|
for node in nodes:
|
||||||
|
node_id = str(node.get("id", ""))
|
||||||
|
role = node.get("properties", {}).get("lm_marker_role", "")
|
||||||
|
if role.startswith(_META_MARK_PREFIX):
|
||||||
|
mark_type = role[len(_META_MARK_PREFIX):]
|
||||||
|
if mark_type in marks:
|
||||||
|
logger.warning(
|
||||||
|
"Duplicate meta hint '%s': node %s (previous: %s), "
|
||||||
|
"last match wins",
|
||||||
|
mark_type, node_id, marks[mark_type],
|
||||||
|
)
|
||||||
|
marks[mark_type] = node_id
|
||||||
|
|
||||||
|
# Fallback: try prompt.original_prompt (API-only submissions may not have workflow)
|
||||||
|
if not marks:
|
||||||
|
prompt = metadata.get("current_prompt")
|
||||||
|
if prompt and getattr(prompt, "original_prompt", None):
|
||||||
|
for node_id, node_data in prompt.original_prompt.items():
|
||||||
|
role = node_data.get("properties", {}).get("lm_marker_role", "")
|
||||||
|
if role.startswith(_META_MARK_PREFIX):
|
||||||
|
mark_type = role[len(_META_MARK_PREFIX):]
|
||||||
|
marks[mark_type] = node_id
|
||||||
|
|
||||||
|
return marks
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def find_primary_sampler(metadata, downstream_id=None):
|
def find_primary_sampler(metadata, downstream_id=None):
|
||||||
"""
|
"""
|
||||||
@@ -471,20 +524,57 @@ class MetadataProcessor:
|
|||||||
"checkpoint": None,
|
"checkpoint": None,
|
||||||
"loras": "",
|
"loras": "",
|
||||||
"size": None,
|
"size": None,
|
||||||
"clip_skip": None
|
"clip_skip": None,
|
||||||
|
"additional_data": "",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Get the prompt object for node relationship tracing
|
# Get the prompt object for node relationship tracing
|
||||||
prompt = metadata.get("current_prompt")
|
prompt = metadata.get("current_prompt")
|
||||||
|
|
||||||
# Find the primary KSampler node
|
# ---- User marks: override heuristic inference with user-assigned hints ----
|
||||||
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
user_marks = MetadataProcessor._get_user_marks(metadata)
|
||||||
|
|
||||||
# Directly get checkpoint from metadata instead of tracing
|
# Find the primary KSampler node (user mark takes priority)
|
||||||
# Pass primary_sampler_id to avoid redundant calculation
|
primary_sampler_id = None
|
||||||
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
primary_sampler = None
|
||||||
if checkpoint:
|
if _MARK_PRIMARY_SAMPLER in user_marks:
|
||||||
params["checkpoint"] = checkpoint
|
marked_id = user_marks[_MARK_PRIMARY_SAMPLER]
|
||||||
|
sampler_data = metadata.get(SAMPLING, {}).get(marked_id)
|
||||||
|
if sampler_data and sampler_data.get(IS_SAMPLER):
|
||||||
|
primary_sampler_id = marked_id
|
||||||
|
primary_sampler = sampler_data
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"User-marked primary sampler %s has no runtime metadata, "
|
||||||
|
"falling back to heuristic",
|
||||||
|
marked_id,
|
||||||
|
)
|
||||||
|
if primary_sampler is None:
|
||||||
|
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
||||||
|
|
||||||
|
# Resolve checkpoint / model (user mark takes priority)
|
||||||
|
if _MARK_PRIMARY_MODEL in user_marks:
|
||||||
|
marked_id = user_marks[_MARK_PRIMARY_MODEL]
|
||||||
|
if marked_id in metadata.get(MODELS, {}):
|
||||||
|
params["checkpoint"] = metadata[MODELS][marked_id].get("name")
|
||||||
|
else:
|
||||||
|
extra_data = metadata.get("extra_data")
|
||||||
|
extra_pnginfo = extra_data.get("extra_pnginfo", {}) if extra_data and isinstance(extra_data, dict) else {}
|
||||||
|
workflow = extra_pnginfo.get("workflow", {}) if isinstance(extra_pnginfo, dict) else {}
|
||||||
|
node_type = "unknown"
|
||||||
|
for n in workflow.get("nodes", []):
|
||||||
|
if str(n.get("id", "")) == marked_id:
|
||||||
|
node_type = n.get("type", "unknown")
|
||||||
|
break
|
||||||
|
logger.warning(
|
||||||
|
"User-marked primary model %s (type=%s, registered=%s) has no runtime metadata, "
|
||||||
|
"falling back to heuristic",
|
||||||
|
marked_id, node_type, node_type in NODE_EXTRACTORS,
|
||||||
|
)
|
||||||
|
if params["checkpoint"] is None:
|
||||||
|
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
||||||
|
if checkpoint:
|
||||||
|
params["checkpoint"] = checkpoint
|
||||||
|
|
||||||
# Check if guidance parameter exists in any sampling node
|
# Check if guidance parameter exists in any sampling node
|
||||||
for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
|
for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
|
||||||
@@ -539,7 +629,22 @@ class MetadataProcessor:
|
|||||||
|
|
||||||
# For SamplerCustom, handle any additional parameters
|
# For SamplerCustom, handle any additional parameters
|
||||||
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
||||||
|
|
||||||
|
# ---- User marks: override prompts with explicitly tagged nodes ----
|
||||||
|
prompts_data = metadata.get(PROMPTS, {})
|
||||||
|
if _MARK_POSITIVE_PROMPT in user_marks:
|
||||||
|
pos_id = user_marks[_MARK_POSITIVE_PROMPT]
|
||||||
|
if pos_id in prompts_data:
|
||||||
|
prompt_text = prompts_data[pos_id].get("text") or prompts_data[pos_id].get("positive_text")
|
||||||
|
if prompt_text:
|
||||||
|
params["prompt"] = prompt_text
|
||||||
|
if _MARK_NEGATIVE_PROMPT in user_marks:
|
||||||
|
neg_id = user_marks[_MARK_NEGATIVE_PROMPT]
|
||||||
|
if neg_id in prompts_data:
|
||||||
|
prompt_text = prompts_data[neg_id].get("text") or prompts_data[neg_id].get("negative_text")
|
||||||
|
if prompt_text:
|
||||||
|
params["negative_prompt"] = prompt_text
|
||||||
|
|
||||||
# Size extraction is same for all sampler types
|
# Size extraction is same for all sampler types
|
||||||
# Check if the sampler itself has size information (from latent_image)
|
# Check if the sampler itself has size information (from latent_image)
|
||||||
if primary_sampler_id in metadata.get(SIZE, {}):
|
if primary_sampler_id in metadata.get(SIZE, {}):
|
||||||
@@ -568,7 +673,21 @@ class MetadataProcessor:
|
|||||||
break
|
break
|
||||||
if params["clip_skip"] is None:
|
if params["clip_skip"] is None:
|
||||||
params["clip_skip"] = "1"
|
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
|
||||||
|
|
||||||
|
# Bridge: the overwrite node exposes the field as "model" (more accurate),
|
||||||
|
# but the internal pipeline key remains "checkpoint" for backward compatibility
|
||||||
|
# with A1111 metadata format and downstream consumers.
|
||||||
|
if params.get("model"):
|
||||||
|
params["checkpoint"] = params["model"]
|
||||||
|
del params["model"]
|
||||||
|
|
||||||
return params
|
return params
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import time
|
import time
|
||||||
from nodes import NODE_CLASS_MAPPINGS # type: ignore
|
from nodes import NODE_CLASS_MAPPINGS # type: ignore
|
||||||
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
||||||
from .constants import METADATA_CATEGORIES, IMAGES
|
from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
|
||||||
|
|
||||||
|
|
||||||
class MetadataRegistry:
|
class MetadataRegistry:
|
||||||
@@ -61,6 +61,7 @@ class MetadataRegistry:
|
|||||||
{
|
{
|
||||||
"execution_order": [],
|
"execution_order": [],
|
||||||
"current_prompt": None, # Will store the prompt object
|
"current_prompt": None, # Will store the prompt object
|
||||||
|
"extra_data": None, # Will store the API extra_data for workflow metadata
|
||||||
"timestamp": time.time(),
|
"timestamp": time.time(),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@@ -75,6 +76,11 @@ class MetadataRegistry:
|
|||||||
# Store the prompt in the metadata for later relationship tracing
|
# Store the prompt in the metadata for later relationship tracing
|
||||||
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
||||||
|
|
||||||
|
def set_extra_data(self, extra_data):
|
||||||
|
"""Store the API extra_data (contains extra_pnginfo.workflow with node properties)"""
|
||||||
|
if self.current_prompt_id and self.current_prompt_id in self.prompt_metadata:
|
||||||
|
self.prompt_metadata[self.current_prompt_id]["extra_data"] = extra_data
|
||||||
|
|
||||||
def get_metadata(self, prompt_id=None):
|
def get_metadata(self, prompt_id=None):
|
||||||
"""Get collected metadata for a prompt"""
|
"""Get collected metadata for a prompt"""
|
||||||
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
||||||
@@ -122,20 +128,28 @@ class MetadataRegistry:
|
|||||||
cache_key = f"{node_id}:{class_type}"
|
cache_key = f"{node_id}:{class_type}"
|
||||||
|
|
||||||
# Check if this node type is relevant for metadata collection
|
# Check if this node type is relevant for metadata collection
|
||||||
if class_type in NODE_EXTRACTORS:
|
if class_type in NODE_EXTRACTORS or cache_key in self.node_cache:
|
||||||
# Check if we have cached metadata for this node
|
# Check if we have cached metadata for this node
|
||||||
if cache_key in self.node_cache:
|
if cache_key in self.node_cache:
|
||||||
cached_data = self.node_cache[cache_key]
|
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
|
# Apply cached metadata to the current metadata
|
||||||
for category in self.metadata_categories:
|
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 category in cached_data and node_id in cached_data[category]:
|
||||||
if node_id not in metadata[category]:
|
if node_id not in metadata[category]:
|
||||||
metadata[category][node_id] = cached_data[category][
|
metadata[category][node_id] = cached_data[category][
|
||||||
node_id
|
node_id
|
||||||
]
|
]
|
||||||
|
|
||||||
def record_node_execution(self, node_id, class_type, inputs, outputs):
|
def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None):
|
||||||
"""Record information about a node's execution"""
|
"""Record information about a node's execution"""
|
||||||
if not self.current_prompt_id:
|
if not self.current_prompt_id:
|
||||||
return
|
return
|
||||||
@@ -158,17 +172,18 @@ class MetadataRegistry:
|
|||||||
|
|
||||||
# Extract node-specific metadata
|
# Extract node-specific metadata
|
||||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||||
extractor.extract(
|
if extractor is GenericNodeExtractor:
|
||||||
node_id,
|
extractor.extract(node_id, processed_inputs, outputs,
|
||||||
processed_inputs,
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
outputs,
|
return_types=return_types)
|
||||||
self.prompt_metadata[self.current_prompt_id],
|
else:
|
||||||
)
|
extractor.extract(node_id, processed_inputs, outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id])
|
||||||
|
|
||||||
# Cache this node's metadata
|
# Cache this node's metadata
|
||||||
self._cache_node_metadata(node_id, class_type)
|
self._cache_node_metadata(node_id, class_type)
|
||||||
|
|
||||||
def update_node_execution(self, node_id, class_type, outputs):
|
def update_node_execution(self, node_id, class_type, outputs, return_types=None):
|
||||||
"""Update node metadata with output information"""
|
"""Update node metadata with output information"""
|
||||||
if not self.current_prompt_id:
|
if not self.current_prompt_id:
|
||||||
return
|
return
|
||||||
@@ -179,9 +194,17 @@ class MetadataRegistry:
|
|||||||
# Use the same extractor to update with outputs
|
# Use the same extractor to update with outputs
|
||||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||||
if hasattr(extractor, "update"):
|
if hasattr(extractor, "update"):
|
||||||
extractor.update(
|
if extractor is GenericNodeExtractor:
|
||||||
node_id, processed_outputs, self.prompt_metadata[self.current_prompt_id]
|
extractor.update(
|
||||||
)
|
node_id, processed_outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
|
return_types=return_types,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
extractor.update(
|
||||||
|
node_id, processed_outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
|
)
|
||||||
|
|
||||||
# Update the cached metadata for this node
|
# Update the cached metadata for this node
|
||||||
self._cache_node_metadata(node_id, class_type)
|
self._cache_node_metadata(node_id, class_type)
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
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):
|
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||||
@@ -31,11 +31,78 @@ class NodeMetadataExtractor:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
class GenericNodeExtractor(NodeMetadataExtractor):
|
class GenericNodeExtractor(NodeMetadataExtractor):
|
||||||
"""Default extractor for nodes without specific handling"""
|
"""Fallback extractor with type-signature-based detection.
|
||||||
|
|
||||||
|
When a node is not in the NODE_EXTRACTORS registry, the hook layer
|
||||||
|
passes ``return_types`` from ``obj.RETURN_TYPES``:
|
||||||
|
|
||||||
|
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
|
||||||
|
are checked for a model file name and stored as checkpoint metadata.
|
||||||
|
* ``CONDITIONING`` output: common text input fields are checked for
|
||||||
|
prompt text and stored as prompt metadata.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Input field names that carry a model path in loader-style nodes.
|
||||||
|
_MODEL_NAME_FIELDS = (
|
||||||
|
"ckpt_name", "unet_name", "model_path", "model_name", "gguf_name",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Extensions used by checkpoint_scanner.py — only record values that look
|
||||||
|
# like real model filenames to avoid capturing unrelated string fields.
|
||||||
|
_MODEL_EXTENSIONS = {
|
||||||
|
".ckpt", ".pt", ".pt2", ".bin", ".pth", ".safetensors", ".pkl", ".sft", ".gguf",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Input field names that may carry prompt text in encoder-style nodes.
|
||||||
|
_TEXT_FIELDS = ("text", "clip_l", "t5xxl", "prompt", "positive", "negative")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def extract(node_id, inputs, outputs, metadata):
|
def extract(node_id, inputs, outputs, metadata, return_types=None):
|
||||||
pass
|
if return_types is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
# — MODEL loader detection (checkpoint / UNET / GGUF) —
|
||||||
|
if "MODEL" in return_types or any("MODEL" in str(t) for t in return_types):
|
||||||
|
for field in GenericNodeExtractor._MODEL_NAME_FIELDS:
|
||||||
|
val = inputs.get(field)
|
||||||
|
if val and isinstance(val, str) and val.strip():
|
||||||
|
name = val.strip()
|
||||||
|
if not any(name.lower().endswith(ext) for ext in GenericNodeExtractor._MODEL_EXTENSIONS):
|
||||||
|
continue
|
||||||
|
_store_checkpoint_metadata(metadata, node_id, name)
|
||||||
|
return
|
||||||
|
|
||||||
|
# — CONDITIONING encoder detection (CLIPTextEncode, Flux, custom) —
|
||||||
|
if "CONDITIONING" in return_types or any("CONDITIONING" in str(t) for t in return_types):
|
||||||
|
text = None
|
||||||
|
for field in GenericNodeExtractor._TEXT_FIELDS:
|
||||||
|
val = inputs.get(field)
|
||||||
|
if val and isinstance(val, str) and val.strip():
|
||||||
|
text = val.strip()
|
||||||
|
break
|
||||||
|
if text:
|
||||||
|
prompt_data = metadata.setdefault(PROMPTS, {})
|
||||||
|
prompt_data[node_id] = {
|
||||||
|
"text": text,
|
||||||
|
"node_id": node_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def update(node_id, outputs, metadata, return_types=None):
|
||||||
|
if return_types is None:
|
||||||
|
return
|
||||||
|
if "CONDITIONING" not in return_types and not any(
|
||||||
|
"CONDITIONING" in str(t) for t in return_types
|
||||||
|
):
|
||||||
|
return
|
||||||
|
if node_id not in metadata.get(PROMPTS, {}):
|
||||||
|
return
|
||||||
|
if outputs and isinstance(outputs, list) and len(outputs) > 0:
|
||||||
|
if isinstance(outputs[0], tuple) and len(outputs[0]) > 0:
|
||||||
|
cond = outputs[0][0]
|
||||||
|
if cond is not None:
|
||||||
|
metadata[PROMPTS][node_id]["conditioning"] = cond
|
||||||
|
|
||||||
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def extract(node_id, inputs, outputs, metadata):
|
def extract(node_id, inputs, outputs, metadata):
|
||||||
@@ -1154,6 +1221,32 @@ class CR_ApplyControlNetStackExtractor(NodeMetadataExtractor):
|
|||||||
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
||||||
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
|
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
|
# Registry of node-specific extractors
|
||||||
# Keys are node class names
|
# Keys are node class names
|
||||||
NODE_EXTRACTORS = {
|
NODE_EXTRACTORS = {
|
||||||
@@ -1221,5 +1314,7 @@ NODE_EXTRACTORS = {
|
|||||||
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
||||||
# Image
|
# Image
|
||||||
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
||||||
|
# Metadata overwrite
|
||||||
|
"MetadataOverwriteLM": MetadataOverwriteExtractor,
|
||||||
# Add other nodes as needed
|
# Add other nodes as needed
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,157 @@
|
|||||||
|
"""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.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"model": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"tooltip": (
|
||||||
|
"The checkpoint or diffusion model (UNet) used "
|
||||||
|
"for generation. Only overwrites when non-empty."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"loras": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"multiline": True,
|
||||||
|
"tooltip": (
|
||||||
|
"LoRA syntax, e.g. <lora:name:strength> "
|
||||||
|
"or <lora:name:model_strength:clip_strength>, "
|
||||||
|
"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,)
|
||||||
@@ -471,6 +471,9 @@ class SaveImageLM:
|
|||||||
params.append(f"Clip skip: {abs(cs)}")
|
params.append(f"Clip skip: {abs(cs)}")
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
pass
|
pass
|
||||||
|
additional_data = metadata_dict.get("additional_data", "")
|
||||||
|
if additional_data:
|
||||||
|
params.append(additional_data)
|
||||||
if ckpt_hash:
|
if ckpt_hash:
|
||||||
params.append(f"Model hash: {ckpt_hash[:10].upper()}")
|
params.append(f"Model hash: {ckpt_hash[:10].upper()}")
|
||||||
if ckpt_display_name:
|
if ckpt_display_name:
|
||||||
|
|||||||
@@ -30,10 +30,10 @@ def test_metadata_hook_installs_and_traces_execution(monkeypatch, metadata_regis
|
|||||||
|
|
||||||
calls = []
|
calls = []
|
||||||
|
|
||||||
def record_stub(self, node_id, class_type, inputs, outputs):
|
def record_stub(self, node_id, class_type, inputs, outputs, return_types=None):
|
||||||
calls.append(("record", node_id, class_type, inputs))
|
calls.append(("record", node_id, class_type, inputs))
|
||||||
|
|
||||||
def update_stub(self, node_id, class_type, outputs):
|
def update_stub(self, node_id, class_type, outputs, return_types=None):
|
||||||
calls.append(("update", node_id, class_type, outputs))
|
calls.append(("update", node_id, class_type, outputs))
|
||||||
|
|
||||||
monkeypatch.setattr(MetadataRegistry, "record_node_execution", record_stub)
|
monkeypatch.setattr(MetadataRegistry, "record_node_execution", record_stub)
|
||||||
@@ -820,3 +820,220 @@ def test_lora_manager_checkpoint_and_unet_loaders_extract_models(metadata_regist
|
|||||||
"type": "checkpoint",
|
"type": "checkpoint",
|
||||||
"node_id": "unet_node",
|
"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": "",
|
||||||
|
"model": "myModel.safetensors",
|
||||||
|
"loras": "<lora:detail:0.8>",
|
||||||
|
"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["model"] == "myModel.safetensors"
|
||||||
|
assert params["loras"] == "<lora:detail:0.8>"
|
||||||
|
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": "", "model": "",
|
||||||
|
"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": "", "model": "",
|
||||||
|
"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()
|
||||||
|
|||||||
+43
-10
@@ -7,12 +7,16 @@ import { app } from "../../scripts/app.js";
|
|||||||
// Roles are stored in ``node.properties.lm_marker_role`` and automatically
|
// Roles are stored in ``node.properties.lm_marker_role`` and automatically
|
||||||
// persist with the workflow JSON.
|
// persist with the workflow JSON.
|
||||||
//
|
//
|
||||||
|
// Two categories:
|
||||||
|
// send_* – consumed by the standalone UI's "Send to Workflow" feature
|
||||||
|
// meta_* – consumed by the metadata processor to override heuristic inference
|
||||||
|
//
|
||||||
// The workflow registry reads these markers and makes them available to the
|
// The workflow registry reads these markers and makes them available to the
|
||||||
// standalone UI (e.g. ``sendEmbeddingToWorkflow`` also considers nodes marked
|
// standalone UI (e.g. ``sendEmbeddingToWorkflow`` also considers nodes marked
|
||||||
// as ``send_prompt_target``).
|
// as ``send_prompt_target``).
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|
||||||
const ROLES = {
|
const SEND_ROLES = {
|
||||||
send_prompt_target: {
|
send_prompt_target: {
|
||||||
label: "Send Prompt Target",
|
label: "Send Prompt Target",
|
||||||
emoji: "\uD83D\uDCDD",
|
emoji: "\uD83D\uDCDD",
|
||||||
@@ -23,6 +27,28 @@ const ROLES = {
|
|||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
const META_ROLES = {
|
||||||
|
meta_primary_model: {
|
||||||
|
label: "Meta hints: Primary Model",
|
||||||
|
emoji: "\uD83D\uDCA1",
|
||||||
|
},
|
||||||
|
meta_primary_sampler: {
|
||||||
|
label: "Meta hints: Primary Sampler",
|
||||||
|
emoji: "\uD83D\uDCA1",
|
||||||
|
},
|
||||||
|
meta_positive_prompt: {
|
||||||
|
label: "Meta hints: Positive Prompt",
|
||||||
|
emoji: "\uD83D\uDCA1",
|
||||||
|
},
|
||||||
|
meta_negative_prompt: {
|
||||||
|
label: "Meta hints: Negative Prompt",
|
||||||
|
emoji: "\uD83D\uDCA1",
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
// Flat lookup for setMarker / getMarker / clearMarker
|
||||||
|
const ROLES = { ...SEND_ROLES, ...META_ROLES };
|
||||||
|
|
||||||
// ---- Helpers ----------------------------------------------------------------
|
// ---- Helpers ----------------------------------------------------------------
|
||||||
|
|
||||||
function getMarker(node) {
|
function getMarker(node) {
|
||||||
@@ -54,7 +80,7 @@ function clearMarker(node) {
|
|||||||
// Restore original title: prefer stripping emoji from current title
|
// Restore original title: prefer stripping emoji from current title
|
||||||
// (captures user renames after marking), fall back to saved original.
|
// (captures user renames after marking), fall back to saved original.
|
||||||
const cleaned = node.title?.replace(
|
const cleaned = node.title?.replace(
|
||||||
/^(\u2709\uFE0F?|\u2699\uFE0F?|\uD83D\uDCDD|\uD83C\uDF9B\uFE0F?|\uD83D\uDD27)\s*/,
|
/^(\u2709\uFE0F?|\u2699\uFE0F?|\uD83D\uDCDD|\uD83C\uDF9B\uFE0F?|\uD83D\uDD27|\uD83D\uDCA1)\s*/,
|
||||||
''
|
''
|
||||||
);
|
);
|
||||||
if (cleaned && cleaned !== node.title) {
|
if (cleaned && cleaned !== node.title) {
|
||||||
@@ -84,16 +110,23 @@ function buildSubmenuOptions(node) {
|
|||||||
const currentRole = getMarker(node);
|
const currentRole = getMarker(node);
|
||||||
const options = [];
|
const options = [];
|
||||||
|
|
||||||
for (const [key, def] of Object.entries(ROLES)) {
|
const buildGroup = (roles) => {
|
||||||
const isActive = currentRole === key;
|
for (const [key, def] of Object.entries(roles)) {
|
||||||
options.push({
|
const isActive = currentRole === key;
|
||||||
content: `${isActive ? "\u2713 " : ""}${def.label}`,
|
options.push({
|
||||||
disabled: isActive,
|
content: `${isActive ? "\u2713 " : ""}${def.label}`,
|
||||||
callback: () => setMarker(node, key),
|
disabled: isActive,
|
||||||
});
|
callback: () => setMarker(node, key),
|
||||||
}
|
});
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
buildGroup(SEND_ROLES);
|
||||||
|
options.push(null); // separator
|
||||||
|
buildGroup(META_ROLES);
|
||||||
|
|
||||||
if (currentRole) {
|
if (currentRole) {
|
||||||
|
options.push(null); // separator
|
||||||
options.push({
|
options.push({
|
||||||
content: "Clear marker",
|
content: "Clear marker",
|
||||||
callback: () => clearMarker(node),
|
callback: () => clearMarker(node),
|
||||||
|
|||||||
Reference in New Issue
Block a user