diff --git a/__init__.py b/__init__.py index 751ec1c..2f5efd4 100644 --- a/__init__.py +++ b/__init__.py @@ -168,6 +168,21 @@ def set_current_unet_offload_device(device): current_unet_offload_device = device logger.debug(f"[MultiGPU Initialization] current_unet_offload_device set to: {device}") + +def get_current_device(): + """Get the current device context for MultiGPU operations at runtime.""" + return current_device + + +def get_current_text_encoder_device(): + """Get the current text encoder device context for CLIP models at runtime.""" + return current_text_encoder_device + + +def get_current_unet_offload_device(): + """Get the current UNet offload device context at runtime.""" + return current_unet_offload_device + def get_torch_device_patched(): """Return MultiGPU-aware device selection for patched mm.get_torch_device.""" device = None diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index 6b5f7ef..973a78e 100644 --- a/checkpoint_multigpu.py +++ b/checkpoint_multigpu.py @@ -34,7 +34,7 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, embedding_directory=None, output_model=True, model_options={}, te_model_options={}, metadata=None): """Patched checkpoint loader with MultiGPU and DisTorch2 device placement support.""" - from . import set_current_device, set_current_text_encoder_device, current_device, current_text_encoder_device + from . import set_current_device, set_current_text_encoder_device, get_current_device, get_current_text_encoder_device sd_size = sum(p.numel() for p in sd.values() if hasattr(p, 'numel')) config_hash = str(sd_size) @@ -54,8 +54,9 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, model = None model_patcher = None - original_main_device = current_device - original_clip_device = current_text_encoder_device + # Capture the current devices at runtime so we can restore them after loading + original_main_device = get_current_device() + original_clip_device = get_current_text_encoder_device() try: diffusion_model_prefix = comfy.model_detection.unet_prefix_from_state_dict(sd) diff --git a/wrappers.py b/wrappers.py index 30e85c6..a132792 100644 --- a/wrappers.py +++ b/wrappers.py @@ -59,6 +59,10 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette device_value = kwargs.get(device_param_name) + # Capture the current device at runtime so we can restore it later + from . import get_current_device, get_current_text_encoder_device + original_device = get_current_device() if device_param_name == "compute_device" else get_current_text_encoder_device() + import comfy.model_management as mm if eject_models: @@ -118,7 +122,11 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}") - return out + try: + return out + finally: + # Restore the device that was in use when the override started + device_setter_func(original_device) return NodeOverrideDisTorchSafetensorV2 @@ -164,7 +172,7 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls): def override_class_with_distorch_gguf(cls): """DisTorch V1 Legacy wrapper - maintains V1 UI but calls V2 backend""" - from . import set_current_device + from . import set_current_device, get_current_device from .distorch_2 import register_patched_safetensor_modelpatcher class NodeOverrideDisTorchGGUFLegacy(cls): @@ -185,6 +193,8 @@ def override_class_with_distorch_gguf(cls): TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (Legacy)" def override(self, *args, device=None, expert_mode_allocations="", use_other_vram=False, virtual_vram_gb=0.0, **kwargs): + # Capture and restore the current device to avoid leaking global state + original_device = get_current_device() if device is not None: set_current_device(device) @@ -220,15 +230,17 @@ def override_class_with_distorch_gguf(cls): if model_to_check and full_allocation: inner_model = model_to_check.model inner_model._distorch_v2_meta = {"full_allocation": full_allocation} - - return out + try: + return out + finally: + set_current_device(original_device) return NodeOverrideDisTorchGGUFLegacy def override_class_with_distorch_gguf_v2(cls): """DisTorch V2 wrapper for GGUF models""" - from . import set_current_device + from . import set_current_device, get_current_device from .distorch_2 import register_patched_safetensor_modelpatcher class NodeOverrideDisTorchGGUFv2(cls): @@ -250,6 +262,7 @@ def override_class_with_distorch_gguf_v2(cls): TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (DisTorch2)" def override(self, *args, compute_device=None, virtual_vram_gb=4.0, donor_device="cpu", expert_mode_allocations="", **kwargs): + original_device = get_current_device() if compute_device is not None: set_current_device(compute_device) @@ -282,15 +295,17 @@ def override_class_with_distorch_gguf_v2(cls): if model_to_check and full_allocation: inner_model = model_to_check.model inner_model._distorch_v2_meta = {"full_allocation": full_allocation} - - return out + try: + return out + finally: + set_current_device(original_device) return NodeOverrideDisTorchGGUFv2 def override_class_with_distorch_clip(cls): """DisTorch V1 wrapper for CLIP models - calls V2 backend""" - from . import set_current_text_encoder_device + from . import set_current_text_encoder_device, get_current_text_encoder_device from .distorch_2 import register_patched_safetensor_modelpatcher class NodeOverrideDisTorchClip(cls): @@ -311,6 +326,7 @@ def override_class_with_distorch_clip(cls): TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (DisTorch)" def override(self, *args, device=None, expert_mode_allocations="", use_other_vram=False, virtual_vram_gb=0.0, **kwargs): + original_text_device = get_current_text_encoder_device() if device is not None: set_current_text_encoder_device(device) @@ -346,15 +362,17 @@ def override_class_with_distorch_clip(cls): if model_to_check and full_allocation: inner_model = model_to_check.model inner_model._distorch_v2_meta = {"full_allocation": full_allocation} - - return out + try: + return out + finally: + set_current_text_encoder_device(original_text_device) return NodeOverrideDisTorchClip def override_class_with_distorch_clip_no_device(cls): """DisTorch V1 wrapper for Triple/Quad CLIP models - calls V2 backend""" - from . import set_current_text_encoder_device + from . import set_current_text_encoder_device, get_current_text_encoder_device from .distorch_2 import register_patched_safetensor_modelpatcher class NodeOverrideDisTorchClipNoDevice(cls): @@ -375,6 +393,7 @@ def override_class_with_distorch_clip_no_device(cls): TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (DisTorch)" def override(self, *args, device=None, expert_mode_allocations="", use_other_vram=False, virtual_vram_gb=0.0, **kwargs): + original_text_device = get_current_text_encoder_device() if device is not None: set_current_text_encoder_device(device) @@ -410,8 +429,10 @@ def override_class_with_distorch_clip_no_device(cls): if model_to_check and full_allocation: inner_model = model_to_check.model inner_model._distorch_v2_meta = {"full_allocation": full_allocation} - - return out + try: + return out + finally: + set_current_text_encoder_device(original_text_device) return NodeOverrideDisTorchClipNoDevice @@ -426,7 +447,7 @@ override_class_with_distorch = override_class_with_distorch_gguf def override_class(cls): """Standard MultiGPU device override for UNet/VAE models""" - from . import set_current_device + from . import set_current_device, get_current_device class NodeOverride(cls): @classmethod @@ -442,17 +463,21 @@ def override_class(cls): FUNCTION = "override" def override(self, *args, device=None, **kwargs): + original_device = get_current_device() if device is not None: set_current_device(device) fn = getattr(super(), cls.FUNCTION) out = fn(*args, **kwargs) - return out + try: + return out + finally: + set_current_device(original_device) return NodeOverride def override_class_offload(cls): """Standard MultiGPU device override for UNet/VAE models""" - from . import set_current_device, set_current_unet_offload_device + from . import set_current_device, set_current_unet_offload_device, get_current_device, get_current_unet_offload_device class NodeOverride(cls): @classmethod @@ -469,13 +494,19 @@ def override_class_offload(cls): FUNCTION = "override" def override(self, *args, device=None, offload_device=None, **kwargs): + original_device = get_current_device() + original_offload_device = get_current_unet_offload_device() if device is not None: set_current_device(device) if offload_device is not None: set_current_unet_offload_device(offload_device) fn = getattr(super(), cls.FUNCTION) out = fn(*args, **kwargs) - return out + try: + return out + finally: + set_current_device(original_device) + set_current_unet_offload_device(original_offload_device) return NodeOverride @@ -483,7 +514,7 @@ def override_class_offload(cls): def override_class_clip(cls): """Standard MultiGPU device override for CLIP models (with device kwarg workaround)""" - from . import set_current_text_encoder_device + from . import set_current_text_encoder_device, get_current_text_encoder_device class NodeOverride(cls): @classmethod @@ -499,19 +530,23 @@ def override_class_clip(cls): FUNCTION = "override" def override(self, *args, device=None, **kwargs): + original_text_device = get_current_text_encoder_device() if device is not None: set_current_text_encoder_device(device) kwargs['device'] = 'default' fn = getattr(super(), cls.FUNCTION) out = fn(*args, **kwargs) - return out + try: + return out + finally: + set_current_text_encoder_device(original_text_device) return NodeOverride def override_class_clip_no_device(cls): """Standard MultiGPU device override for Triple/Quad CLIP models (no device kwarg workaround)""" - from . import set_current_text_encoder_device + from . import set_current_text_encoder_device, get_current_text_encoder_device class NodeOverride(cls): @classmethod @@ -527,10 +562,14 @@ def override_class_clip_no_device(cls): FUNCTION = "override" def override(self, *args, device=None, **kwargs): + original_text_device = get_current_text_encoder_device() if device is not None: set_current_text_encoder_device(device) fn = getattr(super(), cls.FUNCTION) out = fn(*args, **kwargs) - return out + try: + return out + finally: + set_current_text_encoder_device(original_text_device) return NodeOverride