Merge pull request #3 from xmarre/codex/policy-override-loader

Add explicit residency policy loader input
This commit is contained in:
xmarre
2026-04-12 06:49:22 +02:00
committed by GitHub
+142 -45
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import contextlib
import logging
import os
from typing import Any
import folder_paths
@@ -12,7 +13,7 @@ import comfy.utils
from comfy.cli_args import PerformanceFeature, args
from comfy.ldm.modules.attention import attention_pytorch, wrap_attn
from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, REGISTRY
from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, POLICIES, REGISTRY
_LOG = logging.getLogger(__name__)
@@ -244,6 +245,70 @@ def _build_model_options(weight_dtype: str) -> dict[str, Any]:
return model_options
def _normalize_optional_string(value: str | None) -> str | None:
if value is None:
return None
normalized = str(value).strip()
return normalized or None
def _normalize_optional_policy(value: str | None) -> str | None:
normalized = _normalize_optional_string(value)
return None if normalized is None else normalized.lower()
def _resolve_loader_policy_and_extra_state_dict(
*,
loader_name: str,
policy_override: str | None,
extra_state_dict: str | None,
) -> tuple[str | None, str | None]:
normalized_policy = _normalize_optional_policy(policy_override)
if normalized_policy is not None and normalized_policy not in POLICIES:
raise ValueError(
f"{loader_name}: unsupported policy_override {normalized_policy!r}. "
f"Expected one of: {', '.join(POLICIES)}."
)
normalized_extra = _normalize_optional_string(extra_state_dict)
if normalized_extra is None:
return normalized_policy, None
if normalized_policy is None and normalized_extra.lower() in POLICIES and not os.path.exists(normalized_extra):
_LOG.warning(
"GPU Resident Loader: interpreting legacy extra_state_dict value %r as policy_override for %s",
normalized_extra,
loader_name,
)
return normalized_extra.lower(), None
if not os.path.isfile(normalized_extra):
raise FileNotFoundError(
f"{loader_name}: extra_state_dict must point to an existing state-dict file, got {normalized_extra!r}. "
f"If you meant to pass a residency policy, connect that STRING to policy_override instead."
)
return normalized_policy, normalized_extra
@contextlib.contextmanager
def _apply_policy_override(policy_override: str | None):
if policy_override is None:
yield
return
previous_policy = REGISTRY.get_policy()
if previous_policy == policy_override:
yield
return
REGISTRY.set_policy(policy_override)
try:
yield
finally:
REGISTRY.set_policy(previous_policy)
class DiffusionModelSelectorResident:
@classmethod
def INPUT_TYPES(cls):
@@ -306,6 +371,13 @@ class DiffusionModelLoaderResident:
"forceInput": True,
"tooltip": "Optional absolute path to a second state dict merged into the main diffusion state dict before model detection.",
},
),
"policy_override": (
"STRING",
{
"forceInput": True,
"tooltip": "Optional residency policy override. Connect Set Global Residency Policy here, not to extra_state_dict.",
},
)
},
}
@@ -328,28 +400,36 @@ class DiffusionModelLoaderResident:
sage_attention: str,
enable_fp16_accumulation: bool,
extra_state_dict: str | None = None,
policy_override: str | None = None,
):
with _temporary_backend_flags(
cublas=patch_cublaslinear,
fp16_accumulation=enable_fp16_accumulation,
):
model_options = _build_model_options(weight_dtype)
unet_path = folder_paths.get_full_path_or_raise("diffusion_models", model_name)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_MODEL, source_path=unet_path)
policy_override, extra_state_dict = _resolve_loader_policy_and_extra_state_dict(
loader_name="Diffusion Model Loader Resident",
policy_override=policy_override,
extra_state_dict=extra_state_dict,
)
with REGISTRY.load_context(kind=KIND_MODEL, source_path=unet_path, explicit_device=explicit_device):
sd, metadata = comfy.utils.load_torch_file(unet_path, return_metadata=True)
if extra_state_dict:
extra_sd = comfy.utils.load_torch_file(extra_state_dict)
sd.update(extra_sd)
del extra_sd
with _apply_policy_override(policy_override):
with _temporary_backend_flags(
cublas=patch_cublaslinear,
fp16_accumulation=enable_fp16_accumulation,
):
model_options = _build_model_options(weight_dtype)
unet_path = folder_paths.get_full_path_or_raise("diffusion_models", model_name)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_MODEL, source_path=unet_path)
diffusion_model_prefix = comfy.sd.model_detection.unet_prefix_from_state_dict(sd)
sd = comfy.utils.state_dict_prefix_replace(sd, {diffusion_model_prefix: ""}, filter_keys=False)
model = comfy.sd.load_diffusion_model_state_dict(sd, model_options=model_options, metadata=metadata)
_apply_model_postload_options(model, compute_dtype=compute_dtype, sage_attention=sage_attention)
REGISTRY.bind_object(model, source_path=unet_path, kind=KIND_MODEL)
return (model,)
with REGISTRY.load_context(kind=KIND_MODEL, source_path=unet_path, explicit_device=explicit_device):
sd, metadata = comfy.utils.load_torch_file(unet_path, return_metadata=True)
if extra_state_dict:
extra_sd = comfy.utils.load_torch_file(extra_state_dict)
sd.update(extra_sd)
del extra_sd
diffusion_model_prefix = comfy.sd.model_detection.unet_prefix_from_state_dict(sd)
sd = comfy.utils.state_dict_prefix_replace(sd, {diffusion_model_prefix: ""}, filter_keys=False)
model = comfy.sd.load_diffusion_model_state_dict(sd, model_options=model_options, metadata=metadata)
_apply_model_postload_options(model, compute_dtype=compute_dtype, sage_attention=sage_attention)
REGISTRY.bind_object(model, source_path=unet_path, kind=KIND_MODEL)
return (model,)
class CheckpointLoaderResident:
@@ -378,7 +458,16 @@ class CheckpointLoaderResident:
"BOOLEAN",
{"default": False, "tooltip": "Set torch.backends.cuda.matmul.allow_fp16_accumulation."},
),
}
},
"optional": {
"policy_override": (
"STRING",
{
"forceInput": True,
"tooltip": "Optional residency policy override. Connect Set Global Residency Policy here.",
},
)
},
}
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
@@ -397,30 +486,38 @@ class CheckpointLoaderResident:
patch_cublaslinear: bool,
sage_attention: str,
enable_fp16_accumulation: bool,
policy_override: str | None = None,
):
with _temporary_backend_flags(
cublas=patch_cublaslinear,
fp16_accumulation=enable_fp16_accumulation,
):
model_options = _build_model_options(weight_dtype)
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CHECKPOINT, source_path=ckpt_path)
policy_override, _ = _resolve_loader_policy_and_extra_state_dict(
loader_name="Checkpoint Loader Resident",
policy_override=policy_override,
extra_state_dict=None,
)
with REGISTRY.load_context(kind=KIND_CHECKPOINT, source_path=ckpt_path, explicit_device=explicit_device):
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
with _apply_policy_override(policy_override):
with _temporary_backend_flags(
cublas=patch_cublaslinear,
fp16_accumulation=enable_fp16_accumulation,
):
model_options = _build_model_options(weight_dtype)
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CHECKPOINT, source_path=ckpt_path)
model, clip, vae, _ = comfy.sd.load_state_dict_guess_config(
sd,
output_vae=True,
output_clip=True,
embedding_directory=folder_paths.get_folder_paths("embeddings"),
metadata=metadata,
model_options=model_options,
)
_apply_model_postload_options(model, compute_dtype=compute_dtype, sage_attention=sage_attention)
REGISTRY.bind_object(model, source_path=ckpt_path, kind=KIND_MODEL, note="checkpoint model")
if clip is not None and getattr(clip, "patcher", None) is not None:
REGISTRY.bind_object(clip.patcher, source_path=ckpt_path, kind=KIND_CLIP, note="checkpoint clip")
if vae is not None and getattr(vae, "patcher", None) is not None:
REGISTRY.bind_object(vae.patcher, source_path=ckpt_path, kind=KIND_VAE, note="checkpoint vae")
return model, clip, vae
with REGISTRY.load_context(kind=KIND_CHECKPOINT, source_path=ckpt_path, explicit_device=explicit_device):
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
model, clip, vae, _ = comfy.sd.load_state_dict_guess_config(
sd,
output_vae=True,
output_clip=True,
embedding_directory=folder_paths.get_folder_paths("embeddings"),
metadata=metadata,
model_options=model_options,
)
_apply_model_postload_options(model, compute_dtype=compute_dtype, sage_attention=sage_attention)
REGISTRY.bind_object(model, source_path=ckpt_path, kind=KIND_MODEL, note="checkpoint model")
if clip is not None and getattr(clip, "patcher", None) is not None:
REGISTRY.bind_object(clip.patcher, source_path=ckpt_path, kind=KIND_CLIP, note="checkpoint clip")
if vae is not None and getattr(vae, "patcher", None) is not None:
REGISTRY.bind_object(vae.patcher, source_path=ckpt_path, kind=KIND_VAE, note="checkpoint vae")
return model, clip, vae