Address PR review comments

This commit is contained in:
xmarre
2026-04-15 02:29:52 +02:00
parent f7b8fe5cc9
commit 19e69a8ecc
2 changed files with 38 additions and 19 deletions
+21 -13
View File
@@ -576,13 +576,12 @@ def _load_checkpoint_clip_only(
return reused_clip
_warn_pickle_checkpoint_gpu_compatibility(loader_name, ckpt_path)
is_safetensors = _is_safetensors_path(ckpt_path)
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
model_config = None if header_info is None else header_info.get("model_config")
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CLIP, source_path=ckpt_path)
with REGISTRY.load_context(kind=KIND_CLIP, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
header_info = checkpoint_component_info_from_header(ckpt_path) if _is_safetensors_path(ckpt_path) else None
model_config = None if header_info is None else header_info.get("model_config")
if model_config is None:
if _is_safetensors_path(ckpt_path):
if is_safetensors and model_config is None:
requested_device = explicit_device if explicit_device is not None else torch.device("cpu")
sd, metadata, _, _ = load_safetensors_state_dict(
ckpt_path,
@@ -590,6 +589,9 @@ def _load_checkpoint_clip_only(
return_metadata=True,
selected_keys=None,
)
else:
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
if model_config is None:
_, clip, _, _ = comfy.sd.load_state_dict_guess_config(
sd,
output_vae=False,
@@ -602,7 +604,7 @@ def _load_checkpoint_clip_only(
scaled_fp8_list = []
for key in list(sd.keys()):
if key.endswith(".scaled_fp8"):
scaled_fp8_list.append(key[:-len("scaled_fp8")])
scaled_fp8_list.append(key[:-len(".scaled_fp8")])
if scaled_fp8_list:
clip_source_sd: dict[str, Any] = {}
@@ -611,7 +613,7 @@ def _load_checkpoint_clip_only(
continue
clip_source_sd[key] = value
for prefix in scaled_fp8_list:
quant_sd, _ = comfy.utils.convert_old_quants(sd, prefix, metadata={})
quant_sd, _ = comfy.utils.convert_old_quants(sd, prefix, metadata=metadata or {})
clip_source_sd.update(quant_sd)
else:
clip_source_sd = sd
@@ -652,13 +654,12 @@ def _load_checkpoint_vae_only(
return reused_vae
_warn_pickle_checkpoint_gpu_compatibility(loader_name, ckpt_path)
is_safetensors = _is_safetensors_path(ckpt_path)
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
model_config = None if header_info is None else header_info.get("model_config")
explicit_device = REGISTRY.explicit_load_device(kind=KIND_VAE, source_path=ckpt_path)
with REGISTRY.load_context(kind=KIND_VAE, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
header_info = checkpoint_component_info_from_header(ckpt_path) if _is_safetensors_path(ckpt_path) else None
model_config = None if header_info is None else header_info.get("model_config")
if model_config is None:
if _is_safetensors_path(ckpt_path):
if is_safetensors and model_config is None:
requested_device = explicit_device if explicit_device is not None else torch.device("cpu")
sd, metadata, _, _ = load_safetensors_state_dict(
ckpt_path,
@@ -666,6 +667,9 @@ def _load_checkpoint_vae_only(
return_metadata=True,
selected_keys=None,
)
else:
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
if model_config is None:
_, _, vae, _ = comfy.sd.load_state_dict_guess_config(
sd,
output_vae=True,
@@ -681,7 +685,11 @@ def _load_checkpoint_vae_only(
filter_keys=True,
)
vae_sd = model_config.process_vae_state_dict(vae_sd)
vae = comfy.sd.VAE(sd=vae_sd, metadata=metadata)
if len(vae_sd) == 0:
_LOG.warning("%s: no VAE weights found in %s after selective checkpoint load", loader_name, ckpt_path)
vae = None
else:
vae = comfy.sd.VAE(sd=vae_sd, metadata=metadata)
_bind_vae_for_reuse(vae, source_path=ckpt_path, note="checkpoint vae", loader_key=loader_key)
return vae
+17 -6
View File
@@ -436,15 +436,26 @@ class ResidencyRegistry:
entry.last_touched = _now()
return obj
def _entry_id_for_object(self, obj: Any) -> str | None:
current = obj
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
entry_id = getattr(current, "__gpu_resident_loader_entry_id__", None)
if entry_id is None:
try:
entry_id = self._object_to_entry.get(current)
except TypeError:
entry_id = None
if entry_id is not None:
return entry_id
current = getattr(current, "parent", None)
return None
def entry_for_object(self, obj: Any) -> ResidencyEntry | None:
if obj is None:
return None
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
if entry_id is None:
try:
entry_id = self._object_to_entry.get(obj)
except TypeError:
entry_id = None
entry_id = self._entry_id_for_object(obj)
if entry_id is None:
return None
with self._lock: