Files
pollockjj-ComfyUI-MultiGPU/checkpoint_multigpu.py
T
John Pollock 0b1511edee refactor: Simplify checkpoint loading and fix text encoder device
This commit introduces two main improvements: refactoring the checkpoint loading mechanism and fixing the initial device placement for the text encoder (CLIP).

1.  **Fix Text Encoder Device Handling:**
    - A new patch is applied to `mm.text_encoder_initial_device` to gain control over the device used when the text encoder is first loaded.
    - The `CLIPLoader` override now forces `device='default'` to ensure ComfyUI's patching mechanism is triggered correctly, preventing the text encoder from being incorrectly placed on the wrong GPU.

2.  **Refactor Checkpoint Loaders:**
    - Removed the global stores (`checkpoint_dtype_store`, `checkpoint_half_store`, `checkpoint_config_store`).
    - The `CheckpointLoaderSimpleMultiGPU` and `AdvCheckpointLoaderMultiGPU` nodes now use arguments and ComfyUI's internal defaults directly. This simplifies the logic, reduces global state, and makes the code easier to follow.

Additionally, log message prefixes have been updated to be more descriptive, aiding in debugging.
2025-08-31 01:00:53 -05:00

279 lines
14 KiB
Python

"""
Advanced Checkpoint Loaders for MultiGPU
Provides device-specific and DisTorch2 sharding for checkpoint components
"""
import torch
import logging
import hashlib
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, safetensor_settings_store, create_safetensor_model_hash, register_patched_safetensor_modelpatcher
logger = logging.getLogger("MultiGPU")
checkpoint_device_config = {}
checkpoint_distorch_config = {}
original_load_state_dict_guess_config = None
def patch_load_state_dict_guess_config():
"""
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:
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
comfy.sd.load_state_dict_guess_config = 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
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:
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}")
clip = None
clipvision = None
vae = None
model = None
model_patcher = None
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:
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)
unet_compute_device = device_config.get('unet_device', original_main_device)
manual_cast_dtype = mm.unet_manual_cast(unet_dtype, torch.device(unet_compute_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}")
if model_config.clip_vision_prefix is not None and output_clipvision:
clipvision = comfy.clip_vision.load_clipvision_from_sd(sd, model_config.clip_vision_prefix, True)
if output_model:
unet_compute_device = device_config.get('unet_device', original_main_device)
set_current_device(unet_compute_device)
inital_load_device = mm.unet_inital_load_device(parameters, unet_dtype)
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)
if output_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
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)
vae = VAE(sd=vae_sd, metadata=metadata)
if output_clip:
clip_target_device = device_config.get('clip_device', original_clip_device)
set_current_text_encoder_device(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:
@classmethod
def INPUT_TYPES(s):
import folder_paths
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
return {
"required": {
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"unet_device": (devices, {"default": default_device}),
"clip_device": (devices, {"default": default_device}),
"vae_device": (devices, {"default": default_device}),
}
}
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
FUNCTION = "load_checkpoint"
CATEGORY = "multigpu"
TITLE = "Checkpoint Loader Advanced (MultiGPU)"
def load_checkpoint(self, ckpt_name, unet_device, clip_device, vae_device):
patch_load_state_dict_guess_config()
import folder_paths
import comfy.utils
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
sd = comfy.utils.load_torch_file(ckpt_path)
sd_size = sum(p.numel() for p in sd.values() if hasattr(p, 'numel'))
config_hash = str(sd_size)
checkpoint_device_config[config_hash] = {
'unet_device': unet_device, 'clip_device': clip_device, 'vae_device': vae_device
}
# Load using standard loader, our patch will intercept
from nodes import CheckpointLoaderSimple
return CheckpointLoaderSimple().load_checkpoint(ckpt_name)
class CheckpointLoaderAdvancedDisTorch2MultiGPU:
@classmethod
def INPUT_TYPES(s):
import folder_paths
devices = get_device_list()
compute_device = devices[1] if len(devices) > 1 else devices[0]
return {
"required": {
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"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": ("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": ("STRING", {"default": "cpu"}),
"vae_device": (devices, {"default": compute_device}),
}, "optional": {
"unet_expert_mode_allocations": ("STRING", {"multiline": False, "default": ""}),
"clip_expert_mode_allocations": ("STRING", {"multiline": False, "default": ""}),
"high_precision_loras": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
FUNCTION = "load_checkpoint"
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):
patch_load_state_dict_guess_config()
import folder_paths
import comfy.utils
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
sd = comfy.utils.load_torch_file(ckpt_path)
sd_size = sum(p.numel() for p in sd.values() if hasattr(p, 'numel'))
config_hash = str(sd_size)
checkpoint_device_config[config_hash] = {
'unet_device': unet_compute_device,
'clip_device': clip_compute_device,
'vae_device': vae_device
}
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_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(),
}
from nodes import CheckpointLoaderSimple
return CheckpointLoaderSimple().load_checkpoint(ckpt_name)