This commit removes several utility modules used for debugging, memory inspection, and hardware information gathering. These tools are no longer required and their removal simplifies the codebase. The following files have been deleted: - `debug_utils.py` - `device_memory_audit.py` - `hardware_info.py` - `model_sig.py` Additionally, the call to log memory usage on startup has been removed from `__init__.py`.
371 lines
16 KiB
Python
371 lines
16 KiB
Python
import torch
|
|
import logging
|
|
import os
|
|
import copy
|
|
from pathlib import Path
|
|
import folder_paths
|
|
import comfy.model_management as mm
|
|
from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS
|
|
|
|
# --- DisTorch V2 Logging Configuration ---
|
|
# Set to "E" for Engineering (DEBUG) or "P" for Production (INFO)
|
|
LOG_LEVEL = "P"
|
|
|
|
# Configure logger
|
|
logger = logging.getLogger("MultiGPU")
|
|
logger.propagate = False
|
|
|
|
if not logger.handlers:
|
|
log_level = logging.DEBUG if LOG_LEVEL == "E" else logging.INFO
|
|
handler = logging.StreamHandler()
|
|
formatter = logging.Formatter('%(message)s')
|
|
handler.setFormatter(formatter)
|
|
logger.addHandler(handler)
|
|
logger.setLevel(log_level)
|
|
|
|
# --- End Logging Configuration ---
|
|
|
|
# Global device state management
|
|
current_device = mm.get_torch_device()
|
|
current_text_encoder_device = mm.text_encoder_device()
|
|
|
|
def _has_xpu():
|
|
try:
|
|
return hasattr(torch, "xpu") and hasattr(torch.xpu, "is_available") and torch.xpu.is_available()
|
|
except Exception:
|
|
return False
|
|
|
|
def get_device_list():
|
|
devs = ["cpu"]
|
|
try:
|
|
if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_available") and torch.cuda.is_available():
|
|
devs += [f"cuda:{i}" for i in range(torch.cuda.device_count())]
|
|
except Exception:
|
|
pass
|
|
try:
|
|
if _has_xpu():
|
|
devs += [f"xpu:{i}" for i in range(torch.xpu.device_count())]
|
|
except Exception:
|
|
pass
|
|
try:
|
|
if torch.backends.mps.is_available():
|
|
devs += ["mps"]
|
|
except Exception:
|
|
pass
|
|
return devs
|
|
|
|
def set_current_device(device):
|
|
global current_device
|
|
current_device = device
|
|
logger.info(f"[MultiGPU] current_device set to: {device}")
|
|
|
|
def set_current_text_encoder_device(device):
|
|
global current_text_encoder_device
|
|
current_text_encoder_device = device
|
|
logger.info(f"[MultiGPU] current_text_encoder_device set to: {device}")
|
|
|
|
def override_class(cls):
|
|
class NodeOverride(cls):
|
|
@classmethod
|
|
def INPUT_TYPES(s):
|
|
inputs = copy.deepcopy(cls.INPUT_TYPES())
|
|
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=None, **kwargs):
|
|
logger.debug(f"[MultiGPU] override_class called for {cls.__name__} with device={device}")
|
|
|
|
if device is not None:
|
|
set_current_device(device)
|
|
|
|
fn = getattr(super(), cls.FUNCTION)
|
|
out = fn(*args, **kwargs)
|
|
logger.debug(f"[MultiGPU] override_class for {cls.__name__} completed successfully")
|
|
|
|
return out
|
|
|
|
return NodeOverride
|
|
|
|
def override_class_clip(cls):
|
|
class NodeOverride(cls):
|
|
@classmethod
|
|
def INPUT_TYPES(s):
|
|
inputs = copy.deepcopy(cls.INPUT_TYPES())
|
|
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=None, **kwargs):
|
|
if device is not None:
|
|
set_current_text_encoder_device(device)
|
|
|
|
fn = getattr(super(), cls.FUNCTION)
|
|
out = fn(*args, **kwargs)
|
|
|
|
return out
|
|
|
|
return NodeOverride
|
|
|
|
def get_torch_device_patched():
|
|
device = None
|
|
if (not (torch.cuda.is_available() or _has_xpu()) or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_device).lower()):
|
|
device = torch.device("cpu")
|
|
else:
|
|
devs = set(get_device_list())
|
|
device = torch.device(current_device) if str(current_device) in devs else torch.device("cpu")
|
|
logger.debug(f"[MultiGPU] get_torch_device_patched returning device: {device} (current_device={current_device})")
|
|
return device
|
|
|
|
def text_encoder_device_patched():
|
|
device = None
|
|
if (not (torch.cuda.is_available() or _has_xpu()) or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_text_encoder_device).lower()):
|
|
device = torch.device("cpu")
|
|
else:
|
|
devs = set(get_device_list())
|
|
device = torch.device(current_text_encoder_device) if str(current_text_encoder_device) in devs else torch.device("cpu")
|
|
logger.debug(f"[MultiGPU] text_encoder_device_patched returning device: {device} (current_text_encoder_device={current_text_encoder_device})")
|
|
return device
|
|
|
|
# Apply patches
|
|
logger.info(f"[MultiGPU] Patching mm.get_torch_device and mm.text_encoder_device")
|
|
logger.debug(f"[MultiGPU] Initial current_device: {current_device}")
|
|
logger.debug(f"[MultiGPU] Initial current_text_encoder_device: {current_text_encoder_device}")
|
|
mm.get_torch_device = get_torch_device_patched
|
|
mm.text_encoder_device = text_encoder_device_patched
|
|
logger.debug(f"[MultiGPU] Patches applied successfully")
|
|
|
|
def check_module_exists(module_path):
|
|
full_path = os.path.join(folder_paths.get_folder_paths("custom_nodes")[0], module_path)
|
|
logger.debug(f"[MultiGPU] Checking for module at {full_path}")
|
|
if not os.path.exists(full_path):
|
|
logger.debug(f"[MultiGPU] Module {module_path} not found - skipping")
|
|
return False
|
|
logger.debug(f"[MultiGPU] Found {module_path}, creating compatible MultiGPU nodes")
|
|
return True
|
|
|
|
# Import from nodes.py
|
|
from .nodes import (
|
|
DeviceSelectorMultiGPU,
|
|
HunyuanVideoEmbeddingsAdapter,
|
|
UnetLoaderGGUF,
|
|
UnetLoaderGGUFAdvanced,
|
|
CLIPLoaderGGUF,
|
|
DualCLIPLoaderGGUF,
|
|
TripleCLIPLoaderGGUF,
|
|
QuadrupleCLIPLoaderGGUF,
|
|
LTXVLoader,
|
|
Florence2ModelLoader,
|
|
DownloadAndLoadFlorence2Model,
|
|
CheckpointLoaderNF4,
|
|
LoadFluxControlNet,
|
|
MMAudioModelLoader,
|
|
MMAudioFeatureUtilsLoader,
|
|
MMAudioSampler,
|
|
PulidModelLoader,
|
|
PulidInsightFaceLoader,
|
|
PulidEvaClipLoader,
|
|
HyVideoModelLoader,
|
|
HyVideoVAELoader,
|
|
DownloadAndLoadHyVideoTextEncoder,
|
|
)
|
|
|
|
# Import from wanvideo.py
|
|
from .wanvideo import (
|
|
WanVideoModelLoader,
|
|
WanVideoModelLoader_2,
|
|
WanVideoVAELoader,
|
|
LoadWanVideoT5TextEncoder,
|
|
LoadWanVideoClipTextEncoder,
|
|
WanVideoTextEncode,
|
|
WanVideoBlockSwap,
|
|
WanVideoSampler
|
|
)
|
|
|
|
# Import from distorch.py
|
|
from .distorch import (
|
|
model_allocation_store,
|
|
create_model_hash,
|
|
register_patched_ggufmodelpatcher,
|
|
analyze_ggml_loading,
|
|
calculate_vvram_allocation_string,
|
|
override_class_with_distorch_gguf,
|
|
override_class_with_distorch_gguf_v2,
|
|
override_class_with_distorch_clip,
|
|
override_class_with_distorch
|
|
)
|
|
|
|
# Import from distorch_2.py for DisTorch v2 SafeTensor support
|
|
from .distorch_2 import (
|
|
safetensor_allocation_store,
|
|
create_safetensor_model_hash,
|
|
register_patched_safetensor_modelpatcher,
|
|
analyze_safetensor_loading,
|
|
calculate_safetensor_vvram_allocation,
|
|
override_class_with_distorch_safetensor_v2
|
|
)
|
|
|
|
# Initialize NODE_CLASS_MAPPINGS
|
|
NODE_CLASS_MAPPINGS = {
|
|
"DeviceSelectorMultiGPU": DeviceSelectorMultiGPU,
|
|
"HunyuanVideoEmbeddingsAdapter": HunyuanVideoEmbeddingsAdapter,
|
|
}
|
|
|
|
# Standard MultiGPU nodes
|
|
NODE_CLASS_MAPPINGS["UNETLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
|
|
NODE_CLASS_MAPPINGS["VAELoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"])
|
|
NODE_CLASS_MAPPINGS["CLIPLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
|
|
NODE_CLASS_MAPPINGS["DualCLIPLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
|
|
if "TripleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["TripleCLIPLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
|
|
if "QuadrupleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"])
|
|
NODE_CLASS_MAPPINGS["CLIPVisionLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"])
|
|
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
|
|
NODE_CLASS_MAPPINGS["ControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"])
|
|
if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["DiffusersLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"])
|
|
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["DiffControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
|
|
|
|
# DisTorch 2 SafeTensor nodes for FLUX and other safetensor models
|
|
NODE_CLASS_MAPPINGS["UNETLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
|
|
NODE_CLASS_MAPPINGS["VAELoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"])
|
|
NODE_CLASS_MAPPINGS["CLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
|
|
NODE_CLASS_MAPPINGS["DualCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
|
|
if "TripleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["TripleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
|
|
if "QuadrupleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"])
|
|
NODE_CLASS_MAPPINGS["CLIPVisionLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"])
|
|
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
|
|
NODE_CLASS_MAPPINGS["ControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"])
|
|
if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["DiffusersLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"])
|
|
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
|
NODE_CLASS_MAPPINGS["DiffControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
|
|
|
|
# --- Registration Table ---
|
|
logger.info("[MultiGPU] Initiating custom_node Registration. . .")
|
|
dash_line = "-" * 47
|
|
fmt_reg = "{:<30}{:>5}{:>10}"
|
|
logger.info(dash_line)
|
|
logger.info(fmt_reg.format("custom_node", "Found", "Nodes"))
|
|
logger.info(dash_line)
|
|
|
|
registration_data = []
|
|
|
|
def register_and_count(module_names, node_map):
|
|
found = False
|
|
for name in module_names:
|
|
if check_module_exists(name):
|
|
found = True
|
|
break
|
|
|
|
count = 0
|
|
if found:
|
|
initial_len = len(NODE_CLASS_MAPPINGS)
|
|
for key, value in node_map.items():
|
|
NODE_CLASS_MAPPINGS[key] = value
|
|
count = len(NODE_CLASS_MAPPINGS) - initial_len
|
|
|
|
registration_data.append({"name": module_names[0], "found": "Y" if found else "N", "count": count})
|
|
return found
|
|
|
|
# ComfyUI-LTXVideo
|
|
ltx_nodes = {"LTXVLoaderMultiGPU": override_class(LTXVLoader)}
|
|
register_and_count(["ComfyUI-LTXVideo", "comfyui-ltxvideo"], ltx_nodes)
|
|
|
|
# ComfyUI-Florence2
|
|
florence_nodes = {
|
|
"Florence2ModelLoaderMultiGPU": override_class(Florence2ModelLoader),
|
|
"DownloadAndLoadFlorence2ModelMultiGPU": override_class(DownloadAndLoadFlorence2Model)
|
|
}
|
|
register_and_count(["ComfyUI-Florence2", "comfyui-florence2"], florence_nodes)
|
|
|
|
# ComfyUI_bitsandbytes_NF4
|
|
nf4_nodes = {"CheckpointLoaderNF4MultiGPU": override_class(CheckpointLoaderNF4)}
|
|
register_and_count(["ComfyUI_bitsandbytes_NF4", "comfyui_bitsandbytes_nf4"], nf4_nodes)
|
|
|
|
# x-flux-comfyui
|
|
flux_controlnet_nodes = {"LoadFluxControlNetMultiGPU": override_class(LoadFluxControlNet)}
|
|
register_and_count(["x-flux-comfyui"], flux_controlnet_nodes)
|
|
|
|
# ComfyUI-MMAudio
|
|
mmaudio_nodes = {
|
|
"MMAudioModelLoaderMultiGPU": override_class(MMAudioModelLoader),
|
|
"MMAudioFeatureUtilsLoaderMultiGPU": override_class(MMAudioFeatureUtilsLoader),
|
|
"MMAudioSamplerMultiGPU": override_class(MMAudioSampler)
|
|
}
|
|
register_and_count(["ComfyUI-MMAudio", "comfyui-mmaudio"], mmaudio_nodes)
|
|
|
|
# ComfyUI-GGUF
|
|
gguf_nodes = {
|
|
"UnetLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_gguf(UnetLoaderGGUF),
|
|
"UnetLoaderGGUFAdvancedDisTorchMultiGPU": override_class_with_distorch_gguf(UnetLoaderGGUFAdvanced),
|
|
"CLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(CLIPLoaderGGUF),
|
|
"DualCLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(DualCLIPLoaderGGUF),
|
|
"TripleCLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(TripleCLIPLoaderGGUF),
|
|
"QuadrupleCLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF),
|
|
"UnetLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(UnetLoaderGGUF),
|
|
"UnetLoaderGGUFAdvancedDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(UnetLoaderGGUFAdvanced),
|
|
"CLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(CLIPLoaderGGUF),
|
|
"DualCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(DualCLIPLoaderGGUF),
|
|
"TripleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(TripleCLIPLoaderGGUF),
|
|
"QuadrupleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(QuadrupleCLIPLoaderGGUF),
|
|
"UnetLoaderGGUFMultiGPU": override_class(UnetLoaderGGUF),
|
|
"UnetLoaderGGUFAdvancedMultiGPU": override_class(UnetLoaderGGUFAdvanced),
|
|
"CLIPLoaderGGUFMultiGPU": override_class_clip(CLIPLoaderGGUF),
|
|
"DualCLIPLoaderGGUFMultiGPU": override_class_clip(DualCLIPLoaderGGUF),
|
|
"TripleCLIPLoaderGGUFMultiGPU": override_class_clip(TripleCLIPLoaderGGUF),
|
|
"QuadrupleCLIPLoaderGGUFMultiGPU": override_class_clip(QuadrupleCLIPLoaderGGUF)
|
|
}
|
|
register_and_count(["ComfyUI-GGUF", "comfyui-gguf"], gguf_nodes)
|
|
|
|
# PuLID_ComfyUI
|
|
pulid_nodes = {
|
|
"PulidModelLoaderMultiGPU": override_class(PulidModelLoader),
|
|
"PulidInsightFaceLoaderMultiGPU": override_class(PulidInsightFaceLoader),
|
|
"PulidEvaClipLoaderMultiGPU": override_class(PulidEvaClipLoader)
|
|
}
|
|
register_and_count(["PuLID_ComfyUI", "pulid_comfyui"], pulid_nodes)
|
|
|
|
# ComfyUI-HunyuanVideoWrapper
|
|
hunyuan_nodes = {
|
|
"HyVideoModelLoaderMultiGPU": override_class(HyVideoModelLoader),
|
|
"HyVideoVAELoaderMultiGPU": override_class(HyVideoVAELoader),
|
|
"DownloadAndLoadHyVideoTextEncoderMultiGPU": override_class(DownloadAndLoadHyVideoTextEncoder)
|
|
}
|
|
register_and_count(["ComfyUI-HunyuanVideoWrapper", "comfyui-hunyuanvideowrapper"], hunyuan_nodes)
|
|
|
|
# ComfyUI-WanVideoWrapper
|
|
wanvideo_nodes = {
|
|
"WanVideoModelLoaderMultiGPU": WanVideoModelLoader,
|
|
"WanVideoModelLoaderMultiGPU_2": WanVideoModelLoader_2,
|
|
"WanVideoVAELoaderMultiGPU": WanVideoVAELoader,
|
|
"LoadWanVideoT5TextEncoderMultiGPU": LoadWanVideoT5TextEncoder,
|
|
"LoadWanVideoClipTextEncoderMultiGPU": LoadWanVideoClipTextEncoder,
|
|
"WanVideoTextEncodeMultiGPU": WanVideoTextEncode,
|
|
"WanVideoBlockSwapMultiGPU": WanVideoBlockSwap,
|
|
"WanVideoSamplerMultiGPU": WanVideoSampler
|
|
}
|
|
register_and_count(["ComfyUI-WanVideoWrapper", "comfyui-wanvideowrapper"], wanvideo_nodes)
|
|
|
|
# Print the registration table
|
|
for item in registration_data:
|
|
logger.info(fmt_reg.format(item['name'], item['found'], str(item['count'])))
|
|
logger.info(dash_line)
|
|
|
|
|
|
logger.info(f"[MultiGPU] Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
|