diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index 7e398f8..7f892e3 100644 --- a/checkpoint_multigpu.py +++ b/checkpoint_multigpu.py @@ -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)