diff --git a/kj_loader.py b/kj_loader.py index 348f513..991bd36 100644 --- a/kj_loader.py +++ b/kj_loader.py @@ -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