from __future__ import annotations import contextlib import json import logging import os from typing import Any import folder_paths import torch import comfy.sd import comfy.utils from comfy.cli_args import PerformanceFeature, args from comfy.ldm.modules.attention import attention_pytorch, wrap_attn from .cleanup import trim_resident_vram_for_load from .patches import ( checkpoint_component_info_from_header, estimate_checkpoint_component_bytes, estimate_safetensors_tensor_bytes, infer_unet_prefix_from_keys, load_safetensors_state_dict, ) from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, POLICIES, REGISTRY _LOG = logging.getLogger(__name__) SAGE_ATTN_MODES = [ "disabled", "auto", "sageattn_qk_int8_pv_fp16_cuda", "sageattn_qk_int8_pv_fp16_triton", "sageattn_qk_int8_pv_fp8_cuda", "sageattn_qk_int8_pv_fp8_cuda++", "sageattn3", "sageattn3_per_block_mean", ] DTYPE_MAP = { "fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e5m2": torch.float8_e5m2, "fp16": torch.float16, "bf16": torch.bfloat16, "fp32": torch.float32, } UNET_PREFIX_CANDIDATES = ( "model.diffusion_model.", "model.model.", "net.", "model.", ) def _set_cublas_linear(enabled: bool) -> None: if enabled: args.fast.add(PerformanceFeature.CublasOps) else: args.fast.discard(PerformanceFeature.CublasOps) def _set_fp16_accumulation(enabled: bool) -> None: if not hasattr(torch.backends.cuda, "matmul"): if enabled: raise RuntimeError( "Failed to enable fp16 accumulation. This requires a PyTorch build exposing " "torch.backends.cuda.matmul.allow_fp16_accumulation." ) _LOG.warning( "GPU Resident Loader: fp16 accumulation toggle is unavailable on this PyTorch build; leaving default behavior." ) return flag = getattr(torch.backends.cuda.matmul, "allow_fp16_accumulation", None) if flag is None: if enabled: raise RuntimeError( "Failed to enable fp16 accumulation. This requires a PyTorch build exposing " "torch.backends.cuda.matmul.allow_fp16_accumulation." ) _LOG.warning( "GPU Resident Loader: fp16 accumulation toggle is unavailable on this PyTorch build; leaving default behavior." ) return torch.backends.cuda.matmul.allow_fp16_accumulation = bool(enabled) def _get_fp16_accumulation_state() -> bool | None: matmul = getattr(torch.backends.cuda, "matmul", None) return getattr(matmul, "allow_fp16_accumulation", None) @contextlib.contextmanager def _temporary_backend_flags(*, cublas: bool, fp16_accumulation: bool): prev_cublas = PerformanceFeature.CublasOps in args.fast prev_fp16 = _get_fp16_accumulation_state() try: _set_cublas_linear(cublas) _set_fp16_accumulation(fp16_accumulation) yield finally: _set_cublas_linear(prev_cublas) if prev_fp16 is not None: torch.backends.cuda.matmul.allow_fp16_accumulation = prev_fp16 def get_sage_func(sage_attention: str, allow_compile: bool = False): _LOG.info("GPU Resident Loader: using sage attention mode %s", sage_attention) if sage_attention == "auto": from sageattention import sageattn def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"): return sageattn(q, k, v, is_causal=is_causal, attn_mask=attn_mask, tensor_layout=tensor_layout) elif sage_attention == "sageattn_qk_int8_pv_fp16_cuda": from sageattention import sageattn_qk_int8_pv_fp16_cuda def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"): return sageattn_qk_int8_pv_fp16_cuda( q, k, v, is_causal=is_causal, attn_mask=attn_mask, pv_accum_dtype="fp32", tensor_layout=tensor_layout, ) elif sage_attention == "sageattn_qk_int8_pv_fp16_triton": from sageattention import sageattn_qk_int8_pv_fp16_triton def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"): return sageattn_qk_int8_pv_fp16_triton( q, k, v, is_causal=is_causal, attn_mask=attn_mask, tensor_layout=tensor_layout, ) elif sage_attention == "sageattn_qk_int8_pv_fp8_cuda": from sageattention import sageattn_qk_int8_pv_fp8_cuda def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"): return sageattn_qk_int8_pv_fp8_cuda( q, k, v, is_causal=is_causal, attn_mask=attn_mask, pv_accum_dtype="fp32+fp32", tensor_layout=tensor_layout, ) elif sage_attention == "sageattn_qk_int8_pv_fp8_cuda++": from sageattention import sageattn_qk_int8_pv_fp8_cuda def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"): return sageattn_qk_int8_pv_fp8_cuda( q, k, v, is_causal=is_causal, attn_mask=attn_mask, pv_accum_dtype="fp32+fp16", tensor_layout=tensor_layout, ) elif "sageattn3" in sage_attention: from sageattn3 import sageattn3_blackwell def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD", **kwargs): q, k, v = [x.transpose(1, 2) if tensor_layout == "NHD" else x for x in (q, k, v)] out = sageattn3_blackwell( q, k, v, is_causal=is_causal, attn_mask=attn_mask, per_block_mean=(sage_attention == "sageattn3_per_block_mean"), ) return out.transpose(1, 2) if tensor_layout == "NHD" else out else: raise ValueError(f"Unsupported sage attention mode: {sage_attention}") compiler = getattr(torch, "compiler", None) disable = getattr(compiler, "disable", None) if not allow_compile and callable(disable): sage_func = disable()(sage_func) @wrap_attn def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): if kwargs.get("low_precision_attention", True) is False: return attention_pytorch( q, k, v, heads, mask=mask, skip_reshape=skip_reshape, skip_output_reshape=skip_output_reshape, **kwargs, ) in_dtype = v.dtype if q.dtype == torch.float32 or k.dtype == torch.float32 or v.dtype == torch.float32: q, k, v = q.to(torch.float16), k.to(torch.float16), v.to(torch.float16) if skip_reshape: batch, _, _, dim_head = q.shape tensor_layout = "HND" else: batch, _, dim_head = q.shape dim_head //= heads q, k, v = [tensor.view(batch, -1, heads, dim_head) for tensor in (q, k, v)] tensor_layout = "NHD" if mask is not None: if mask.ndim == 2: mask = mask.unsqueeze(0) if mask.ndim == 3: mask = mask.unsqueeze(1) out = sage_func(q, k, v, attn_mask=mask, is_causal=False, tensor_layout=tensor_layout).to(in_dtype) if tensor_layout == "HND": if not skip_output_reshape: out = out.transpose(1, 2).reshape(batch, -1, heads * dim_head) else: if skip_output_reshape: out = out.transpose(1, 2) else: out = out.reshape(batch, -1, heads * dim_head) return out return attention_sage def _apply_model_postload_options(model, *, compute_dtype: str, sage_attention: str) -> None: if dtype := DTYPE_MAP.get(compute_dtype): model.set_model_compute_dtype(dtype) model.force_cast_weights = False _LOG.info("GPU Resident Loader: set compute dtype to %s", dtype) if sage_attention != "disabled": new_attention = get_sage_func(sage_attention) def attention_override_sage(func, *args, **kwargs): return new_attention.__wrapped__(*args, **kwargs) model.model_options.setdefault("transformer_options", {})["optimized_attention_override"] = attention_override_sage def _build_model_options(weight_dtype: str) -> dict[str, Any]: model_options: dict[str, Any] = {} if dtype := DTYPE_MAP.get(weight_dtype): model_options["dtype"] = dtype _LOG.info("GPU Resident Loader: set weight dtype to %s", dtype) if weight_dtype == "fp8_e4m3fn_fast": model_options["dtype"] = torch.float8_e4m3fn model_options["fp8_optimizations"] = True 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 _effective_policy_name(policy_override: str | None) -> str: return REGISTRY.get_policy() if policy_override is None else policy_override def _make_loader_key(loader_name: str, **payload: Any) -> str: normalized_payload = {"loader": loader_name} normalized_payload.update(payload) return json.dumps(normalized_payload, sort_keys=True, separators=(",", ":")) 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 def _is_safetensors_path(path: str) -> bool: lowered = path.lower() return lowered.endswith(".safetensors") or lowered.endswith(".sft") def _warn_pickle_checkpoint_gpu_compatibility(loader_name: str, path: str) -> None: if _is_safetensors_path(path): return if not REGISTRY.wants_gpu_ingest(KIND_MODEL): return _LOG.warning( "%s: %s is using the compatibility path through torch.load() on CPU before tensor-by-tensor copies to GPU. " "Convert hot checkpoints to safetensors with scripts/convert_checkpoint_to_safetensors.py for the real fast path.", loader_name, path, ) def _fallback_file_size_bytes(path: str) -> int: try: return int(os.path.getsize(path)) except OSError: return 0 def _weight_dtype_override(weight_dtype: str) -> torch.dtype | None: if weight_dtype == "fp8_e4m3fn_fast": return torch.float8_e4m3fn return DTYPE_MAP.get(weight_dtype) def _normalize_keep_models(*objects: Any) -> tuple[Any, ...]: keep_models: list[Any] = [] for obj in objects: if obj is None: continue patcher = getattr(obj, "patcher", None) keep = patcher if patcher is not None else obj if not any(existing is keep for existing in keep_models): keep_models.append(keep) return tuple(keep_models) def _estimate_extra_state_dict_bytes(extra_state_dict: str | None, *, weight_dtype: str) -> int: if not extra_state_dict: return 0 dtype_override = _weight_dtype_override(weight_dtype) if _is_safetensors_path(extra_state_dict): estimated = estimate_safetensors_tensor_bytes(extra_state_dict, dtype_override=dtype_override) if estimated is not None: return int(estimated) return _fallback_file_size_bytes(extra_state_dict) def _estimate_model_load_bytes( source_path: str, *, cache_scope: str, weight_dtype: str, extra_state_dict: str | None, ) -> int: dtype_override = _weight_dtype_override(weight_dtype) estimated = None if _is_safetensors_path(source_path): if cache_scope == "checkpoint_model": estimated = estimate_checkpoint_component_bytes( source_path, KIND_MODEL, dtype_override=dtype_override, ) else: estimated = estimate_safetensors_tensor_bytes(source_path, dtype_override=dtype_override) if estimated is None: estimated = _fallback_file_size_bytes(source_path) return int(estimated) + _estimate_extra_state_dict_bytes(extra_state_dict, weight_dtype=weight_dtype) def _estimate_checkpoint_aux_component_bytes(ckpt_path: str, *, kind: str) -> int: estimated = None if _is_safetensors_path(ckpt_path): estimated = estimate_checkpoint_component_bytes(ckpt_path, kind) if estimated is None: estimated = _fallback_file_size_bytes(ckpt_path) return int(estimated) def _maybe_trim_before_load( *, loader_name: str, reason: str, explicit_device: torch.device | None, required_bytes: int, keep_models: tuple[Any, ...] = (), ) -> None: if explicit_device is None or explicit_device.type != "cuda": return if required_bytes <= 0: return report = trim_resident_vram_for_load( required_bytes=required_bytes, reason=reason, device=explicit_device, keep_models=keep_models, ) if report["freed_vram_bytes"] > 0: _LOG.info( "%s: adaptively freed %.2f GiB before load (%s, estimated %.2f GiB + %.2f GiB headroom)", loader_name, report["freed_vram_bytes"] / (1024 ** 3), report["stopped_reason"], report["estimated_load_bytes"] / (1024 ** 3), report["adaptive_headroom_bytes"] / (1024 ** 3), ) if not report["target_met"]: _LOG.warning( "%s: adaptive trim could not reach the estimated headroom for %s (%s, free %.2f GiB / target %.2f GiB)", loader_name, reason, report["stopped_reason"], report["free_after_bytes"] / (1024 ** 3), report["target_free_vram_bytes"] / (1024 ** 3), ) def _selected_unet_key_map_from_header( path: str, *, known_unet_keys: set[str] | None = None, ) -> tuple[dict[str, str], str | None]: from safetensors import safe_open with safe_open(path, framework="pt", device="cpu") as handle: all_keys = list(handle.keys()) prefix = infer_unet_prefix_from_keys(all_keys) if known_unet_keys is None: selected = {key: key[len(prefix):] for key in all_keys if key.startswith(prefix)} return (selected or {key: key for key in all_keys}), prefix prefixes_to_try: list[str] = [] for candidate in (prefix, *UNET_PREFIX_CANDIDATES): if candidate and candidate not in prefixes_to_try: prefixes_to_try.append(candidate) selected: dict[str, str] = {} for key in all_keys: if key in known_unet_keys: selected[key] = key continue for candidate in prefixes_to_try: if key.startswith(candidate): stripped = key[len(candidate):] if stripped in known_unet_keys: selected[key] = stripped break return selected, prefix def _load_matching_extra_unet_state_dict( extra_state_dict_path: str, *, requested_device: torch.device, known_unet_keys: set[str], ) -> dict[str, Any]: if _is_safetensors_path(extra_state_dict_path): selected_keys, _ = _selected_unet_key_map_from_header( extra_state_dict_path, known_unet_keys=known_unet_keys, ) extra_sd, _, _, _ = load_safetensors_state_dict( extra_state_dict_path, requested_device, selected_keys=selected_keys, ) return extra_sd extra_sd = comfy.utils.load_torch_file(extra_state_dict_path) return _extract_unet_state_dict(extra_sd, known_unet_keys=known_unet_keys) def _extract_unet_state_dict( sd: dict[str, Any], *, diffusion_model_prefix: str | None = None, known_unet_keys: set[str] | None = None, ) -> dict[str, Any]: if known_unet_keys is None: if diffusion_model_prefix is None: diffusion_model_prefix = comfy.sd.model_detection.unet_prefix_from_state_dict(sd) if diffusion_model_prefix: prefix_len = len(diffusion_model_prefix) extracted = {key[prefix_len:]: value for key, value in sd.items() if key.startswith(diffusion_model_prefix)} if extracted: return extracted return dict(sd) prefixes_to_try: list[str] = [] def add_prefix(prefix: str | None) -> None: if prefix and prefix not in prefixes_to_try: prefixes_to_try.append(prefix) add_prefix(diffusion_model_prefix) add_prefix(comfy.sd.model_detection.unet_prefix_from_state_dict(sd)) for prefix in UNET_PREFIX_CANDIDATES: add_prefix(prefix) extracted: dict[str, Any] = {} for key, value in sd.items(): if key in known_unet_keys: extracted[key] = value continue for prefix in prefixes_to_try: if not key.startswith(prefix): continue stripped = key[len(prefix):] if stripped in known_unet_keys: extracted[stripped] = value break return extracted @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) def _bind_model_for_reuse(model, *, source_path: str, note: str, loader_key: str) -> None: REGISTRY.bind_object(model, source_path=source_path, kind=KIND_MODEL, note=note, loader_key=loader_key) def _bind_clip_for_reuse(clip, *, source_path: str, note: str, loader_key: str) -> None: if clip is not None and getattr(clip, "patcher", None) is not None: REGISTRY.bind_object( clip.patcher, source_path=source_path, kind=KIND_CLIP, note=note, loader_key=loader_key, reusable_obj=clip, ) def _bind_vae_for_reuse(vae, *, source_path: str, note: str, loader_key: str) -> None: if vae is not None and getattr(vae, "patcher", None) is not None: REGISTRY.bind_object( vae.patcher, source_path=source_path, kind=KIND_VAE, note=note, loader_key=loader_key, reusable_obj=vae, ) def _load_resident_diffusion_model( *, loader_name: str, cache_scope: str, source_path: str, note: str, weight_dtype: str, compute_dtype: str, patch_cublaslinear: bool, sage_attention: str, enable_fp16_accumulation: bool, extra_state_dict: str | None = None, policy_override: str | None = None, keep_models: tuple[Any, ...] = (), ) -> Any: model_options = _build_model_options(weight_dtype) effective_policy = _effective_policy_name(policy_override) loader_key = _make_loader_key( cache_scope, component="model", weight_dtype=weight_dtype, compute_dtype=compute_dtype, patch_cublaslinear=patch_cublaslinear, sage_attention=sage_attention, enable_fp16_accumulation=enable_fp16_accumulation, extra_state_dict=extra_state_dict, policy=effective_policy, ) reused_model = REGISTRY.lookup_live_object(kind=KIND_MODEL, source_path=source_path, loader_key=loader_key) if reused_model is not None: _LOG.info("%s: reusing live model for %s", loader_name, source_path) return reused_model _warn_pickle_checkpoint_gpu_compatibility(loader_name, source_path) explicit_device = REGISTRY.explicit_load_device(kind=KIND_MODEL, source_path=source_path) _maybe_trim_before_load( loader_name=loader_name, reason=f"model load for {os.path.basename(source_path)}", explicit_device=explicit_device, required_bytes=_estimate_model_load_bytes( source_path, cache_scope=cache_scope, weight_dtype=weight_dtype, extra_state_dict=extra_state_dict, ), keep_models=keep_models, ) with _temporary_backend_flags( cublas=patch_cublaslinear, fp16_accumulation=enable_fp16_accumulation, ): with REGISTRY.load_context( kind=KIND_MODEL, source_path=source_path, explicit_device=explicit_device, cache_key=loader_key, ): sd, metadata = comfy.utils.load_torch_file(source_path, return_metadata=True) if not _is_safetensors_path(source_path): sd = _extract_unet_state_dict(sd) if extra_state_dict: requested_device = explicit_device if explicit_device is not None else torch.device("cpu") sd.update( _load_matching_extra_unet_state_dict( extra_state_dict, requested_device=requested_device, known_unet_keys=set(sd), ) ) 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) _bind_model_for_reuse(model, source_path=source_path, note=note, loader_key=loader_key) return model def _checkpoint_model_loader_key( *, weight_dtype: str, compute_dtype: str, patch_cublaslinear: bool, sage_attention: str, enable_fp16_accumulation: bool, policy_override: str | None, ) -> str: return _make_loader_key( "checkpoint_model", component="model", weight_dtype=weight_dtype, compute_dtype=compute_dtype, patch_cublaslinear=patch_cublaslinear, sage_attention=sage_attention, enable_fp16_accumulation=enable_fp16_accumulation, extra_state_dict=None, policy=_effective_policy_name(policy_override), ) def _checkpoint_component_loader_key(component: str, policy_override: str | None) -> str: return _make_loader_key(f"checkpoint_{component}", component=component, policy=_effective_policy_name(policy_override)) def _load_checkpoint_clip_only( *, ckpt_path: str, policy_override: str | None, loader_name: str, keep_models: tuple[Any, ...] = (), ): loader_key = _checkpoint_component_loader_key("clip", policy_override) reused_clip = REGISTRY.lookup_live_object(kind=KIND_CLIP, source_path=ckpt_path, loader_key=loader_key) if reused_clip is not None: _LOG.info("%s: reusing live CLIP for %s", loader_name, ckpt_path) 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) _maybe_trim_before_load( loader_name=loader_name, reason=f"CLIP load for {os.path.basename(ckpt_path)}", explicit_device=explicit_device, required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_CLIP), keep_models=keep_models, ) with REGISTRY.load_context(kind=KIND_CLIP, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key): 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, requested_device, 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, output_clip=True, output_model=False, embedding_directory=folder_paths.get_folder_paths("embeddings"), metadata=metadata, ) else: scaled_fp8_list = [] for key in list(sd.keys()): if key.endswith(".scaled_fp8"): scaled_fp8_list.append(key[:-len(".scaled_fp8")]) if scaled_fp8_list: clip_source_sd: dict[str, Any] = {} for key, value in sd.items(): if any(key.startswith(prefix) for prefix in scaled_fp8_list): continue clip_source_sd[key] = value for prefix in scaled_fp8_list: quant_sd, _ = comfy.utils.convert_old_quants(sd, prefix, metadata=metadata or {}) clip_source_sd.update(quant_sd) else: clip_source_sd = sd clip_target = model_config.clip_target(state_dict=clip_source_sd) if clip_target is None: clip = None else: clip_sd = model_config.process_clip_state_dict(clip_source_sd) if len(clip_sd) == 0: _LOG.warning("%s: no CLIP/text encoder weights found in %s after selective checkpoint load", loader_name, ckpt_path) clip = None else: parameters = comfy.utils.calculate_parameters(clip_sd) clip = comfy.sd.CLIP( clip_target, embedding_directory=folder_paths.get_folder_paths("embeddings"), tokenizer_data=clip_sd, parameters=parameters, state_dict=clip_sd, model_options={}, disable_dynamic=False, ) _bind_clip_for_reuse(clip, source_path=ckpt_path, note="checkpoint clip", loader_key=loader_key) return clip def _load_checkpoint_vae_only( *, ckpt_path: str, policy_override: str | None, loader_name: str, keep_models: tuple[Any, ...] = (), ): loader_key = _checkpoint_component_loader_key("vae", policy_override) reused_vae = REGISTRY.lookup_live_object(kind=KIND_VAE, source_path=ckpt_path, loader_key=loader_key) if reused_vae is not None: _LOG.info("%s: reusing live VAE for %s", loader_name, ckpt_path) 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) _maybe_trim_before_load( loader_name=loader_name, reason=f"VAE load for {os.path.basename(ckpt_path)}", explicit_device=explicit_device, required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_VAE), keep_models=keep_models, ) with REGISTRY.load_context(kind=KIND_VAE, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key): 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, requested_device, 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, output_clip=False, output_model=False, embedding_directory=folder_paths.get_folder_paths("embeddings"), metadata=metadata, ) else: vae_sd = comfy.utils.state_dict_prefix_replace( sd, {prefix: "" for prefix in model_config.vae_key_prefix}, filter_keys=True, ) vae_sd = model_config.process_vae_state_dict(vae_sd) 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 def _load_full_checkpoint( *, ckpt_path: str, weight_dtype: str, compute_dtype: str, patch_cublaslinear: bool, sage_attention: str, enable_fp16_accumulation: bool, policy_override: str | None, loader_name: str, ): model_key = _checkpoint_model_loader_key( weight_dtype=weight_dtype, compute_dtype=compute_dtype, patch_cublaslinear=patch_cublaslinear, sage_attention=sage_attention, enable_fp16_accumulation=enable_fp16_accumulation, policy_override=policy_override, ) clip_key = _checkpoint_component_loader_key("clip", policy_override) vae_key = _checkpoint_component_loader_key("vae", policy_override) model = REGISTRY.lookup_live_object(kind=KIND_MODEL, source_path=ckpt_path, loader_key=model_key) clip = REGISTRY.lookup_live_object(kind=KIND_CLIP, source_path=ckpt_path, loader_key=clip_key) vae = REGISTRY.lookup_live_object(kind=KIND_VAE, source_path=ckpt_path, loader_key=vae_key) missing = [name for name, value in (("model", model), ("clip", clip), ("vae", vae)) if value is None] if not missing: _LOG.info("%s: reusing live checkpoint outputs for %s", loader_name, ckpt_path) return model, clip, vae keep_models = _normalize_keep_models(model, clip, vae) if model is None: model = _load_resident_diffusion_model( loader_name=loader_name, cache_scope="checkpoint_model", source_path=ckpt_path, note="checkpoint model", weight_dtype=weight_dtype, compute_dtype=compute_dtype, patch_cublaslinear=patch_cublaslinear, sage_attention=sage_attention, enable_fp16_accumulation=enable_fp16_accumulation, policy_override=policy_override, keep_models=keep_models, ) keep_models = keep_models + _normalize_keep_models(model) if clip is None: clip = _load_checkpoint_clip_only( ckpt_path=ckpt_path, policy_override=policy_override, loader_name=loader_name, keep_models=keep_models, ) keep_models = keep_models + _normalize_keep_models(clip) if vae is None: vae = _load_checkpoint_vae_only( ckpt_path=ckpt_path, policy_override=policy_override, loader_name=loader_name, keep_models=keep_models, ) return model, clip, vae class DiffusionModelSelectorResident: @classmethod def INPUT_TYPES(cls): ltx2_connector_models = folder_paths.get_filename_list("text_encoders") ltx2_connector_models = [m for m in ltx2_connector_models if "connector" in m.lower()] return { "required": { "model_name": ( folder_paths.get_filename_list("diffusion_models") + ltx2_connector_models, {"tooltip": "The name of the diffusion model or connector to resolve."}, ), } } RETURN_TYPES = ("STRING",) RETURN_NAMES = ("model_path",) FUNCTION = "get_path" CATEGORY = "GPU Resident Loader/loaders" DESCRIPTION = "Returns the absolute model path as a string. Mirrors the selector behavior of KJ's diffusion model selector." def get_path(self, model_name: str): if "connector" in model_name.lower(): model_path = folder_paths.get_full_path_or_raise("text_encoders", model_name) else: model_path = folder_paths.get_full_path_or_raise("diffusion_models", model_name) return (model_path,) class DiffusionModelLoaderResident: @classmethod def INPUT_TYPES(cls): return { "required": { "model_name": ( folder_paths.get_filename_list("diffusion_models"), {"tooltip": "The diffusion model file to load."}, ), "weight_dtype": (["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp16", "bf16", "fp32"],), "compute_dtype": ( ["default", "fp16", "bf16", "fp32"], {"default": "default", "tooltip": "Compute dtype to apply after model creation."}, ), "patch_cublaslinear": ( "BOOLEAN", {"default": False, "tooltip": "Toggle ComfyUI's cublas_ops performance feature."}, ), "sage_attention": ( SAGE_ATTN_MODES, {"default": "disabled", "tooltip": "Patch optimized attention override to a SageAttention variant."}, ), "enable_fp16_accumulation": ( "BOOLEAN", {"default": False, "tooltip": "Set torch.backends.cuda.matmul.allow_fp16_accumulation."}, ), }, "optional": { "extra_state_dict": ( "STRING", { "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.", }, ) }, } RETURN_TYPES = ("MODEL",) FUNCTION = "patch_and_load" CATEGORY = "GPU Resident Loader/loaders" DESCRIPTION = ( "KJ-compatible diffusion-model loader with GPU-resident ingest and residency tracking. " "It mirrors the KJ node's weight dtype, compute dtype, cublas_ops, SageAttention, fp16 accumulation, " "and optional extra-state-dict features." ) def patch_and_load( self, model_name: str, weight_dtype: str, compute_dtype: str, patch_cublaslinear: bool, sage_attention: str, enable_fp16_accumulation: bool, extra_state_dict: str | None = None, policy_override: str | None = None, ): 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 _apply_policy_override(policy_override): unet_path = folder_paths.get_full_path_or_raise("diffusion_models", model_name) model = _load_resident_diffusion_model( loader_name="Diffusion Model Loader Resident", cache_scope="diffusion_model", source_path=unet_path, note="diffusion model", weight_dtype=weight_dtype, compute_dtype=compute_dtype, patch_cublaslinear=patch_cublaslinear, sage_attention=sage_attention, enable_fp16_accumulation=enable_fp16_accumulation, extra_state_dict=extra_state_dict, policy_override=policy_override, ) return (model,) class CheckpointLoaderResident: @classmethod def INPUT_TYPES(cls): return { "required": { "ckpt_name": ( folder_paths.get_filename_list("checkpoints"), {"tooltip": "Checkpoint file to load."}, ), "weight_dtype": (["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp16", "bf16", "fp32"],), "compute_dtype": ( ["default", "fp16", "bf16", "fp32"], {"default": "default", "tooltip": "Compute dtype to apply to the diffusion model patcher after load."}, ), "patch_cublaslinear": ( "BOOLEAN", {"default": False, "tooltip": "Toggle ComfyUI's cublas_ops performance feature."}, ), "sage_attention": ( SAGE_ATTN_MODES, {"default": "disabled", "tooltip": "Patch optimized attention override on the loaded model."}, ), "enable_fp16_accumulation": ( "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") FUNCTION = "load" CATEGORY = "GPU Resident Loader/loaders" DESCRIPTION = ( "Checkpoint loader with the KJ DiffusionModelLoader-style tuning knobs plus GPU-resident ingest. " "It reuses live checkpoint components when possible and composes missing outputs from the model, CLIP, and VAE loaders instead of materializing a broad checkpoint state dict." ) def load( self, ckpt_name: str, weight_dtype: str, compute_dtype: str, patch_cublaslinear: bool, sage_attention: str, enable_fp16_accumulation: bool, policy_override: str | None = None, ): policy_override, _ = _resolve_loader_policy_and_extra_state_dict( loader_name="Checkpoint Loader Resident", policy_override=policy_override, extra_state_dict=None, ) with _apply_policy_override(policy_override): ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) model, clip, vae = _load_full_checkpoint( ckpt_path=ckpt_path, weight_dtype=weight_dtype, compute_dtype=compute_dtype, patch_cublaslinear=patch_cublaslinear, sage_attention=sage_attention, enable_fp16_accumulation=enable_fp16_accumulation, policy_override=policy_override, loader_name="Checkpoint Loader Resident", ) return model, clip, vae class CheckpointModelLoaderResident: @classmethod def INPUT_TYPES(cls): return CheckpointLoaderResident.INPUT_TYPES() RETURN_TYPES = ("MODEL",) FUNCTION = "load" CATEGORY = "GPU Resident Loader/loaders" DESCRIPTION = ( "Checkpoint model-only loader. Safetensors checkpoints take the selective UNet fast path, and live equivalent models are reused when available." ) def load( self, ckpt_name: str, weight_dtype: str, compute_dtype: str, patch_cublaslinear: bool, sage_attention: str, enable_fp16_accumulation: bool, policy_override: str | None = None, ): policy_override, _ = _resolve_loader_policy_and_extra_state_dict( loader_name="Checkpoint Model Loader Resident", policy_override=policy_override, extra_state_dict=None, ) with _apply_policy_override(policy_override): ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) model = _load_resident_diffusion_model( loader_name="Checkpoint Model Loader Resident", cache_scope="checkpoint_model", source_path=ckpt_path, note="checkpoint model", weight_dtype=weight_dtype, compute_dtype=compute_dtype, patch_cublaslinear=patch_cublaslinear, sage_attention=sage_attention, enable_fp16_accumulation=enable_fp16_accumulation, policy_override=policy_override, ) return (model,) class CheckpointClipLoaderResident: @classmethod def INPUT_TYPES(cls): return { "required": { "ckpt_name": ( folder_paths.get_filename_list("checkpoints"), {"tooltip": "Checkpoint file to load."}, ), }, "optional": { "policy_override": ( "STRING", { "forceInput": True, "tooltip": "Optional residency policy override. Connect Set Global Residency Policy here.", }, ) }, } RETURN_TYPES = ("CLIP",) FUNCTION = "load" CATEGORY = "GPU Resident Loader/loaders" DESCRIPTION = "Checkpoint CLIP-only loader with live-object reuse. It avoids rebuilding the text encoder when an equivalent CLIP object is already alive." def load(self, ckpt_name: str, policy_override: str | None = None): policy_override, _ = _resolve_loader_policy_and_extra_state_dict( loader_name="Checkpoint Clip Loader Resident", policy_override=policy_override, extra_state_dict=None, ) with _apply_policy_override(policy_override): ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) clip = _load_checkpoint_clip_only( ckpt_path=ckpt_path, policy_override=policy_override, loader_name="Checkpoint Clip Loader Resident", ) return (clip,) class CheckpointVAELoaderResident: @classmethod def INPUT_TYPES(cls): return CheckpointClipLoaderResident.INPUT_TYPES() RETURN_TYPES = ("VAE",) FUNCTION = "load" CATEGORY = "GPU Resident Loader/loaders" DESCRIPTION = "Checkpoint VAE-only loader with live-object reuse. It avoids rebuilding the VAE when an equivalent object is still alive." def load(self, ckpt_name: str, policy_override: str | None = None): policy_override, _ = _resolve_loader_policy_and_extra_state_dict( loader_name="Checkpoint VAE Loader Resident", policy_override=policy_override, extra_state_dict=None, ) with _apply_policy_override(policy_override): ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) vae = _load_checkpoint_vae_only( ckpt_path=ckpt_path, policy_override=policy_override, loader_name="Checkpoint VAE Loader Resident", ) return (vae,)