From ba24f572eed902dd052faa31d1dd188b0bcd8c01 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Sat, 11 Jan 2025 10:03:40 -0600 Subject: [PATCH] feat: Add DeviceSelectorMultiGPU node and device selection functionality - allowing the linking of one or more MultiGPU nodes to the same cuda device in cases where this would prevent accidental errors if should there be a device mismatch futher along in the pipeline due to non-loader nodes performing device-specific tasks or logic. --- __init__.py | 37 ++++++++++++++++++++++++++++++++----- 1 file changed, 32 insertions(+), 5 deletions(-) diff --git a/__init__.py b/__init__.py index fa186e8..3c9e4e9 100644 --- a/__init__.py +++ b/__init__.py @@ -25,27 +25,54 @@ def get_torch_device_patched(): comfy.model_management.get_torch_device = get_torch_device_patched +def get_device_list(): + import torch + return ["cpu"] + [f"cuda:{i}" for i in range(torch.cuda.device_count())] + +class DeviceSelectorMultiGPU: + @classmethod + def INPUT_TYPES(s): + devices = get_device_list() + return { + "required": { + "device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0]}) + } + } + + RETURN_TYPES = (get_device_list(),) + RETURN_NAMES = ("device",) + FUNCTION = "select_device" + CATEGORY = "multigpu" + + def select_device(self, device): + return (device,) + def override_class(cls): class NodeOverride(cls): @classmethod def INPUT_TYPES(s): inputs = copy.deepcopy(cls.INPUT_TYPES()) - devices = ["cpu"] + [f"cuda:{i}" for i in range(torch.cuda.device_count())] - inputs["required"]["device"] = (devices,) + devices = get_device_list() + default_device = devices[1] if len(devices) > 1 else devices[0] + inputs["optional"] = inputs.get("optional", {}) + inputs["optional"]["device"] = (devices, {"default": default_device}) return inputs CATEGORY = "multigpu" FUNCTION = "override" - def override(self, *args, device, **kwargs): + def override(self, *args, device=None, **kwargs): global current_device - current_device = device + if device is not None: + current_device = device fn = getattr(super(), cls.FUNCTION) return fn(*args, **kwargs) return NodeOverride -NODE_CLASS_MAPPINGS = {} +NODE_CLASS_MAPPINGS = { + "DeviceSelectorMultiGPU": DeviceSelectorMultiGPU +} def check_module_exists(module_path): full_path = os.path.join("custom_nodes", module_path)