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.
279 lines
14 KiB
Python
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)
|