Fix loader state restoration and residency bookkeeping
This commit is contained in:
+58
-34
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
@@ -69,11 +70,32 @@ def _set_fp16_accumulation(enabled: bool) -> None:
|
||||
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)
|
||||
from sageattention import sageattn
|
||||
|
||||
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":
|
||||
@@ -307,24 +329,25 @@ class DiffusionModelLoaderResident:
|
||||
enable_fp16_accumulation: bool,
|
||||
extra_state_dict: str | None = None,
|
||||
):
|
||||
_set_cublas_linear(patch_cublaslinear)
|
||||
_set_fp16_accumulation(enable_fp16_accumulation)
|
||||
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)
|
||||
|
||||
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)
|
||||
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 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)
|
||||
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,)
|
||||
|
||||
@@ -375,25 +398,26 @@ class CheckpointLoaderResident:
|
||||
sage_attention: str,
|
||||
enable_fp16_accumulation: bool,
|
||||
):
|
||||
_set_cublas_linear(patch_cublaslinear)
|
||||
_set_fp16_accumulation(enable_fp16_accumulation)
|
||||
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_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)
|
||||
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 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)
|
||||
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")
|
||||
|
||||
+8
-3
@@ -49,8 +49,13 @@ def _safe_open_device_arg(device: torch.device) -> Any:
|
||||
return device.type
|
||||
|
||||
|
||||
def _copy_tensor_if_needed(tensor: torch.Tensor, target_device: torch.device) -> torch.Tensor:
|
||||
if tensor.device == target_device:
|
||||
def _copy_tensor_if_needed(
|
||||
tensor: torch.Tensor,
|
||||
target_device: torch.device,
|
||||
*,
|
||||
force_copy: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if tensor.device == target_device and not force_copy:
|
||||
return tensor
|
||||
return tensor.to(device=target_device, copy=True)
|
||||
|
||||
@@ -119,7 +124,7 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata
|
||||
for key in handle.keys():
|
||||
tensor = handle.get_tensor(key)
|
||||
if getattr(comfy_utils, "DISABLE_MMAP", False) and tensor.device.type == "cpu":
|
||||
tensor = _copy_tensor_if_needed(tensor, requested_device)
|
||||
tensor = _copy_tensor_if_needed(tensor, requested_device, force_copy=True)
|
||||
sd[key] = tensor
|
||||
if return_metadata:
|
||||
metadata = handle.metadata()
|
||||
|
||||
@@ -255,9 +255,11 @@ class ResidencyRegistry:
|
||||
raise ValueError("Cannot bind None into residency registry")
|
||||
|
||||
with self._lock:
|
||||
old_key: tuple[str, str] | None = None
|
||||
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
|
||||
if entry_id is not None and entry_id in self._entries:
|
||||
entry = self._entries[entry_id]
|
||||
old_key = (entry.kind, entry.source_path)
|
||||
else:
|
||||
entry_id = self._make_entry_id(kind, source_path)
|
||||
entry = ResidencyEntry(
|
||||
@@ -290,6 +292,11 @@ class ResidencyRegistry:
|
||||
entry.priority = int(priority)
|
||||
entry.source_path = source_path
|
||||
entry.kind = kind
|
||||
new_key = (entry.kind, entry.source_path)
|
||||
if old_key is not None and old_key != new_key:
|
||||
if self._path_to_entry.get(old_key) == entry_id:
|
||||
self._path_to_entry.pop(old_key, None)
|
||||
self._path_to_entry[new_key] = entry_id
|
||||
entry.last_touched = _now()
|
||||
if note:
|
||||
entry.notes.append(note)
|
||||
@@ -358,12 +365,14 @@ class ResidencyRegistry:
|
||||
obj = entry.object()
|
||||
if obj is None:
|
||||
continue
|
||||
entry.loaded_bytes = 0
|
||||
load_device = getattr(obj, "load_device", None)
|
||||
offload_device = getattr(obj, "offload_device", None)
|
||||
if load_device is not None:
|
||||
entry.load_device = str(load_device)
|
||||
if offload_device is not None:
|
||||
entry.offload_device = str(offload_device)
|
||||
entry.current_device = entry.offload_device
|
||||
|
||||
for loaded in list(model_management.current_loaded_models):
|
||||
entry = self.entry_for_object(loaded.model)
|
||||
|
||||
Reference in New Issue
Block a user