Merge pull request #3 from xmarre/codex/policy-override-loader
Add explicit residency policy loader input
This commit is contained in:
+142
-45
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user