fix: Overhaul checkpoint loader for proper device handling

This commit is contained in:
John Pollock
2025-08-30 19:51:15 -05:00
parent f07c2d2b89
commit 9e14e4622c
+191 -201
View File
@@ -6,137 +6,194 @@ Provides device-specific and DisTorch2 sharding for checkpoint components
import torch
import logging
import hashlib
import copy
import comfy.sd
import comfy.utils
import comfy.model_management as mm
import comfy.model_detection
import comfy.clip_vision
from comfy.sd import VAE, CLIP
from .device_utils import get_device_list
from .distorch_2 import safetensor_allocation_store, create_safetensor_model_hash
from .distorch_2 import safetensor_allocation_store, safetensor_settings_store, create_safetensor_model_hash, register_patched_safetensor_modelpatcher
logger = logging.getLogger("MultiGPU")
# Store checkpoint loading configurations
# --- Global Stores for Configuration ---
checkpoint_device_config = {}
checkpoint_distorch_config = {}
# Store the original function
# --- Original Function Store ---
original_load_state_dict_guess_config = None
def create_checkpoint_config_hash(checkpoint_name, config_str):
"""Create a unique hash for checkpoint configuration"""
identifier = f"{checkpoint_name}_{config_str}"
return hashlib.sha256(identifier.encode()).hexdigest()
def patch_load_state_dict_guess_config():
"""Monkey patch the load_state_dict_guess_config function to support per-component device selection"""
"""
Monkey patch the load_state_dict_guess_config function to replace its logic
with a MultiGPU-aware implementation.
"""
global original_load_state_dict_guess_config
if original_load_state_dict_guess_config is not None:
return # Already patched
logger.info("[MultiGPU] load_state_dict_guess_config is already patched.")
return
logger.info("[MultiGPU] Patching comfy.sd.load_state_dict_guess_config for advanced MultiGPU loading.")
original_load_state_dict_guess_config = comfy.sd.load_state_dict_guess_config
def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_clipvision=False,
embedding_directory=None, output_model=True, model_options={},
te_model_options={}, metadata=None):
# Import here to avoid circular imports
from . import set_current_device, set_current_text_encoder_device, current_device, current_text_encoder_device
# Check if we have a device configuration for this checkpoint
# We use the state dict size as a simple identifier
sd_size = sum(t.numel() for t in sd.values() if hasattr(t, 'numel'))
config_hash = str(sd_size)
device_config = checkpoint_device_config.get(config_hash)
distorch_config = checkpoint_distorch_config.get(config_hash)
if device_config or distorch_config:
logger.info(f"[MultiGPU] Using custom device configuration for checkpoint")
# Save original devices
original_unet_device = current_device
original_clip_device = current_text_encoder_device
# Handle UNet device/DisTorch config
if device_config and 'unet_device' in device_config:
set_current_device(device_config['unet_device'])
logger.info(f"[MultiGPU] Setting UNet device to: {device_config['unet_device']}")
# Apply DisTorch2 config for UNet if present
if distorch_config and 'unet_allocation' in distorch_config:
# We'll store this for when the model patcher is created
logger.info(f"[MultiGPU] DisTorch2 UNet allocation will be applied: {distorch_config['unet_allocation']}")
# Call original function to load the checkpoint
result = original_load_state_dict_guess_config(
sd, output_vae=output_vae, output_clip=output_clip, output_clipvision=output_clipvision,
embedding_directory=embedding_directory, output_model=output_model,
model_options=model_options, te_model_options=te_model_options, metadata=metadata
)
model_patcher, clip, vae, clipvision = result
# Apply DisTorch2 configurations after loading
if distorch_config:
if model_patcher and 'unet_allocation' in distorch_config:
model_hash = create_safetensor_model_hash(model_patcher, "checkpoint_loader")
safetensor_allocation_store[model_hash] = distorch_config['unet_allocation']
if 'unet_settings' in distorch_config:
from .distorch_2 import safetensor_settings_store
safetensor_settings_store[model_hash] = distorch_config['unet_settings']
logger.info(f"[MultiGPU] Applied DisTorch2 config to UNet: {model_hash[:8]}")
if clip and 'clip_allocation' in distorch_config:
# For CLIP, we need to get the model from the CLIP object
if hasattr(clip, 'patcher'):
clip_hash = create_safetensor_model_hash(clip.patcher, "checkpoint_loader_clip")
safetensor_allocation_store[clip_hash] = distorch_config['clip_allocation']
if 'clip_settings' in distorch_config:
from .distorch_2 import safetensor_settings_store
safetensor_settings_store[clip_hash] = distorch_config['clip_settings']
logger.info(f"[MultiGPU] Applied DisTorch2 config to CLIP: {clip_hash[:8]}")
# Handle CLIP device
if device_config and 'clip_device' in device_config and clip:
set_current_text_encoder_device(device_config['clip_device'])
logger.info(f"[MultiGPU] Setting CLIP device to: {device_config['clip_device']}")
# Force CLIP to load on the specified device
if hasattr(clip, 'patcher'):
clip.patcher.load(force_patch_weights=True)
# Handle VAE device
if device_config and 'vae_device' in device_config and vae:
vae_device = torch.device(device_config['vae_device'])
logger.info(f"[MultiGPU] Setting VAE device to: {device_config['vae_device']}")
# Move VAE to specified device
if hasattr(vae, 'first_stage_model'):
vae.first_stage_model = vae.first_stage_model.to(vae_device)
# Clean up stored configs
if config_hash in checkpoint_device_config:
del checkpoint_device_config[config_hash]
if config_hash in checkpoint_distorch_config:
del checkpoint_distorch_config[config_hash]
return result
else:
# No custom config, use original behavior
return original_load_state_dict_guess_config(
sd, output_vae=output_vae, output_clip=output_clip, output_clipvision=output_clipvision,
embedding_directory=embedding_directory, output_model=output_model,
model_options=model_options, te_model_options=te_model_options, metadata=metadata
)
# Apply the patch
comfy.sd.load_state_dict_guess_config = patched_load_state_dict_guess_config
logger.info("[MultiGPU] Successfully patched load_state_dict_guess_config")
def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_clipvision=False,
embedding_directory=None, output_model=True, model_options={},
te_model_options={}, metadata=None):
from . import set_current_device, set_current_text_encoder_device, current_device, current_text_encoder_device
# --- Check for custom configuration ---
sd_size = sum(p.numel() for p in sd.values() if hasattr(p, 'numel'))
config_hash = str(sd_size)
device_config = checkpoint_device_config.get(config_hash)
distorch_config = checkpoint_distorch_config.get(config_hash)
if not device_config and not distorch_config:
# No config, fall back to original untouched function
return original_load_state_dict_guess_config(sd, output_vae, output_clip, output_clipvision, embedding_directory, output_model, model_options, te_model_options, metadata)
logger.info("--- [MultiGPU] ENTERING Patched Checkpoint Loader ---")
logger.info(f"Received Device Config: {device_config}")
logger.info(f"Received DisTorch2 Config: {distorch_config}")
# --- Start of Rewritten Logic ---
clip = None
clipvision = None
vae = None
model = None
model_patcher = None
# Store original device contexts to restore later
original_main_device = current_device
original_clip_device = current_text_encoder_device
logger.info(f"Saved original device contexts: UNet/VAE='{original_main_device}', CLIP='{original_clip_device}'")
try:
# --- Model Configuration Detection (Replicated from original) ---
diffusion_model_prefix = comfy.model_detection.unet_prefix_from_state_dict(sd)
parameters = comfy.utils.calculate_parameters(sd, diffusion_model_prefix)
weight_dtype = comfy.utils.weight_dtype(sd, diffusion_model_prefix)
model_config = comfy.model_detection.model_config_from_unet(sd, diffusion_model_prefix, metadata=metadata)
if model_config is None:
logger.warning("[MultiGPU] Warning: Not a standard checkpoint file. Trying to load as diffusion model only.")
# Simplified fallback for non-checkpoints
set_current_device(device_config.get('unet_device', original_main_device))
diffusion_model = comfy.sd.load_diffusion_model_state_dict(sd, model_options={})
if diffusion_model is None:
return None
return (diffusion_model, None, VAE(sd={}), None)
logger.info(f"[MultiGPU] Detected Model Config: {type(model_config).__name__}, Parameters: {parameters/10**9:.2f}B")
unet_weight_dtype = list(model_config.supported_inference_dtypes)
if model_config.scaled_fp8 is not None:
weight_dtype = None
model_config.custom_operations = model_options.get("custom_operations", None)
unet_dtype = model_options.get("dtype", model_options.get("weight_dtype", None)))
if unet_dtype is None:
unet_dtype = mm.unet_dtype(model_params=parameters, supported_dtypes=unet_weight_dtype, weight_dtype=weight_dtype)
manual_cast_dtype = mm.unet_manual_cast(unet_dtype, torch.device(device_config.get('unet_device')), model_config.supported_inference_dtypes)
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype)
logger.info(f"UNet DType: {unet_dtype}, Manual Cast: {manual_cast_dtype}")
# --- CLIP Vision Loading ---
if model_config.clip_vision_prefix is not None and output_clipvision:
logger.info("--- Loading CLIP Vision ---")
clipvision = comfy.clip_vision.load_clipvision_from_sd(sd, model_config.clip_vision_prefix, True)
logger.info("CLIP Vision Loaded.")
# --- UNet Loading Block ---
if output_model:
logger.info("--- Loading UNet ---")
unet_compute_device = device_config.get('unet_device', original_main_device)
set_current_device(unet_compute_device)
logger.info(f"Set UNet context to: {unet_compute_device}")
inital_load_device = mm.unet_inital_load_device(parameters, unet_dtype)
logger.info(f"UNet initial load device: {inital_load_device}")
model = model_config.get_model(sd, diffusion_model_prefix, device=inital_load_device)
model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=unet_compute_device, offload_device=mm.unet_offload_device())
if distorch_config and 'unet_allocation' in distorch_config:
register_patched_safetensor_modelpatcher()
model_hash = create_safetensor_model_hash(model_patcher, "checkpoint_loader_unet")
safetensor_allocation_store[model_hash] = distorch_config['unet_allocation']
safetensor_settings_store[model_hash] = distorch_config.get('unet_settings','')
model.is_distorch = True
model._distorch_high_precision_loras = distorch_config.get('high_precision_loras', True)
logger.info(f"Stored DisTorch2 config for UNet (hash {model_hash[:8]}): {distorch_config['unet_allocation']}")
model.load_model_weights(sd, diffusion_model_prefix)
logger.info("UNet Loaded.")
# --- VAE Loading Block ---
if output_vae:
logger.info("--- Loading VAE ---")
vae_target_device = torch.device(device_config.get('vae_device', original_main_device))
set_current_device(vae_target_device) # Use main device context for VAE
logger.info(f"Set VAE context to: {vae_target_device}")
vae_sd = comfy.utils.state_dict_prefix_replace(sd, {k: "" for k in model_config.vae_key_prefix}, filter_keys=True)
vae_sd = model_config.process_vae_state_dict(vae_sd)
# The VAE class itself respects the mm.get_torch_device() patch
vae = VAE(sd=vae_sd, metadata=metadata)
logger.info(f"VAE Loaded. Final device should be: {vae_target_device}")
# --- CLIP Loading Block ---
if output_clip:
logger.info("--- Loading CLIP ---")
clip_target_device = device_config.get('clip_device', original_clip_device)
set_current_text_encoder_device(clip_target_device)
logger.info(f"Set CLIP context to: {clip_target_device}")
clip_target = model_config.clip_target(state_dict=sd)
if clip_target is not None:
clip_sd = model_config.process_clip_state_dict(sd)
if len(clip_sd) > 0:
clip_params = comfy.utils.calculate_parameters(clip_sd)
clip = CLIP(clip_target, embedding_directory=embedding_directory, tokenizer_data=clip_sd, parameters=clip_params, model_options=te_model_options)
if distorch_config and 'clip_allocation' in distorch_config:
if hasattr(clip, 'patcher'):
register_patched_safetensor_modelpatcher()
clip_hash = create_safetensor_model_hash(clip.patcher, "checkpoint_loader_clip")
safetensor_allocation_store[clip_hash] = distorch_config['clip_allocation']
safetensor_settings_store[clip_hash] = distorch_config.get('clip_settings','')
clip.patcher.model.is_distorch = True
clip.patcher.model._distorch_high_precision_loras = distorch_config.get('high_precision_loras', True)
logger.info(f"Stored DisTorch2 config for CLIP (hash {clip_hash[:8]}): {distorch_config['clip_allocation']}")
m, u = clip.load_sd(clip_sd, full_model=True) # This respects the patched text_encoder_device
if len(m) > 0: logger.warning(f"CLIP missing keys: {m}")
if len(u) > 0: logger.debug(f"CLIP unexpected keys: {u}")
logger.info("CLIP Loaded.")
else:
logger.warning("No CLIP/text encoder weights in checkpoint.")
else:
logger.warning("CLIP target not found in model config.")
finally:
# --- Restore original device contexts and clean up ---
set_current_device(original_main_device)
set_current_text_encoder_device(original_clip_device)
if config_hash in checkpoint_device_config:
del checkpoint_device_config[config_hash]
if config_hash in checkpoint_distorch_config:
del checkpoint_distorch_config[config_hash]
logger.info(f"Restored original device contexts. UNet/VAE='{original_main_device}', CLIP='{original_clip_device}'")
logger.info("--- [MultiGPU] EXITING Patched Checkpoint Loader ---")
return (model_patcher, clip, vae, clipvision)
class CheckpointLoaderAdvancedMultiGPU:
"""
Checkpoint loader that allows loading UNet, CLIP, and VAE to different devices
"""
@classmethod
def INPUT_TYPES(s):
import folder_paths
@@ -158,39 +215,26 @@ class CheckpointLoaderAdvancedMultiGPU:
TITLE = "Checkpoint Loader Advanced (MultiGPU)"
def load_checkpoint(self, ckpt_name, unet_device, clip_device, vae_device):
# Apply the patch if not already applied
patch_load_state_dict_guess_config()
# Store device configuration
import folder_paths
import comfy.utils
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
sd = comfy.utils.load_torch_file(ckpt_path)
# Use state dict size as identifier
sd_size = sum(t.numel() for t in sd.values() if hasattr(t, 'numel'))
sd_size = sum(p.numel() for p in sd.values() if hasattr(p, 'numel'))
config_hash = str(sd_size)
# Store the device configuration
checkpoint_device_config[config_hash] = {
'unet_device': unet_device,
'clip_device': clip_device,
'vae_device': vae_device
'unet_device': unet_device, 'clip_device': clip_device, 'vae_device': vae_device
}
logger.info(f"[MultiGPU] CheckpointLoaderAdvanced configured - UNet: {unet_device}, CLIP: {clip_device}, VAE: {vae_device}")
# Load the checkpoint - our patched function will handle device placement
# Load using standard loader, our patch will intercept
from nodes import CheckpointLoaderSimple
loader = CheckpointLoaderSimple()
return loader.load_checkpoint(ckpt_name)
return CheckpointLoaderSimple().load_checkpoint(ckpt_name)
class CheckpointLoaderAdvancedDisTorch2MultiGPU:
"""
Checkpoint loader with full DisTorch2 sharding for UNet and CLIP, device selection for VAE
"""
@classmethod
def INPUT_TYPES(s):
import folder_paths
@@ -200,18 +244,14 @@ class CheckpointLoaderAdvancedDisTorch2MultiGPU:
return {
"required": {
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
# UNet DisTorch2 settings
"unet_compute_device": (devices, {"default": compute_device}),
"unet_virtual_vram_gb": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 128.0, "step": 0.1}),
"unet_donor_device": (devices, {"default": "cpu"}),
# CLIP DisTorch2 settings
"clip_compute_device": (devices, {"default": compute_device}),
"unet_donor_device": ("STRING", {"default": "cpu"}),
"clip_compute_device": (devices, {"default": "cpu"}),
"clip_virtual_vram_gb": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 128.0, "step": 0.1}),
"clip_donor_device": (devices, {"default": "cpu"}),
# VAE simple device
"clip_donor_device": ("STRING", {"default": "cpu"}),
"vae_device": (devices, {"default": compute_device}),
},
"optional": {
}, "optional": {
"unet_expert_mode_allocations": ("STRING", {"multiline": False, "default": ""}),
"clip_expert_mode_allocations": ("STRING", {"multiline": False, "default": ""}),
"high_precision_loras": ("BOOLEAN", {"default": True}),
@@ -223,88 +263,38 @@ class CheckpointLoaderAdvancedDisTorch2MultiGPU:
CATEGORY = "multigpu/distorch_2"
TITLE = "Checkpoint Loader Advanced (DisTorch2)"
def load_checkpoint(self, ckpt_name,
unet_compute_device, unet_virtual_vram_gb, unet_donor_device,
clip_compute_device, clip_virtual_vram_gb, clip_donor_device,
vae_device,
unet_expert_mode_allocations="", clip_expert_mode_allocations="",
high_precision_loras=True):
def load_checkpoint(self, ckpt_name, unet_compute_device, unet_virtual_vram_gb, unet_donor_device,
clip_compute_device, clip_virtual_vram_gb, clip_donor_device, vae_device,
unet_expert_mode_allocations="", clip_expert_mode_allocations="", high_precision_loras=True):
# Apply the patch if not already applied
patch_load_state_dict_guess_config()
patch_load_state_dict_guess_config()
# Register DisTorch2 model patcher
from .distorch_2 import register_patched_safetensor_modelpatcher
register_patched_safetensor_modelpatcher()
# Store device configuration
import folder_paths
import comfy.utils
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
sd = comfy.utils.load_torch_file(ckpt_path)
# Use state dict size as identifier
sd_size = sum(t.numel() for t in sd.values() if hasattr(t, 'numel'))
sd_size = sum(p.numel() for p in sd.values() if hasattr(p, 'numel'))
config_hash = str(sd_size)
# Store device configuration
checkpoint_device_config[config_hash] = {
'unet_device': unet_compute_device,
'clip_device': clip_compute_device,
'vae_device': vae_device
}
# Build DisTorch2 allocation strings
unet_vram_string = ""
if unet_virtual_vram_gb > 0:
unet_vram_string = f"{unet_compute_device};{unet_virtual_vram_gb};{unet_donor_device}"
elif unet_expert_mode_allocations:
unet_vram_string = unet_compute_device
unet_allocation = f"{unet_expert_mode_allocations}#{unet_vram_string}" if unet_expert_mode_allocations or unet_vram_string else ""
clip_vram_string = ""
if clip_virtual_vram_gb > 0:
clip_vram_string = f"{clip_compute_device};{clip_virtual_vram_gb};{clip_donor_device}"
elif clip_expert_mode_allocations:
clip_vram_string = clip_compute_device
clip_allocation = f"{clip_expert_mode_allocations}#{clip_vram_string}" if clip_expert_mode_allocations or clip_vram_string else ""
# Create settings hashes for DisTorch2
unet_settings_str = f"{unet_compute_device}{unet_virtual_vram_gb}{unet_donor_device}{unet_expert_mode_allocations}{high_precision_loras}"
unet_settings_hash = hashlib.sha256(unet_settings_str.encode()).hexdigest()
clip_settings_str = f"{clip_compute_device}{clip_virtual_vram_gb}{clip_donor_device}{clip_expert_mode_allocations}{high_precision_loras}"
clip_settings_hash = hashlib.sha256(clip_settings_str.encode()).hexdigest()
# Store DisTorch2 configuration
unet_vram_str = f"{unet_compute_device};{unet_virtual_vram_gb};{unet_donor_device}"
unet_alloc = f"{unet_expert_mode_allocations}#{unet_vram_str}"
clip_vram_str = f"{clip_compute_device};{clip_virtual_vram_gb};{clip_donor_device}"
clip_alloc = f"{clip_expert_mode_allocations}#{clip_vram_str}"
checkpoint_distorch_config[config_hash] = {
'unet_allocation': unet_allocation,
'unet_settings': unet_settings_hash,
'clip_allocation': clip_allocation,
'clip_settings': clip_settings_hash,
'high_precision_loras': high_precision_loras
'unet_allocation': unet_alloc,
'clip_allocation': clip_alloc,
'high_precision_loras': high_precision_loras,
'unet_settings': hashlib.sha256(f"{unet_alloc}{high_precision_loras}".encode()).hexdigest(),
'clip_settings': hashlib.sha256(f"{clip_alloc}{high_precision_loras}".encode()).hexdigest(),
}
logger.info(f"[MultiGPU] CheckpointLoaderDisTorch2 configured:")
logger.info(f" UNet: compute={unet_compute_device}, vram={unet_virtual_vram_gb}GB, donor={unet_donor_device}")
logger.info(f" CLIP: compute={clip_compute_device}, vram={clip_virtual_vram_gb}GB, donor={clip_donor_device}")
logger.info(f" VAE: device={vae_device}")
# Load the checkpoint - our patched function will handle device placement and DisTorch2
from nodes import CheckpointLoaderSimple
loader = CheckpointLoaderSimple()
# Set high precision loras flag
result = loader.load_checkpoint(ckpt_name)
# Store high_precision_loras in the models
model_patcher, clip, vae = result
if model_patcher and hasattr(model_patcher, 'model'):
model_patcher.model._distorch_high_precision_loras = high_precision_loras
if clip and hasattr(clip, 'patcher') and hasattr(clip.patcher, 'model'):
clip.patcher.model._distorch_high_precision_loras = high_precision_loras
return result
return CheckpointLoaderSimple().load_checkpoint(ckpt_name)