From 24e4e17a97ca1992c4bcf21955ed353cef13c65c Mon Sep 17 00:00:00 2001 From: John Pollock Date: Fri, 6 Mar 2026 05:19:03 -0600 Subject: [PATCH] feat: enhance MultiGPU support with device guards and runtime management --- __init__.py | 169 ++++++++++++++++++++++++++++++++++++++++++++++++---- wrappers.py | 20 +++++-- 2 files changed, 170 insertions(+), 19 deletions(-) diff --git a/__init__.py b/__init__.py index b2c4148..14abdbd 100644 --- a/__init__.py +++ b/__init__.py @@ -5,11 +5,14 @@ import os import copy import json import importlib +from contextlib import contextmanager from datetime import datetime from pathlib import Path import folder_paths import comfy.model_management as mm +import comfy.memory_management import comfy.model_patcher +import comfy.sample as comfy_sample from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS from .device_utils import ( get_device_list, @@ -171,6 +174,9 @@ def check_module_exists(module_path): current_device = mm.get_torch_device() current_text_encoder_device = mm.text_encoder_device() current_unet_offload_device = mm.unet_offload_device() +_aimdo_initialized_devices = set() +if isinstance(current_device, torch.device) and current_device.type == "cuda" and current_device.index is not None: + _aimdo_initialized_devices.add(current_device.index) def set_current_device(device): """Set the current device context for MultiGPU operations.""" @@ -205,6 +211,73 @@ def get_current_unet_offload_device(): """Get the current UNet offload device context at runtime.""" return current_unet_offload_device +def _coerce_torch_device(device): + """Best-effort conversion to torch.device for guard and patch helpers.""" + if device is None: + return None + if isinstance(device, torch.device): + return device + try: + return torch.device(device) + except (TypeError, RuntimeError, ValueError): + return None + +@contextmanager +def cuda_device_guard(device, reason="runtime"): + """Temporarily switch the real CUDA current device for non-primary execution paths.""" + target_device = _coerce_torch_device(device) + previous_device_index = None + switched_device = False + + if ( + target_device is not None + and target_device.type == "cuda" + and target_device.index is not None + and torch.cuda.is_available() + ): + previous_device_index = torch.cuda.current_device() + if previous_device_index != target_device.index: + logger.info( + f"[MultiGPU CUDA Guard] Switching CUDA current device {previous_device_index} -> {target_device.index} ({reason})" + ) + torch.cuda.set_device(target_device.index) + switched_device = True + + try: + yield target_device + finally: + if switched_device and previous_device_index is not None: + torch.cuda.set_device(previous_device_index) + logger.info( + f"[MultiGPU CUDA Guard] Restored CUDA current device {target_device.index} -> {previous_device_index} ({reason})" + ) + +def _get_runtime_device_from_model(model): + """Resolve the actual execution device from a model or patcher wrapper.""" + if hasattr(model, "load_device"): + return getattr(model, "load_device") + patcher = getattr(model, "patcher", None) + if patcher is not None and hasattr(patcher, "load_device"): + return patcher.load_device + inner_model = getattr(model, "model", None) + if inner_model is not None and hasattr(inner_model, "load_device"): + return inner_model.load_device + return None + +@contextmanager +def multigpu_runtime_device_guard(device, reason="runtime"): + """Align MultiGPU logical device state with the real runtime device for inference.""" + original_device = get_current_device() + target_device = _coerce_torch_device(device) or device + if target_device is not None: + set_current_device(target_device) + logger.info(f"[MultiGPU Runtime] Using runtime device {target_device} ({reason})") + try: + with cuda_device_guard(target_device, reason=reason): + yield _coerce_torch_device(target_device) + finally: + set_current_device(original_device) + def get_torch_device_patched(): """Return MultiGPU-aware device selection for patched mm.get_torch_device.""" device = None @@ -238,6 +311,85 @@ def unet_offload_device_patched(): logger.debug(f"[MultiGPU Core Patching] unet_offload_device_patched returning device: {device} (current_unet_offload_device={current_unet_offload_device})") return device +def _patch_model_management_current_stream(): + """Make ComfyUI stream lookup honor the requested CUDA device.""" + current_stream = getattr(mm, "current_stream", None) + if current_stream is None: + return False + if getattr(current_stream, "_multigpu_cuda_device_aware", False): + return True + + def current_stream_device_aware(device): + target_device = _coerce_torch_device(device) + if target_device is not None and target_device.type == "cuda": + return torch.cuda.current_stream(device=target_device) + return current_stream(device) + + current_stream_device_aware._multigpu_cuda_device_aware = True + current_stream_device_aware._multigpu_original = current_stream + mm.current_stream = current_stream_device_aware + logger.info("[MultiGPU] Patched comfy.model_management.current_stream to honor CUDA device arguments") + return True + +def _initialize_aimdo_visible_cuda_devices(): + """Ensure DynamicVRAM initializes every visible CUDA device once when enabled.""" + if not getattr(comfy.memory_management, "aimdo_enabled", False): + logger.info("[MultiGPU] DynamicVRAM not enabled; skipping multi-device aimdo initialization") + return False + if not torch.cuda.is_available(): + logger.info("[MultiGPU] CUDA unavailable; skipping multi-device aimdo initialization") + return False + + try: + from comfy_aimdo import control as aimdo_control + except ImportError: + logger.warning("[MultiGPU] comfy_aimdo unavailable during multi-device initialization") + return False + + init_device = getattr(aimdo_control, "init_device", None) + if not callable(init_device): + logger.warning("[MultiGPU] comfy_aimdo.control.init_device missing; skipping multi-device initialization") + return False + + initialized_any = False + for device_index in range(torch.cuda.device_count()): + if device_index in _aimdo_initialized_devices: + continue + logger.info(f"[MultiGPU] Initializing comfy_aimdo for CUDA device {device_index}") + initialized = bool(init_device(device_index)) + logger.info(f"[MultiGPU] comfy_aimdo init_device({device_index}) -> {initialized}") + if initialized: + _aimdo_initialized_devices.add(device_index) + initialized_any = True + + return initialized_any + +def _patch_comfy_sample_runtime_device(): + """Wrap Comfy sampling entrypoints so runtime device state matches the model load device.""" + sample_fn = getattr(comfy_sample, "sample", None) + if callable(sample_fn) and not getattr(sample_fn, "_multigpu_runtime_device_guard", False): + def sample_with_runtime_device(model, *args, **kwargs): + runtime_device = _get_runtime_device_from_model(model) + with multigpu_runtime_device_guard(runtime_device, reason=f"comfy.sample.sample:{type(model).__name__}"): + return sample_fn(model, *args, **kwargs) + + sample_with_runtime_device._multigpu_runtime_device_guard = True + sample_with_runtime_device._multigpu_original = sample_fn + comfy_sample.sample = sample_with_runtime_device + logger.info("[MultiGPU] Patched comfy.sample.sample with runtime device guard") + + sample_custom_fn = getattr(comfy_sample, "sample_custom", None) + if callable(sample_custom_fn) and not getattr(sample_custom_fn, "_multigpu_runtime_device_guard", False): + def sample_custom_with_runtime_device(model, *args, **kwargs): + runtime_device = _get_runtime_device_from_model(model) + with multigpu_runtime_device_guard(runtime_device, reason=f"comfy.sample.sample_custom:{type(model).__name__}"): + return sample_custom_fn(model, *args, **kwargs) + + sample_custom_with_runtime_device._multigpu_runtime_device_guard = True + sample_custom_with_runtime_device._multigpu_original = sample_custom_fn + comfy_sample.sample_custom = sample_custom_with_runtime_device + logger.info("[MultiGPU] Patched comfy.sample.sample_custom with runtime device guard") + def _patch_comfy_kitchen_dlpack_device_guard(): """Guard comfy_kitchen DLPack export by switching to the tensor's CUDA device.""" try: @@ -256,20 +408,8 @@ def _patch_comfy_kitchen_dlpack_device_guard(): def wrap_for_dlpack_with_device_guard(*args, **kwargs): tensor = args[0] if args else kwargs.get("tensor") - previous_device_index = None - switched_device = False - - if isinstance(tensor, torch.Tensor) and tensor.is_cuda and tensor.device.index is not None: - previous_device_index = torch.cuda.current_device() - if previous_device_index != tensor.device.index: - torch.cuda.set_device(tensor.device.index) - switched_device = True - - try: + with cuda_device_guard(getattr(tensor, "device", None), reason="comfy_kitchen._wrap_for_dlpack"): return wrap_for_dlpack(*args, **kwargs) - finally: - if switched_device and previous_device_index is not None: - torch.cuda.set_device(previous_device_index) wrap_for_dlpack_with_device_guard._multigpu_cuda_device_guard = True comfy_kitchen_cuda._wrap_for_dlpack = wrap_for_dlpack_with_device_guard @@ -284,7 +424,10 @@ logger.info(f"[MultiGPU DEBUG] Initial current_unet_offload_device: {current_une mm.get_torch_device = get_torch_device_patched mm.text_encoder_device = text_encoder_device_patched mm.unet_offload_device = unet_offload_device_patched +_patch_model_management_current_stream() +_patch_comfy_sample_runtime_device() _patch_comfy_kitchen_dlpack_device_guard() +_initialize_aimdo_visible_cuda_devices() from .nodes import ( DeviceSelectorMultiGPU, diff --git a/wrappers.py b/wrappers.py index a132792..304b85c 100644 --- a/wrappers.py +++ b/wrappers.py @@ -447,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, get_current_device + from . import set_current_device, get_current_device, cuda_device_guard class NodeOverride(cls): @classmethod @@ -466,10 +466,11 @@ def override_class(cls): original_device = get_current_device() if device is not None: set_current_device(device) + target_device = device if device is not None else get_current_device() fn = getattr(super(), cls.FUNCTION) - out = fn(*args, **kwargs) try: - return out + with cuda_device_guard(target_device, reason=f"{type(self).__name__}.{cls.FUNCTION}"): + return fn(*args, **kwargs) finally: set_current_device(original_device) @@ -477,7 +478,13 @@ def override_class(cls): def override_class_offload(cls): """Standard MultiGPU device override for UNet/VAE models""" - from . import set_current_device, set_current_unet_offload_device, get_current_device, get_current_unet_offload_device + from . import ( + set_current_device, + set_current_unet_offload_device, + get_current_device, + get_current_unet_offload_device, + cuda_device_guard, + ) class NodeOverride(cls): @classmethod @@ -500,10 +507,11 @@ def override_class_offload(cls): set_current_device(device) if offload_device is not None: set_current_unet_offload_device(offload_device) + target_device = device if device is not None else get_current_device() fn = getattr(super(), cls.FUNCTION) - out = fn(*args, **kwargs) try: - return out + with cuda_device_guard(target_device, reason=f"{type(self).__name__}.{cls.FUNCTION}"): + return fn(*args, **kwargs) finally: set_current_device(original_device) set_current_unet_offload_device(original_offload_device)