refactor: Introduce DisTorch V2 architecture
This commit introduces a major architectural refactoring, laying the groundwork for DisTorch V2. The changes focus on improving modularity, memory management, and diagnostics. Key changes include: - Renaming `distorch_safetensor.py` to `distorch_2.py` to house the new core logic. - Deleting the legacy `block_swap.py` module. - Adding `device_memory_audit.py` for more sophisticated analysis of GPU memory usage. - Implementing a centralized and configurable logging system in `__init__.py` to provide standardized and level-controlled (DEBUG/INFO) output for better debugging.
This commit is contained in:
+120
-88
@@ -7,6 +7,16 @@ 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
|
||||
log_level = logging.DEBUG if LOG_LEVEL == "E" else logging.INFO
|
||||
logging.basicConfig(level=log_level, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
# --- End Logging Configuration ---
|
||||
|
||||
# Global device state management
|
||||
current_device = mm.get_torch_device()
|
||||
current_text_encoder_device = mm.text_encoder_device()
|
||||
@@ -34,12 +44,12 @@ def get_device_list():
|
||||
def set_current_device(device):
|
||||
global current_device
|
||||
current_device = device
|
||||
logging.info(f"[MultiGPU] current_device set to: {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
|
||||
logging.info(f"[MultiGPU] current_text_encoder_device set to: {device}")
|
||||
logger.info(f"[MultiGPU] current_text_encoder_device set to: {device}")
|
||||
|
||||
def override_class(cls):
|
||||
class NodeOverride(cls):
|
||||
@@ -56,16 +66,14 @@ def override_class(cls):
|
||||
FUNCTION = "override"
|
||||
|
||||
def override(self, *args, device=None, **kwargs):
|
||||
logging.info(f"[MultiGPU override_class] Called with device={device}")
|
||||
logger.debug(f"[MultiGPU] override_class called for {cls.__name__} with device={device}")
|
||||
|
||||
if device is not None:
|
||||
set_current_device(device)
|
||||
logging.info(f"[MultiGPU override_class] Setting current_device to {device}")
|
||||
|
||||
fn = getattr(super(), cls.FUNCTION)
|
||||
logging.info(f"[MultiGPU override_class] Calling wrapped function: {cls.__name__}.{cls.FUNCTION}")
|
||||
out = fn(*args, **kwargs)
|
||||
logging.info(f"[MultiGPU override_class] Wrapped function completed successfully")
|
||||
logger.debug(f"[MultiGPU] override_class for {cls.__name__} completed successfully")
|
||||
|
||||
return out
|
||||
|
||||
@@ -103,7 +111,7 @@ def get_torch_device_patched():
|
||||
else:
|
||||
devs = set(get_device_list())
|
||||
device = torch.device(current_device) if str(current_device) in devs else torch.device("cpu")
|
||||
logging.info(f"[MultiGPU get_torch_device_patched] Returning device: {device} (current_device={current_device})")
|
||||
logger.debug(f"[MultiGPU] get_torch_device_patched returning device: {device} (current_device={current_device})")
|
||||
return device
|
||||
|
||||
def text_encoder_device_patched():
|
||||
@@ -113,24 +121,24 @@ def text_encoder_device_patched():
|
||||
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")
|
||||
logging.info(f"[MultiGPU text_encoder_device_patched] Returning device: {device} (current_text_encoder_device={current_text_encoder_device})")
|
||||
logger.debug(f"[MultiGPU] text_encoder_device_patched returning device: {device} (current_text_encoder_device={current_text_encoder_device})")
|
||||
return device
|
||||
|
||||
# Apply patches
|
||||
logging.info(f"[MultiGPU] Patching mm.get_torch_device and mm.text_encoder_device")
|
||||
logging.info(f"[MultiGPU] Initial current_device: {current_device}")
|
||||
logging.info(f"[MultiGPU] Initial current_text_encoder_device: {current_text_encoder_device}")
|
||||
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
|
||||
logging.info(f"[MultiGPU] Patches applied successfully")
|
||||
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)
|
||||
logging.info(f"MultiGPU: Checking for module at {full_path}")
|
||||
logger.debug(f"[MultiGPU] Checking for module at {full_path}")
|
||||
if not os.path.exists(full_path):
|
||||
logging.info(f"MultiGPU: Module {module_path} not found - skipping")
|
||||
logger.debug(f"[MultiGPU] Module {module_path} not found - skipping")
|
||||
return False
|
||||
logging.info(f"MultiGPU: Found {module_path}, creating compatible MultiGPU nodes")
|
||||
logger.debug(f"[MultiGPU] Found {module_path}, creating compatible MultiGPU nodes")
|
||||
return True
|
||||
|
||||
# Import from nodes.py
|
||||
@@ -184,16 +192,8 @@ from .distorch import (
|
||||
override_class_with_distorch
|
||||
)
|
||||
|
||||
# Import from block_swap.py
|
||||
from .block_swap import (
|
||||
analyze_safetensor_distorch,
|
||||
apply_block_swap,
|
||||
override_class_with_distorch_safetensor,
|
||||
override_class_with_distorch_bs
|
||||
)
|
||||
|
||||
# Import from distorch_safetensor.py for FLUX support
|
||||
from .distorch_safetensor import (
|
||||
# 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,
|
||||
@@ -242,89 +242,121 @@ if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
NODE_CLASS_MAPPINGS["DiffControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
|
||||
|
||||
# DisTorch 2 FLUX-specific nodes (these are the ones users will use for FLUX)
|
||||
logging.info("[DISTORCH_SAFETENSOR] Registering FLUX DisTorch2 nodes")
|
||||
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
|
||||
NODE_CLASS_MAPPINGS["UNETLoaderFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
|
||||
if "DualCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
NODE_CLASS_MAPPINGS["DualCLIPLoaderFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
|
||||
if "TripleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
NODE_CLASS_MAPPINGS["TripleCLIPLoaderFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
|
||||
# --- 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
|
||||
if check_module_exists("ComfyUI-LTXVideo") or check_module_exists("comfyui-ltxvideo"):
|
||||
NODE_CLASS_MAPPINGS["LTXVLoaderMultiGPU"] = override_class(LTXVLoader)
|
||||
ltx_nodes = {"LTXVLoaderMultiGPU": override_class(LTXVLoader)}
|
||||
register_and_count(["ComfyUI-LTXVideo", "comfyui-ltxvideo"], ltx_nodes)
|
||||
|
||||
# ComfyUI-Florence2
|
||||
if check_module_exists("ComfyUI-Florence2") or check_module_exists("comfyui-florence2"):
|
||||
NODE_CLASS_MAPPINGS["Florence2ModelLoaderMultiGPU"] = override_class(Florence2ModelLoader)
|
||||
NODE_CLASS_MAPPINGS["DownloadAndLoadFlorence2ModelMultiGPU"] = override_class(DownloadAndLoadFlorence2Model)
|
||||
florence_nodes = {
|
||||
"Florence2ModelLoaderMultiGPU": override_class(Florence2ModelLoader),
|
||||
"DownloadAndLoadFlorence2ModelMultiGPU": override_class(DownloadAndLoadFlorence2Model)
|
||||
}
|
||||
register_and_count(["ComfyUI-Florence2", "comfyui-florence2"], florence_nodes)
|
||||
|
||||
# ComfyUI_bitsandbytes_NF4
|
||||
if check_module_exists("ComfyUI_bitsandbytes_NF4") or check_module_exists("comfyui_bitsandbytes_nf4"):
|
||||
NODE_CLASS_MAPPINGS["CheckpointLoaderNF4MultiGPU"] = override_class(CheckpointLoaderNF4)
|
||||
nf4_nodes = {"CheckpointLoaderNF4MultiGPU": override_class(CheckpointLoaderNF4)}
|
||||
register_and_count(["ComfyUI_bitsandbytes_NF4", "comfyui_bitsandbytes_nf4"], nf4_nodes)
|
||||
|
||||
# x-flux-comfyui
|
||||
if check_module_exists("x-flux-comfyui") or check_module_exists("x-flux-comfyui"):
|
||||
NODE_CLASS_MAPPINGS["LoadFluxControlNetMultiGPU"] = override_class(LoadFluxControlNet)
|
||||
flux_controlnet_nodes = {"LoadFluxControlNetMultiGPU": override_class(LoadFluxControlNet)}
|
||||
register_and_count(["x-flux-comfyui"], flux_controlnet_nodes)
|
||||
|
||||
# ComfyUI-MMAudio
|
||||
if check_module_exists("ComfyUI-MMAudio") or check_module_exists("comfyui-mmaudio"):
|
||||
NODE_CLASS_MAPPINGS["MMAudioModelLoaderMultiGPU"] = override_class(MMAudioModelLoader)
|
||||
NODE_CLASS_MAPPINGS["MMAudioFeatureUtilsLoaderMultiGPU"] = override_class(MMAudioFeatureUtilsLoader)
|
||||
NODE_CLASS_MAPPINGS["MMAudioSamplerMultiGPU"] = override_class(MMAudioSampler)
|
||||
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
|
||||
if check_module_exists("ComfyUI-GGUF") or check_module_exists("comfyui-gguf"):
|
||||
# Legacy DisTorch GGUF nodes
|
||||
NODE_CLASS_MAPPINGS["UnetLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_gguf(UnetLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedDisTorchMultiGPU"] = override_class_with_distorch_gguf(UnetLoaderGGUFAdvanced)
|
||||
NODE_CLASS_MAPPINGS["CLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(CLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(DualCLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(TripleCLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF)
|
||||
|
||||
# DisTorch 2 GGUF nodes
|
||||
NODE_CLASS_MAPPINGS["UnetLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_gguf_v2(UnetLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedDisTorch2MultiGPU"] = override_class_with_distorch_gguf_v2(UnetLoaderGGUFAdvanced)
|
||||
NODE_CLASS_MAPPINGS["CLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(CLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(DualCLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(TripleCLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF)
|
||||
|
||||
# Standard MultiGPU nodes
|
||||
NODE_CLASS_MAPPINGS["UnetLoaderGGUFMultiGPU"] = override_class(UnetLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedMultiGPU"] = override_class(UnetLoaderGGUFAdvanced)
|
||||
NODE_CLASS_MAPPINGS["CLIPLoaderGGUFMultiGPU"] = override_class_clip(CLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFMultiGPU"] = override_class_clip(DualCLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFMultiGPU"] = override_class_clip(TripleCLIPLoaderGGUF)
|
||||
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFMultiGPU"] = override_class_clip(QuadrupleCLIPLoaderGGUF)
|
||||
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_gguf_v2(UnetLoaderGGUF),
|
||||
"UnetLoaderGGUFAdvancedDisTorch2MultiGPU": override_class_with_distorch_gguf_v2(UnetLoaderGGUFAdvanced),
|
||||
"CLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(CLIPLoaderGGUF),
|
||||
"DualCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(DualCLIPLoaderGGUF),
|
||||
"TripleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(TripleCLIPLoaderGGUF),
|
||||
"QuadrupleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(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
|
||||
if check_module_exists("PuLID_ComfyUI") or check_module_exists("pulid_comfyui"):
|
||||
NODE_CLASS_MAPPINGS["PulidModelLoaderMultiGPU"] = override_class(PulidModelLoader)
|
||||
NODE_CLASS_MAPPINGS["PulidInsightFaceLoaderMultiGPU"] = override_class(PulidInsightFaceLoader)
|
||||
NODE_CLASS_MAPPINGS["PulidEvaClipLoaderMultiGPU"] = override_class(PulidEvaClipLoader)
|
||||
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
|
||||
if check_module_exists("ComfyUI-HunyuanVideoWrapper") or check_module_exists("comfyui-hunyuanvideowrapper"):
|
||||
NODE_CLASS_MAPPINGS["HyVideoModelLoaderMultiGPU"] = override_class(HyVideoModelLoader)
|
||||
NODE_CLASS_MAPPINGS["HyVideoVAELoaderMultiGPU"] = override_class(HyVideoVAELoader)
|
||||
NODE_CLASS_MAPPINGS["DownloadAndLoadHyVideoTextEncoderMultiGPU"] = override_class(DownloadAndLoadHyVideoTextEncoder)
|
||||
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
|
||||
if check_module_exists("ComfyUI-WanVideoWrapper") or check_module_exists("comfyui-wanvideowrapper"):
|
||||
NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU"] = WanVideoModelLoader
|
||||
NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU_2"] = WanVideoModelLoader_2
|
||||
NODE_CLASS_MAPPINGS["WanVideoVAELoaderMultiGPU"] = WanVideoVAELoader
|
||||
NODE_CLASS_MAPPINGS["LoadWanVideoT5TextEncoderMultiGPU"] = LoadWanVideoT5TextEncoder
|
||||
NODE_CLASS_MAPPINGS["LoadWanVideoClipTextEncoderMultiGPU"] = LoadWanVideoClipTextEncoder
|
||||
NODE_CLASS_MAPPINGS["WanVideoTextEncodeMultiGPU"] = WanVideoTextEncode
|
||||
NODE_CLASS_MAPPINGS["WanVideoBlockSwapMultiGPU"] = WanVideoBlockSwap
|
||||
NODE_CLASS_MAPPINGS["WanVideoSamplerMultiGPU"] = WanVideoSampler
|
||||
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)
|
||||
|
||||
logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
|
||||
# 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())}")
|
||||
|
||||
# --- Memory Logging Test ---
|
||||
from .debug_utils import log_memory_usage
|
||||
logger.debug("ComfyUI Startup Memory Log:")
|
||||
log_memory_usage("ComfyUI Startup")
|
||||
|
||||
-468
@@ -1,468 +0,0 @@
|
||||
"""
|
||||
Block Swap Module for SafeTensor Models
|
||||
Contains all SafeTensor DisTorch code for block-swap memory optimization
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
import copy
|
||||
from collections import defaultdict
|
||||
import comfy.model_management as mm
|
||||
import torch.nn as nn
|
||||
from .model_sig import get_model_type
|
||||
|
||||
|
||||
class FluxBlockSwapManager:
|
||||
"""
|
||||
Manages block-swapping for FLUX models using the original, fast `forward` patching method.
|
||||
"""
|
||||
def __init__(self, model_patcher):
|
||||
self.model_patcher = model_patcher
|
||||
self.patched_blocks = {}
|
||||
|
||||
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
|
||||
logging.info(f"[FluxBlockSwapManager] Applying forward patch to {len(blocks_to_swap)} blocks.")
|
||||
for i, block in enumerate(blocks_to_swap):
|
||||
block.to(swap_device)
|
||||
original_forward = block.forward
|
||||
|
||||
def create_patched_forward(original_f, b, block_index, cd, sd):
|
||||
def patched_forward(*args, **kwargs):
|
||||
b.to(cd, non_blocking=True)
|
||||
result = original_f(*args, **kwargs)
|
||||
b.to(sd, non_blocking=True)
|
||||
return result
|
||||
return patched_forward
|
||||
|
||||
block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device))
|
||||
self.patched_blocks[block] = original_forward
|
||||
|
||||
def cleanup(self):
|
||||
logging.info(f"[FluxBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
|
||||
for block, original_forward in self.patched_blocks.items():
|
||||
block.forward = original_forward
|
||||
self.patched_blocks = {}
|
||||
|
||||
|
||||
class QwenBlockSwapManager:
|
||||
"""
|
||||
Manages block-swapping for Qwen models.
|
||||
"""
|
||||
def __init__(self, model_patcher):
|
||||
self.model_patcher = model_patcher
|
||||
self.patched_blocks = {}
|
||||
|
||||
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
|
||||
logging.info(f"[QwenBlockSwapManager] Applying forward patch to {len(blocks_to_swap)} blocks.")
|
||||
for i, block in enumerate(blocks_to_swap):
|
||||
block.to(swap_device)
|
||||
original_forward = block.forward
|
||||
|
||||
def create_patched_forward(original_f, b, block_index, cd, sd):
|
||||
def patched_forward(*args, **kwargs):
|
||||
b.to(cd, non_blocking=True)
|
||||
result = original_f(*args, **kwargs)
|
||||
b.to(sd, non_blocking=True)
|
||||
return result
|
||||
return patched_forward
|
||||
|
||||
block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device))
|
||||
self.patched_blocks[block] = original_forward
|
||||
|
||||
def cleanup(self):
|
||||
logging.info(f"[QwenBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
|
||||
for block, original_forward in self.patched_blocks.items():
|
||||
block.forward = original_forward
|
||||
self.patched_blocks = {}
|
||||
|
||||
|
||||
class WanVideoBlockSwapManager:
|
||||
"""
|
||||
Manages block-swapping for WanVideo models using a pre-allocated GPU shell block.
|
||||
"""
|
||||
def __init__(self, model_patcher, gpu_shell_block):
|
||||
self.model_patcher = model_patcher
|
||||
self.gpu_shell_block = gpu_shell_block
|
||||
self.patched_blocks = {}
|
||||
|
||||
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
|
||||
logging.info(f"[WanVideoBlockSwapManager] Applying state_dict patch to {len(blocks_to_swap)} blocks.")
|
||||
for i, block in enumerate(blocks_to_swap):
|
||||
block.to(swap_device) # Ensure the source block is on the swap device
|
||||
original_forward = block.forward
|
||||
|
||||
def create_patched_forward(cpu_block, gpu_shell):
|
||||
def patched_forward(*args, **kwargs):
|
||||
logging.info(f"[DEBUG WANVIDEO SWAP] Loading state_dict from CPU block into GPU shell.")
|
||||
gpu_shell.load_state_dict(cpu_block.state_dict())
|
||||
logging.info(f"[DEBUG WANVIDEO SWAP] Executing forward pass on GPU shell.")
|
||||
result = gpu_shell.forward(*args, **kwargs)
|
||||
return result
|
||||
return patched_forward
|
||||
|
||||
block.forward = create_patched_forward(block, self.gpu_shell_block)
|
||||
self.patched_blocks[block] = original_forward
|
||||
|
||||
def cleanup(self):
|
||||
logging.info(f"[WanVideoBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
|
||||
for block, original_forward in self.patched_blocks.items():
|
||||
block.forward = original_forward
|
||||
self.patched_blocks = {}
|
||||
|
||||
|
||||
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap):
|
||||
"""Provides a detailed analysis of the block swap configuration, mimicking the GGUF DisTorch style."""
|
||||
|
||||
eq_line = "=" * 60
|
||||
dash_line = "-" * 60
|
||||
|
||||
logging.info(eq_line)
|
||||
logging.info(" DisTorch SafeTensor Memory Analysis")
|
||||
logging.info(eq_line)
|
||||
|
||||
fmt_assign = "{:<12}{:>15}{:>15}{:>15}"
|
||||
logging.info(fmt_assign.format("Device", "Role", "Total Mem (GB)", "Config (GB)"))
|
||||
logging.info(dash_line)
|
||||
|
||||
compute_total_gb = mm.get_total_memory(torch.device(compute_device)) / (1024**3)
|
||||
swap_total_gb = mm.get_total_memory(torch.device(swap_device)) / (1024**3)
|
||||
|
||||
logging.info(fmt_assign.format(compute_device, "Compute", f"{compute_total_gb:.2f}", ""))
|
||||
logging.info(fmt_assign.format(swap_device, "Swap", f"{swap_total_gb:.2f}", f"Offload: {virtual_vram_gb:.2f}"))
|
||||
logging.info(dash_line)
|
||||
|
||||
block_summary = defaultdict(lambda: {'count': 0, 'memory': 0})
|
||||
total_memory = 0
|
||||
|
||||
for block in all_blocks:
|
||||
block_type = type(block).__name__
|
||||
block_memory = sum(p.numel() * p.element_size() for p in block.parameters())
|
||||
block_summary[block_type]['count'] += 1
|
||||
block_summary[block_type]['memory'] += block_memory
|
||||
total_memory += block_memory
|
||||
|
||||
logging.info(" DisTorch SafeTensor Block Analysis")
|
||||
logging.info(dash_line)
|
||||
fmt_layer = "{:<20}{:>10}{:>15}{:>12}"
|
||||
logging.info(fmt_layer.format("Block Type", "Count", "Memory (MB)", "% Total"))
|
||||
logging.info(dash_line)
|
||||
|
||||
sorted_blocks = sorted(block_summary.items(), key=lambda x: x[1]['memory'], reverse=True)
|
||||
|
||||
for block_type, data in sorted_blocks:
|
||||
mem_mb = data['memory'] / (1024 * 1024)
|
||||
mem_percent = (data['memory'] / total_memory) * 100 if total_memory > 0 else 0
|
||||
logging.info(fmt_layer.format(block_type, str(data['count']), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
|
||||
logging.info(dash_line)
|
||||
|
||||
logging.info(" DisTorch Final Block Assignments")
|
||||
logging.info(dash_line)
|
||||
fmt_final = "{:<5} {:<25} {:>15} {:>15}"
|
||||
logging.info(fmt_final.format("ID", "Block Type", "Size (MB)", "Assignment"))
|
||||
logging.info(dash_line)
|
||||
|
||||
total_swapped_size_mb = 0
|
||||
swapped_block_ids = {id(b) for b in blocks_to_swap}
|
||||
|
||||
for i, block in enumerate(all_blocks):
|
||||
block_type = type(block).__name__
|
||||
size_mb = sum(p.numel() * p.element_size() for p in block.parameters()) / (1024**2)
|
||||
|
||||
assignment = "SWAP" if id(block) in swapped_block_ids else "COMPUTE"
|
||||
if assignment == "SWAP":
|
||||
total_swapped_size_mb += size_mb
|
||||
|
||||
logging.info(fmt_final.format(i, block_type, f"{size_mb:.2f}", assignment))
|
||||
|
||||
logging.info(dash_line)
|
||||
logging.info(f"Total Blocks Swapped: {len(blocks_to_swap)} of {len(all_blocks)}")
|
||||
logging.info(f"Total VRAM Offloaded: {total_swapped_size_mb / 1024:.2f} GB")
|
||||
logging.info(eq_line)
|
||||
|
||||
|
||||
def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
|
||||
virtual_vram_gb=4.0, expert_mode_allocations=""):
|
||||
"""
|
||||
Identifies the model type and applies the appropriate block swapping strategy.
|
||||
"""
|
||||
model_type = get_model_type(model_patcher)
|
||||
logging.info(f"[BlockSwap] Detected model type: {model_type}")
|
||||
|
||||
if model_type == "FLUX":
|
||||
manager = FluxBlockSwapManager(model_patcher)
|
||||
|
||||
model_to_patch = model_patcher.model.diffusion_model
|
||||
|
||||
all_blocks = []
|
||||
if hasattr(model_to_patch, 'double_blocks'):
|
||||
all_blocks.extend(model_to_patch.double_blocks)
|
||||
if hasattr(model_to_patch, 'single_blocks'):
|
||||
all_blocks.extend(model_to_patch.single_blocks)
|
||||
|
||||
if not all_blocks:
|
||||
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for FLUX model.")
|
||||
return
|
||||
|
||||
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
|
||||
|
||||
if virtual_vram_gb > model_size_gb:
|
||||
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
|
||||
virtual_vram_gb = model_size_gb
|
||||
|
||||
vram_target_bytes = virtual_vram_gb * (1024**3)
|
||||
current_swap_size = 0
|
||||
blocks_to_swap = []
|
||||
|
||||
for block in reversed(all_blocks):
|
||||
if current_swap_size < vram_target_bytes:
|
||||
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
|
||||
blocks_to_swap.append(block)
|
||||
current_swap_size += block_size
|
||||
else:
|
||||
break
|
||||
|
||||
blocks_to_swap.reverse()
|
||||
|
||||
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
|
||||
|
||||
if not blocks_to_swap:
|
||||
logging.warning("[BlockSwap] No blocks designated for swapping.")
|
||||
return
|
||||
|
||||
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
|
||||
|
||||
if not hasattr(model_patcher, 'block_swap_managers'):
|
||||
model_patcher.block_swap_managers = []
|
||||
model_patcher.block_swap_managers.append(manager)
|
||||
|
||||
logging.info("[BlockSwap] FLUX block swap setup complete.")
|
||||
|
||||
elif model_type == "QWEN":
|
||||
manager = QwenBlockSwapManager(model_patcher)
|
||||
|
||||
model_to_patch = model_patcher.model.diffusion_model
|
||||
|
||||
if not hasattr(model_to_patch, 'transformer_blocks'):
|
||||
logging.error("[BlockSwap] CRITICAL: Could not find 'transformer_blocks' in Qwen model. Please analyze model structure.")
|
||||
log_unsupported_model_analysis(model_patcher)
|
||||
return
|
||||
|
||||
all_blocks = model_to_patch.transformer_blocks
|
||||
|
||||
if not all_blocks:
|
||||
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for QWEN model.")
|
||||
return
|
||||
|
||||
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
|
||||
|
||||
if virtual_vram_gb > model_size_gb:
|
||||
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
|
||||
virtual_vram_gb = model_size_gb
|
||||
|
||||
vram_target_bytes = virtual_vram_gb * (1024**3)
|
||||
current_swap_size = 0
|
||||
blocks_to_swap = []
|
||||
|
||||
for block in reversed(all_blocks):
|
||||
if current_swap_size < vram_target_bytes:
|
||||
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
|
||||
blocks_to_swap.append(block)
|
||||
current_swap_size += block_size
|
||||
else:
|
||||
break
|
||||
|
||||
blocks_to_swap.reverse()
|
||||
|
||||
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
|
||||
|
||||
if not blocks_to_swap:
|
||||
logging.warning("[BlockSwap] No blocks designated for swapping for QWEN model.")
|
||||
return
|
||||
|
||||
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
|
||||
|
||||
if not hasattr(model_patcher, 'block_swap_managers'):
|
||||
model_patcher.block_swap_managers = []
|
||||
model_patcher.block_swap_managers.append(manager)
|
||||
|
||||
logging.info("[BlockSwap] QWEN block swap setup complete.")
|
||||
|
||||
elif model_type == "WANVIDEO":
|
||||
model_to_patch = model_patcher.model.diffusion_model
|
||||
|
||||
if not hasattr(model_to_patch, 'blocks'):
|
||||
logging.error("[BlockSwap] CRITICAL: Could not find 'blocks' in WanVideo model. Please analyze model structure.")
|
||||
log_unsupported_model_analysis(model_patcher)
|
||||
return
|
||||
|
||||
all_blocks = model_to_patch.blocks
|
||||
|
||||
if not all_blocks:
|
||||
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for WanVideo model.")
|
||||
return
|
||||
|
||||
# --- Pre-allocation Strategy ---
|
||||
# 1. Find the largest block to create a shell
|
||||
largest_block = max(all_blocks, key=lambda b: sum(p.numel() * p.element_size() for p in b.parameters()))
|
||||
gpu_shell_block = copy.deepcopy(largest_block).to(compute_device)
|
||||
shell_size_mb = sum(p.numel() * p.element_size() for p in gpu_shell_block.parameters()) / (1024**2)
|
||||
logging.info(f"[BlockSwap] Created GPU shell block for WanVideo on {compute_device}, size: {shell_size_mb:.2f} MB")
|
||||
|
||||
manager = WanVideoBlockSwapManager(model_patcher, gpu_shell_block)
|
||||
# --- End Pre-allocation ---
|
||||
|
||||
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
|
||||
|
||||
if virtual_vram_gb > model_size_gb:
|
||||
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
|
||||
virtual_vram_gb = model_size_gb
|
||||
|
||||
vram_target_bytes = virtual_vram_gb * (1024**3)
|
||||
current_swap_size = 0
|
||||
blocks_to_swap = []
|
||||
|
||||
# We still need to identify which blocks to swap (i.e., which ones will use the shell)
|
||||
for block in reversed(all_blocks):
|
||||
if current_swap_size < vram_target_bytes:
|
||||
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
|
||||
blocks_to_swap.append(block)
|
||||
current_swap_size += block_size
|
||||
else:
|
||||
# The rest of the blocks will remain on the compute device and not be patched
|
||||
block.to(compute_device)
|
||||
|
||||
blocks_to_swap.reverse()
|
||||
|
||||
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
|
||||
|
||||
if not blocks_to_swap:
|
||||
logging.warning("[BlockSwap] No blocks designated for swapping for WanVideo model.")
|
||||
return
|
||||
|
||||
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
|
||||
|
||||
if not hasattr(model_patcher, 'block_swap_managers'):
|
||||
model_patcher.block_swap_managers = []
|
||||
model_patcher.block_swap_managers.append(manager)
|
||||
|
||||
logging.info("[BlockSwap] WanVideo block swap setup complete using pre-allocation strategy.")
|
||||
|
||||
else:
|
||||
logging.warning(f"[BlockSwap] Model type '{model_type}' is not yet supported. Logging model structure for analysis.")
|
||||
log_unsupported_model_analysis(model_patcher)
|
||||
|
||||
|
||||
def log_unsupported_model_analysis(model_patcher):
|
||||
"""
|
||||
Logs the structure of an unsupported model for development purposes.
|
||||
This is a diagnostic tool and does not modify the model.
|
||||
"""
|
||||
logging.info("========================================================================")
|
||||
logging.info(" INTERNAL MODEL ANALYZER (UNSUPPORTED MODEL DETECTED)")
|
||||
logging.info("========================================================================")
|
||||
|
||||
if not hasattr(model_patcher, 'model'):
|
||||
logging.error("[ModelAnalyzer] Model patcher does not contain a 'model' attribute.")
|
||||
return
|
||||
|
||||
model = model_patcher.model
|
||||
logging.info(f"[ModelAnalyzer] Root Model Type: {type(model).__name__}")
|
||||
|
||||
if not hasattr(model, 'diffusion_model'):
|
||||
logging.warning("[ModelAnalyzer] Model does not have a 'diffusion_model' attribute. Dumping root model attributes.")
|
||||
_recursive_log_attrs(model, "model")
|
||||
else:
|
||||
diffusion_model = model.diffusion_model
|
||||
logging.info(f"[ModelAnalyzer] Diffusion Model Type: {type(diffusion_model).__name__}")
|
||||
_recursive_log_attrs(diffusion_model, "diffusion_model")
|
||||
|
||||
logging.info("========================================================================")
|
||||
logging.info(" MODEL ANALYSIS COMPLETE")
|
||||
logging.info("========================================================================")
|
||||
|
||||
def _recursive_log_attrs(module, path, seen_modules=None):
|
||||
"""Helper function to recursively log model attributes."""
|
||||
if seen_modules is None:
|
||||
seen_modules = set()
|
||||
|
||||
if id(module) in seen_modules:
|
||||
return
|
||||
seen_modules.add(id(module))
|
||||
|
||||
logging.info(f"--- Analyzing path: '{path}' (Type: {type(module).__name__}) ---")
|
||||
|
||||
# Log named children first
|
||||
children_found = False
|
||||
for name, submodule in module.named_children():
|
||||
children_found = True
|
||||
new_path = f"{path}.{name}" if path else name
|
||||
|
||||
# Heuristic check for potential block lists
|
||||
if isinstance(submodule, (torch.nn.ModuleList, list)) and submodule and all(isinstance(x, torch.nn.Module) for x in submodule):
|
||||
logging.info(f" > [POTENTIAL BLOCK LIST] '{new_path}' | Type: {type(submodule).__name__}, Length: {len(submodule)}")
|
||||
# Also inspect the first block in the list for more detail
|
||||
if len(submodule) > 0:
|
||||
_recursive_log_attrs(submodule[0], f"{new_path}[0]", seen_modules)
|
||||
else:
|
||||
logging.info(f" - Child: '{new_path}' | Type: {type(submodule).__name__}")
|
||||
# Recurse into non-list children
|
||||
_recursive_log_attrs(submodule, new_path, seen_modules)
|
||||
|
||||
if not children_found:
|
||||
logging.info(" No named children found at this level.")
|
||||
|
||||
|
||||
def override_class_with_distorch_safetensor(cls):
|
||||
"""DisTorch 2.0 wrapper for SafeTensor models, providing block-swap memory optimization."""
|
||||
from .nodes import get_device_list
|
||||
|
||||
class NodeOverrideDisTorchSafeTensorv2(cls):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
inputs = copy.deepcopy(cls.INPUT_TYPES())
|
||||
devices = get_device_list()
|
||||
compute_device = devices[1] if len(devices) > 1 else devices[0]
|
||||
|
||||
inputs["optional"] = inputs.get("optional", {})
|
||||
|
||||
inputs["optional"]["compute_device"] = (devices, {"default": compute_device})
|
||||
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 128.0, "step": 0.1})
|
||||
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
|
||||
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
|
||||
return inputs
|
||||
|
||||
CATEGORY = "multigpu/distorch_2"
|
||||
FUNCTION = "override"
|
||||
|
||||
def override(self, *args, compute_device=None, virtual_vram_gb=4.0,
|
||||
donor_device="cpu", expert_mode_allocations="", **kwargs):
|
||||
from . import set_current_device
|
||||
|
||||
logging.info(f"[DisTorch SafeTensor] Override called with: compute_device={compute_device}, donor_device={donor_device}, virtual_vram_gb={virtual_vram_gb}")
|
||||
|
||||
if compute_device is not None:
|
||||
set_current_device(compute_device)
|
||||
|
||||
fn = getattr(super(), cls.FUNCTION)
|
||||
out = fn(*args, **kwargs)
|
||||
|
||||
model = out[0]
|
||||
if hasattr(model, 'model'):
|
||||
logging.info("[DisTorch SafeTensor] Model has 'model' attribute, applying block swap.")
|
||||
apply_block_swap(
|
||||
model,
|
||||
compute_device=compute_device,
|
||||
swap_device=donor_device,
|
||||
virtual_vram_gb=virtual_vram_gb,
|
||||
expert_mode_allocations=expert_mode_allocations
|
||||
)
|
||||
else:
|
||||
logging.warning("[DisTorch SafeTensor] Loaded object does not have a 'model' attribute, skipping block swap.")
|
||||
|
||||
return out
|
||||
|
||||
return NodeOverrideDisTorchSafeTensorv2
|
||||
|
||||
|
||||
# For backwards compatibility, keep the old name pointing to the new safetensor wrapper
|
||||
override_class_with_distorch_bs = override_class_with_distorch_safetensor
|
||||
+4
-4
@@ -11,7 +11,7 @@ def log_memory_usage(label=""):
|
||||
ram_used_gb = status.get('ram_used', 0) / (1024**3)
|
||||
ram_total_gb = status.get('ram_total', 0) / (1024**3)
|
||||
|
||||
log_message = f"[MEM_DEBUG] {label} | RAM Used: {ram_used_gb:.2f}/{ram_total_gb:.2f} GB"
|
||||
log_message = f"[MultiGPU] {label} | RAM Used: {ram_used_gb:.2f}/{ram_total_gb:.2f} GB"
|
||||
|
||||
if 'gpus' in status:
|
||||
for i, gpu in enumerate(status['gpus']):
|
||||
@@ -19,9 +19,9 @@ def log_memory_usage(label=""):
|
||||
vram_total_gb = gpu.get('vram_total', 0) / (1024**3)
|
||||
log_message += f" | VRAM cuda:{i}: {vram_used_gb:.2f}/{vram_total_gb:.2f} GB"
|
||||
|
||||
logging.info(log_message)
|
||||
logging.debug(log_message)
|
||||
|
||||
except ImportError:
|
||||
logging.warning("[MEM_DEBUG] Could not import local CHardwareInfo. Cannot log memory usage.")
|
||||
logging.warning("[MultiGPU] Could not import local CHardwareInfo. Cannot log memory usage.")
|
||||
except Exception as e:
|
||||
logging.error(f"[MEM_DEBUG] Error getting memory usage: {e}")
|
||||
logging.error(f"[MultiGPU] Error getting memory usage: {e}")
|
||||
|
||||
@@ -0,0 +1,351 @@
|
||||
"""
|
||||
Device Memory Audit Utility for ComfyUI MultiGPU
|
||||
Provides tools to inspect and monitor GPU/CPU memory usage and model placement
|
||||
"""
|
||||
|
||||
import torch
|
||||
import gc
|
||||
import psutil
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from typing import Dict, List, Tuple, Any
|
||||
import comfy.model_management as mm
|
||||
|
||||
def format_bytes(bytes_value: int) -> str:
|
||||
"""Convert bytes to human-readable format"""
|
||||
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
|
||||
if bytes_value < 1024.0:
|
||||
return f"{bytes_value:.2f} {unit}"
|
||||
bytes_value /= 1024.0
|
||||
return f"{bytes_value:.2f} PB"
|
||||
|
||||
def get_device_memory_info() -> Dict[str, Dict[str, Any]]:
|
||||
"""
|
||||
Get current memory usage for all available devices
|
||||
|
||||
Returns:
|
||||
Dict with device names as keys and memory info as values
|
||||
"""
|
||||
memory_info = {}
|
||||
|
||||
# Check CUDA devices
|
||||
if torch.cuda.is_available():
|
||||
for i in range(torch.cuda.device_count()):
|
||||
device_name = f"cuda:{i}"
|
||||
torch.cuda.synchronize(i) # Ensure accurate memory reading
|
||||
|
||||
allocated = torch.cuda.memory_allocated(i)
|
||||
reserved = torch.cuda.memory_reserved(i)
|
||||
total = torch.cuda.get_device_properties(i).total_memory
|
||||
free = total - allocated
|
||||
|
||||
memory_info[device_name] = {
|
||||
"allocated": allocated,
|
||||
"allocated_str": format_bytes(allocated),
|
||||
"reserved": reserved,
|
||||
"reserved_str": format_bytes(reserved),
|
||||
"total": total,
|
||||
"total_str": format_bytes(total),
|
||||
"free": free,
|
||||
"free_str": format_bytes(free),
|
||||
"usage_percent": (allocated / total * 100) if total > 0 else 0
|
||||
}
|
||||
|
||||
# Check CPU/System memory
|
||||
vm = psutil.virtual_memory()
|
||||
memory_info["cpu"] = {
|
||||
"allocated": vm.used,
|
||||
"allocated_str": format_bytes(vm.used),
|
||||
"total": vm.total,
|
||||
"total_str": format_bytes(vm.total),
|
||||
"free": vm.available,
|
||||
"free_str": format_bytes(vm.available),
|
||||
"usage_percent": vm.percent
|
||||
}
|
||||
|
||||
return memory_info
|
||||
|
||||
def audit_torch_tensors() -> Dict[str, List[Dict[str, Any]]]:
|
||||
"""
|
||||
Find all torch tensors in memory and group by device
|
||||
|
||||
Returns:
|
||||
Dict with device names as keys and list of tensor info as values
|
||||
"""
|
||||
tensors_by_device = defaultdict(list)
|
||||
|
||||
# Iterate through all objects in memory
|
||||
for obj in gc.get_objects():
|
||||
try:
|
||||
if torch.is_tensor(obj):
|
||||
device_str = str(obj.device)
|
||||
size_bytes = obj.element_size() * obj.numel()
|
||||
|
||||
tensor_info = {
|
||||
"id": id(obj),
|
||||
"shape": list(obj.shape),
|
||||
"dtype": str(obj.dtype),
|
||||
"size_bytes": size_bytes,
|
||||
"size_str": format_bytes(size_bytes),
|
||||
"requires_grad": obj.requires_grad,
|
||||
"is_leaf": obj.is_leaf
|
||||
}
|
||||
|
||||
tensors_by_device[device_str].append(tensor_info)
|
||||
except:
|
||||
# Some objects might not be accessible
|
||||
pass
|
||||
|
||||
# Sort tensors by size for each device
|
||||
for device in tensors_by_device:
|
||||
tensors_by_device[device].sort(key=lambda x: x["size_bytes"], reverse=True)
|
||||
|
||||
return dict(tensors_by_device)
|
||||
|
||||
def audit_model_placement(model) -> Dict[str, List[Tuple[str, str, int]]]:
|
||||
"""
|
||||
Audit where each layer of a model is placed
|
||||
|
||||
Args:
|
||||
model: PyTorch model to audit
|
||||
|
||||
Returns:
|
||||
Dict with device names as keys and list of (layer_name, layer_type, size) tuples
|
||||
"""
|
||||
device_layers = defaultdict(list)
|
||||
|
||||
for name, module in model.named_modules():
|
||||
# Skip container modules
|
||||
if len(list(module.children())) > 0:
|
||||
continue
|
||||
|
||||
# Find device for this module
|
||||
device = None
|
||||
size_bytes = 0
|
||||
|
||||
# Check parameters
|
||||
for param_name, param in module.named_parameters(recurse=False):
|
||||
if param is not None:
|
||||
device = str(param.device)
|
||||
size_bytes += param.element_size() * param.numel()
|
||||
|
||||
if device:
|
||||
layer_type = type(module).__name__
|
||||
device_layers[device].append((name, layer_type, size_bytes))
|
||||
|
||||
return dict(device_layers)
|
||||
|
||||
def audit_comfy_models() -> Dict[str, Any]:
|
||||
"""
|
||||
Audit ComfyUI's loaded models and their placement
|
||||
|
||||
Returns:
|
||||
Dict with model information
|
||||
"""
|
||||
model_info = {}
|
||||
|
||||
# Try to get loaded models from ComfyUI's model management
|
||||
try:
|
||||
# Access currently loaded models if available
|
||||
if hasattr(mm, 'current_loaded_models'):
|
||||
loaded_models = mm.current_loaded_models
|
||||
|
||||
for i, model_data in enumerate(loaded_models):
|
||||
model_name = f"model_{i}"
|
||||
if hasattr(model_data, 'model'):
|
||||
model = model_data.model
|
||||
placement = audit_model_placement(model)
|
||||
|
||||
total_size = 0
|
||||
for device_layers in placement.values():
|
||||
for _, _, size in device_layers:
|
||||
total_size += size
|
||||
|
||||
model_info[model_name] = {
|
||||
"type": type(model).__name__,
|
||||
"device_placement": placement,
|
||||
"total_size": total_size,
|
||||
"total_size_str": format_bytes(total_size)
|
||||
}
|
||||
except Exception as e:
|
||||
logging.warning(f"[MultiGPU] Could not audit ComfyUI models: {e}")
|
||||
|
||||
return model_info
|
||||
|
||||
def print_memory_audit(detailed=False):
|
||||
"""
|
||||
Print a formatted memory audit report
|
||||
|
||||
Args:
|
||||
detailed: If True, include tensor-level details
|
||||
"""
|
||||
logging.info("\n" + "=" * 60)
|
||||
logging.info(" DEVICE MEMORY AUDIT REPORT")
|
||||
logging.info("=" * 60)
|
||||
|
||||
# Device memory overview
|
||||
memory_info = get_device_memory_info()
|
||||
logging.info("\n📊 MEMORY USAGE BY DEVICE:")
|
||||
logging.info("-" * 60)
|
||||
|
||||
for device, info in memory_info.items():
|
||||
logging.info(f"\n{device.upper()}:")
|
||||
logging.info(f" Total: {info['total_str']:>12}")
|
||||
logging.info(f" Allocated: {info['allocated_str']:>12} ({info['usage_percent']:.1f}%)")
|
||||
logging.info(f" Free: {info['free_str']:>12}")
|
||||
if 'reserved_str' in info:
|
||||
logging.info(f" Reserved: {info['reserved_str']:>12}")
|
||||
|
||||
# Tensor audit
|
||||
if detailed:
|
||||
tensors = audit_torch_tensors()
|
||||
logging.info("\n📦 TENSORS BY DEVICE:")
|
||||
logging.info("-" * 60)
|
||||
|
||||
for device, tensor_list in tensors.items():
|
||||
total_size = sum(t["size_bytes"] for t in tensor_list)
|
||||
logging.info(f"\n{device}: {len(tensor_list)} tensors, {format_bytes(total_size)} total")
|
||||
|
||||
# Show top 5 largest tensors
|
||||
for i, tensor in enumerate(tensor_list[:5]):
|
||||
logging.info(f" #{i+1}: Shape {tensor['shape']}, {tensor['size_str']}, {tensor['dtype']}")
|
||||
|
||||
# ComfyUI model audit
|
||||
models = audit_comfy_models()
|
||||
if models:
|
||||
logging.info("\n🤖 COMFYUI LOADED MODELS:")
|
||||
logging.info("-" * 60)
|
||||
|
||||
for model_name, info in models.items():
|
||||
logging.info(f"\n{model_name} ({info['type']}):")
|
||||
logging.info(f" Total size: {info['total_size_str']}")
|
||||
|
||||
for device, layers in info['device_placement'].items():
|
||||
device_size = sum(size for _, _, size in layers)
|
||||
logging.info(f" {device}: {len(layers)} layers, {format_bytes(device_size)}")
|
||||
|
||||
logging.info("\n" + "=" * 60)
|
||||
|
||||
def get_memory_snapshot() -> Dict[str, Any]:
|
||||
"""
|
||||
Get a complete memory snapshot for benchmarking
|
||||
|
||||
Returns:
|
||||
Dict containing all memory information
|
||||
"""
|
||||
return {
|
||||
"device_memory": get_device_memory_info(),
|
||||
"tensors": audit_torch_tensors(),
|
||||
"models": audit_comfy_models()
|
||||
}
|
||||
|
||||
def compare_snapshots(before: Dict[str, Any], after: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Compare two memory snapshots to see what changed
|
||||
|
||||
Args:
|
||||
before: Snapshot taken before an operation
|
||||
after: Snapshot taken after an operation
|
||||
|
||||
Returns:
|
||||
Dict with differences between snapshots
|
||||
"""
|
||||
differences = {}
|
||||
|
||||
# Compare device memory
|
||||
memory_diff = {}
|
||||
for device in after["device_memory"]:
|
||||
if device in before["device_memory"]:
|
||||
before_mem = before["device_memory"][device]["allocated"]
|
||||
after_mem = after["device_memory"][device]["allocated"]
|
||||
diff = after_mem - before_mem
|
||||
|
||||
memory_diff[device] = {
|
||||
"before": format_bytes(before_mem),
|
||||
"after": format_bytes(after_mem),
|
||||
"difference": format_bytes(abs(diff)),
|
||||
"increased": diff > 0
|
||||
}
|
||||
|
||||
differences["memory_changes"] = memory_diff
|
||||
|
||||
# Compare tensor counts
|
||||
tensor_diff = {}
|
||||
for device in after["tensors"]:
|
||||
after_count = len(after["tensors"][device])
|
||||
before_count = len(before["tensors"].get(device, []))
|
||||
|
||||
if after_count != before_count:
|
||||
tensor_diff[device] = {
|
||||
"before": before_count,
|
||||
"after": after_count,
|
||||
"difference": after_count - before_count
|
||||
}
|
||||
|
||||
differences["tensor_count_changes"] = tensor_diff
|
||||
|
||||
return differences
|
||||
|
||||
def print_snapshot_comparison(before: Dict[str, Any], after: Dict[str, Any], label: str = "Operation"):
|
||||
"""
|
||||
Print a formatted comparison of two snapshots
|
||||
|
||||
Args:
|
||||
before: Snapshot before operation
|
||||
after: Snapshot after operation
|
||||
label: Label for the operation being measured
|
||||
"""
|
||||
diff = compare_snapshots(before, after)
|
||||
|
||||
logging.info(f"\n📈 MEMORY CHANGES AFTER {label}:")
|
||||
logging.info("-" * 60)
|
||||
|
||||
for device, changes in diff["memory_changes"].items():
|
||||
symbol = "↑" if changes["increased"] else "↓"
|
||||
logging.info(f"{device}: {changes['before']} → {changes['after']} ({symbol} {changes['difference']})")
|
||||
|
||||
if diff["tensor_count_changes"]:
|
||||
logging.info("\nTensor count changes:")
|
||||
for device, changes in diff["tensor_count_changes"].items():
|
||||
logging.info(f"{device}: {changes['before']} → {changes['after']} ({changes['difference']:+d})")
|
||||
|
||||
# Example usage functions for benchmarking
|
||||
def benchmark_memory_points():
|
||||
"""
|
||||
Example function showing how to capture the 4 key memory points
|
||||
"""
|
||||
logging.info("\n🔍 Starting Memory Benchmark...")
|
||||
|
||||
# 1. Pre-load memory
|
||||
snapshot_pre = get_memory_snapshot()
|
||||
print_memory_audit(detailed=False)
|
||||
|
||||
# User would load model here
|
||||
logging.info("\n[Load your model now]")
|
||||
input("Press Enter when model is loaded...")
|
||||
|
||||
# 2. Post-load memory
|
||||
snapshot_post_load = get_memory_snapshot()
|
||||
print_snapshot_comparison(snapshot_pre, snapshot_post_load, "MODEL LOAD")
|
||||
|
||||
# User would run inference here
|
||||
logging.info("\n[Run inference now]")
|
||||
input("Press Enter when inference is complete...")
|
||||
|
||||
# 3. Active inference memory (captured during)
|
||||
snapshot_active = get_memory_snapshot()
|
||||
print_snapshot_comparison(snapshot_post_load, snapshot_active, "INFERENCE")
|
||||
|
||||
# 4. Post-generation memory
|
||||
logging.info("\n[Waiting for cleanup...]")
|
||||
torch.cuda.empty_cache() # Force cleanup
|
||||
gc.collect()
|
||||
|
||||
snapshot_post_gen = get_memory_snapshot()
|
||||
print_snapshot_comparison(snapshot_active, snapshot_post_gen, "CLEANUP")
|
||||
|
||||
logging.info("\n✅ Benchmark complete!")
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run a basic audit when script is executed directly
|
||||
print_memory_audit(detailed=True)
|
||||
+8
-7
@@ -22,6 +22,7 @@ def create_model_hash(model, caller):
|
||||
first_layers = str(list(model.model_state_dict().keys())[:3])
|
||||
identifier = f"{model_type}_{model_size}_{first_layers}"
|
||||
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
|
||||
logging.debug(f"[MultiGPU_DisTorch_HASH] Created hash for {caller}: {final_hash[:8]}...")
|
||||
return final_hash
|
||||
|
||||
|
||||
@@ -100,7 +101,7 @@ def analyze_ggml_loading(model, allocations_str):
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
logging.info(eq_line)
|
||||
logging.info(" DisTorch Device Allocations")
|
||||
logging.info(" DisTorch Model Device Allocations")
|
||||
logging.info(eq_line)
|
||||
logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
|
||||
logging.info(dash_line)
|
||||
@@ -133,7 +134,7 @@ def analyze_ggml_loading(model, allocations_str):
|
||||
memory_by_type[layer_type] += layer_memory
|
||||
total_memory += layer_memory
|
||||
|
||||
logging.info(" DisTorch GGML Layer Distribution")
|
||||
logging.info(" DisTorch Model Layer Distribution")
|
||||
logging.info(dash_line)
|
||||
fmt_layer = "{:<12}{:>10}{:>14}{:>10}"
|
||||
logging.info(fmt_layer.format("Layer Type", "Layers", "Memory (MB)", "% Total"))
|
||||
@@ -161,7 +162,7 @@ def analyze_ggml_loading(model, allocations_str):
|
||||
device_assignments[device] = layer_list[start_idx:end_idx]
|
||||
current_layer += device_layer_count
|
||||
|
||||
logging.info(" DisTorch Final Device/Layer Assignments")
|
||||
logging.info("DisTorch Model Final Device/Layer Assignments")
|
||||
logging.info(dash_line)
|
||||
fmt_assign = "{:<12}{:>10}{:>14}{:>10}"
|
||||
logging.info(fmt_assign.format("Device", "Layers", "Memory (MB)", "% Total"))
|
||||
@@ -201,7 +202,7 @@ def calculate_vvram_allocation_string(model, virtual_vram_str):
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
logging.info(eq_line)
|
||||
logging.info(" DisTorch Virtual VRAM Analysis")
|
||||
logging.info(" DisTorch Model Virtual VRAM Analysis")
|
||||
logging.info(eq_line)
|
||||
logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virt(GB)"))
|
||||
logging.info(dash_line)
|
||||
@@ -284,7 +285,7 @@ def calculate_vvram_allocation_string(model, virtual_vram_str):
|
||||
|
||||
allocation_string = ";".join(allocation_parts)
|
||||
fmt_mem = "{:<20}{:>20}"
|
||||
logging.info(fmt_mem.format("\nAllocation String", allocation_string))
|
||||
logging.info(fmt_mem.format("\n v1 Expert String", allocation_string))
|
||||
|
||||
return allocation_string
|
||||
|
||||
@@ -389,7 +390,7 @@ def override_class_with_distorch_gguf_v2(cls):
|
||||
|
||||
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
|
||||
|
||||
logging.info(f"[DisTorch GGUF] Full allocation string: {full_allocation}")
|
||||
logging.info(f"[MultiGPU_DisTorch] Full allocation string: {full_allocation}")
|
||||
|
||||
if hasattr(out[0], 'model'):
|
||||
model_hash = create_model_hash(out[0], "override")
|
||||
@@ -450,7 +451,7 @@ def override_class_with_distorch_clip(cls):
|
||||
|
||||
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
|
||||
|
||||
logging.info(f"[DisTorch] Full allocation string: {full_allocation}")
|
||||
logging.info(f"[MultiGPU_DisTorch] Full allocation string: {full_allocation}")
|
||||
|
||||
if hasattr(out[0], 'model'):
|
||||
model_hash = create_model_hash(out[0], "override")
|
||||
|
||||
@@ -15,6 +15,7 @@ import comfy.model_patcher
|
||||
|
||||
# Global store for safetensor model allocations - EXACTLY like GGUF
|
||||
safetensor_allocation_store = {}
|
||||
safetensor_settings_store = {}
|
||||
|
||||
|
||||
def create_safetensor_model_hash(model, caller):
|
||||
@@ -42,7 +43,7 @@ def create_safetensor_model_hash(model, caller):
|
||||
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
|
||||
|
||||
# DEBUG STATEMENT - ALWAYS LOG THE HASH
|
||||
logging.info(f"[SAFETENSOR_HASH] Created hash for {caller}: {final_hash[:8]}...")
|
||||
logging.debug(f"[MULTIGPU_DISTORCHV2_HASH] Created hash for {caller}: {final_hash[:8]}...")
|
||||
return final_hash
|
||||
|
||||
|
||||
@@ -61,7 +62,7 @@ def register_patched_safetensor_modelpatcher():
|
||||
allocations = safetensor_allocation_store.get(debug_hash)
|
||||
|
||||
if allocations:
|
||||
logging.info(f"[DISTORCH_SAFETENSOR] Using static allocation for model {debug_hash[:8]}")
|
||||
logging.info(f"[MULTIGPU_DISTORCHV2] Using static allocation for model {debug_hash[:8]}")
|
||||
# Parse allocation string and apply static assignment
|
||||
device_assignments = analyze_safetensor_loading(self, allocations)
|
||||
|
||||
@@ -78,7 +79,7 @@ def register_patched_safetensor_modelpatcher():
|
||||
|
||||
if hasattr(module, 'weight') or hasattr(module, 'comfy_cast_weights'):
|
||||
# Move to our assigned device
|
||||
logging.info(f"[DISTORCH_SAFETENSOR] Moving {block_name} to {target_device}")
|
||||
logging.debug(f"[MULTIGPU_DISTORCHV2] Moving {block_name} to {target_device}")
|
||||
module.to(target_device)
|
||||
# Mark for ComfyUI's cast system if not already marked
|
||||
if hasattr(module, 'comfy_cast_weights'):
|
||||
@@ -92,7 +93,7 @@ def register_patched_safetensor_modelpatcher():
|
||||
|
||||
comfy.model_patcher.ModelPatcher.partially_load = new_partially_load
|
||||
comfy.model_patcher.ModelPatcher._distorch_patched = True
|
||||
logging.info("[DISTORCH_SAFETENSOR] Successfully patched ModelPatcher.partially_load")
|
||||
logging.info("[MULTIGPU_DISTORCHV2] Successfully patched ModelPatcher.partially_load")
|
||||
|
||||
|
||||
def analyze_safetensor_loading(model_patcher, allocations_str):
|
||||
@@ -134,7 +135,7 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
|
||||
# IDENTICAL LOGGING TO DISTORCH
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
logging.info(eq_line)
|
||||
logging.info(" DisTorch Safetensor Device Allocations")
|
||||
logging.info(" DisTorch2 Model Device Allocations")
|
||||
logging.info(eq_line)
|
||||
logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
|
||||
logging.info(dash_line)
|
||||
@@ -158,13 +159,29 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
|
||||
# Get the actual model from the patcher
|
||||
model = model_patcher.model if hasattr(model_patcher, 'model') else model_patcher
|
||||
|
||||
# Analyze all modules with weights - matching GGML pattern
|
||||
# First pass: calculate total memory to establish threshold
|
||||
total_memory = 0
|
||||
for name, module in model.named_modules():
|
||||
if hasattr(module, "weight") or hasattr(module, "comfy_cast_weights"):
|
||||
try:
|
||||
block_memory = mm.module_size(module)
|
||||
except:
|
||||
block_memory = 0
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
block_memory += module.weight.numel() * module.weight.element_size()
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
block_memory += module.bias.numel() * module.bias.element_size()
|
||||
total_memory += block_memory
|
||||
|
||||
# Set the minimum block size threshold (0.1% of total model memory)
|
||||
MIN_BLOCK_THRESHOLD = total_memory * 0.001
|
||||
|
||||
# Second pass: analyze and collect all blocks, then filter
|
||||
all_blocks = []
|
||||
for name, module in model.named_modules():
|
||||
if hasattr(module, "weight") or hasattr(module, "comfy_cast_weights"):
|
||||
block_type = type(module).__name__
|
||||
block_summary[block_type] = block_summary.get(block_type, 0) + 1
|
||||
|
||||
# Calculate memory using ComfyUI's module_size or manual calculation
|
||||
try:
|
||||
block_memory = mm.module_size(module)
|
||||
except:
|
||||
@@ -174,12 +191,17 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
block_memory += module.bias.numel() * module.bias.element_size()
|
||||
|
||||
# Populate summary dictionaries with ALL blocks for accurate reporting
|
||||
block_summary[block_type] = block_summary.get(block_type, 0) + 1
|
||||
memory_by_type[block_type] += block_memory
|
||||
total_memory += block_memory
|
||||
block_list.append((name, module, block_type))
|
||||
all_blocks.append((name, module, block_type, block_memory))
|
||||
|
||||
# Filter out tiny blocks from the distribution list
|
||||
block_list = [b for b in all_blocks if b[3] >= MIN_BLOCK_THRESHOLD]
|
||||
tiny_block_list = [b for b in all_blocks if b[3] < MIN_BLOCK_THRESHOLD]
|
||||
|
||||
# Log layer distribution - IDENTICAL FORMAT TO GGML
|
||||
logging.info(" DisTorch Safetensor Layer Distribution")
|
||||
logging.info(" DisTorch2 Model Layer Distribution")
|
||||
logging.info(dash_line)
|
||||
fmt_layer = "{:<12}{:>10}{:>14}{:>10}"
|
||||
logging.info(fmt_layer.format("Layer Type", "Layers", "Memory (MB)", "% Total"))
|
||||
@@ -192,60 +214,102 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
|
||||
|
||||
logging.info(dash_line)
|
||||
|
||||
# Distribute blocks across devices - EXACTLY like GGML
|
||||
nonzero_devices = [d for d, r in DEVICE_RATIOS_DISTORCH.items() if r > 0]
|
||||
nonzero_total_ratio = sum(DEVICE_RATIOS_DISTORCH[d] for d in nonzero_devices)
|
||||
# Distribute blocks sequentially from the tail of the model
|
||||
device_assignments = {device: [] for device in DEVICE_RATIOS_DISTORCH.keys()}
|
||||
block_assignments = {} # Map block name to device
|
||||
|
||||
total_blocks = len(block_list)
|
||||
current_block = 0
|
||||
block_assignments = {}
|
||||
|
||||
for idx, device in enumerate(nonzero_devices):
|
||||
ratio = DEVICE_RATIOS_DISTORCH[device]
|
||||
if idx == len(nonzero_devices) - 1:
|
||||
# Last device gets remaining blocks
|
||||
device_block_count = total_blocks - current_block
|
||||
# Determine the primary compute device (first non-cpu device)
|
||||
compute_device = "cuda:0" # Fallback
|
||||
for dev in sorted_devices:
|
||||
if dev != "cpu":
|
||||
compute_device = dev
|
||||
break
|
||||
|
||||
# Calculate total memory to be offloaded to donor devices
|
||||
total_offload_gb = sum(DEVICE_RATIOS_DISTORCH.get(d, 0) for d in sorted_devices if d != compute_device)
|
||||
total_offload_bytes = total_offload_gb * (1024**3)
|
||||
|
||||
offloaded_bytes = 0
|
||||
|
||||
# Iterate from the TAIL of the model
|
||||
for block_name, module, block_type, block_memory in reversed(block_list):
|
||||
try:
|
||||
# block_memory is already calculated
|
||||
pass
|
||||
except:
|
||||
block_memory = 0
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
block_memory += module.weight.numel() * module.weight.element_size()
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
block_memory += module.bias.numel() * module.bias.element_size()
|
||||
|
||||
# Assign to donor device (currently assumes one donor 'cpu') until target is met
|
||||
if offloaded_bytes < total_offload_bytes:
|
||||
# For now, simple offload to CPU, will expand for multi-donor
|
||||
donor_device = "cpu"
|
||||
for dev in sorted_devices:
|
||||
if dev != compute_device:
|
||||
donor_device = dev
|
||||
break # Use first available donor
|
||||
|
||||
block_assignments[block_name] = donor_device
|
||||
offloaded_bytes += block_memory
|
||||
else:
|
||||
device_block_count = int((ratio / nonzero_total_ratio) * total_blocks)
|
||||
|
||||
start_idx = current_block
|
||||
end_idx = current_block + device_block_count
|
||||
device_blocks = block_list[start_idx:end_idx]
|
||||
device_assignments[device] = device_blocks
|
||||
|
||||
# Track block name to device mapping
|
||||
for block_name, module, block_type in device_blocks:
|
||||
block_assignments[block_name] = device
|
||||
|
||||
current_block += device_block_count
|
||||
# Assign remaining blocks to the primary compute device
|
||||
block_assignments[block_name] = compute_device
|
||||
|
||||
# Explicitly assign tiny blocks to the compute device
|
||||
if tiny_block_list:
|
||||
for block_name, module, block_type, block_memory in tiny_block_list:
|
||||
block_assignments[block_name] = compute_device
|
||||
|
||||
# Populate device_assignments from the final block_assignments
|
||||
for block_name, device in block_assignments.items():
|
||||
# Find the block in the original list to get all its info
|
||||
for b_name, b_module, b_type, b_mem in all_blocks:
|
||||
if b_name == block_name:
|
||||
device_assignments[device].append((b_name, b_module, b_type, b_mem))
|
||||
break
|
||||
|
||||
# Log final assignments - IDENTICAL FORMAT TO GGML
|
||||
logging.info(" DisTorch Safetensor Final Device/Layer Assignments")
|
||||
logging.info("DisTorch2 Model Final Device/Layer Assignments")
|
||||
logging.info(dash_line)
|
||||
logging.info(fmt_assign.format("Device", "Layers", "Memory (MB)", "% Total"))
|
||||
logging.info(dash_line)
|
||||
|
||||
# Calculate and log tiny blocks separately
|
||||
if tiny_block_list:
|
||||
tiny_block_memory = sum(b[3] for b in tiny_block_list)
|
||||
tiny_mem_mb = tiny_block_memory / (1024 * 1024)
|
||||
tiny_mem_percent = (tiny_block_memory / total_memory) * 100 if total_memory > 0 else 0
|
||||
device_label = f"{compute_device}(<0.1%)"
|
||||
logging.info(fmt_assign.format(device_label, str(len(tiny_block_list)), f"{tiny_mem_mb:.2f}", f"{tiny_mem_percent:.1f}%"))
|
||||
|
||||
# Log distributed blocks
|
||||
total_assigned_memory = 0
|
||||
device_memories = {}
|
||||
|
||||
for device, blocks in device_assignments.items():
|
||||
device_memory = 0
|
||||
for block_name, module, block_type in blocks:
|
||||
# Use the memory we calculated earlier
|
||||
if block_summary[block_type] > 0:
|
||||
mem_per_layer = memory_by_type[block_type] / block_summary[block_type]
|
||||
device_memory += mem_per_layer
|
||||
# Exclude tiny blocks from this calculation
|
||||
dist_blocks = [b for b in blocks if b[3] >= MIN_BLOCK_THRESHOLD]
|
||||
if not dist_blocks:
|
||||
continue
|
||||
|
||||
device_memory = sum(b[3] for b in dist_blocks)
|
||||
device_memories[device] = device_memory
|
||||
total_assigned_memory += device_memory
|
||||
|
||||
sorted_assignments = sorted(device_assignments.keys(), key=lambda d: (d == "cpu", d))
|
||||
sorted_assignments = sorted(device_memories.keys(), key=lambda d: (d == "cpu", d))
|
||||
|
||||
for dev in sorted_assignments:
|
||||
blocks = device_assignments[dev]
|
||||
# Get only the distributed blocks for the count
|
||||
dist_blocks = [b for b in device_assignments[dev] if b[3] >= MIN_BLOCK_THRESHOLD]
|
||||
if not dist_blocks:
|
||||
continue
|
||||
|
||||
mem_mb = device_memories[dev] / (1024 * 1024)
|
||||
mem_percent = (device_memories[dev] / total_memory) * 100 if total_memory > 0 else 0
|
||||
logging.info(fmt_assign.format(dev, str(len(blocks)), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
|
||||
logging.info(fmt_assign.format(dev, str(len(dist_blocks)), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
|
||||
|
||||
logging.info(dash_line)
|
||||
|
||||
@@ -267,7 +331,7 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
logging.info(eq_line)
|
||||
logging.info(" DisTorch Safetensor Virtual VRAM Analysis")
|
||||
logging.info(" DisTorch2 Model Virtual VRAM Analysis")
|
||||
logging.info(eq_line)
|
||||
logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virt(GB)"))
|
||||
logging.info(dash_line)
|
||||
@@ -349,7 +413,7 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
|
||||
allocation_string = ";".join(allocation_parts)
|
||||
|
||||
fmt_mem = "{:<20}{:>20}"
|
||||
logging.info(fmt_mem.format("\nAllocation String", allocation_string))
|
||||
logging.info(fmt_mem.format("\n v2 Expert String", allocation_string))
|
||||
|
||||
return allocation_string
|
||||
|
||||
@@ -385,10 +449,6 @@ def override_class_with_distorch_safetensor_v2(cls):
|
||||
|
||||
# Register our patched ModelPatcher
|
||||
register_patched_safetensor_modelpatcher()
|
||||
|
||||
# Call original function
|
||||
fn = getattr(super(), cls.FUNCTION)
|
||||
out = fn(*args, **kwargs)
|
||||
|
||||
# Build allocation string - EXACTLY like GGUF
|
||||
vram_string = ""
|
||||
@@ -397,15 +457,49 @@ def override_class_with_distorch_safetensor_v2(cls):
|
||||
|
||||
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
|
||||
|
||||
logging.info(f"[DisTorch Safetensor] Full allocation string: {full_allocation}")
|
||||
# --- Force Model Reload on Setting Change ---
|
||||
# Create a hash of the DisTorch settings
|
||||
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
|
||||
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()[:8]
|
||||
|
||||
# Temporarily load the model to get its hash, without applying our patch yet
|
||||
fn = getattr(super(), cls.FUNCTION)
|
||||
temp_out = fn(*args, **kwargs)
|
||||
|
||||
model_to_check = None
|
||||
if hasattr(temp_out[0], 'model'):
|
||||
model_to_check = temp_out[0]
|
||||
elif hasattr(temp_out[0], 'patcher') and hasattr(temp_out[0].patcher, 'model'):
|
||||
model_to_check = temp_out[0].patcher
|
||||
|
||||
if model_to_check:
|
||||
model_hash = create_safetensor_model_hash(model_to_check, "override_check")
|
||||
|
||||
last_settings_hash = safetensor_settings_store.get(model_hash)
|
||||
|
||||
if last_settings_hash != settings_hash:
|
||||
logging.info(f"[MultiGPU_DisTorch2] Settings changed for model {model_hash[:8]}. Forcing reload.")
|
||||
mm.unload_model(model_to_check)
|
||||
# Update the settings store *before* reloading
|
||||
safetensor_settings_store[model_hash] = settings_hash
|
||||
# Call the loader again now that the model is unloaded
|
||||
out = fn(*args, **kwargs)
|
||||
else:
|
||||
out = temp_out # Use the already loaded model
|
||||
else:
|
||||
out = temp_out # Should not happen, but as a fallback
|
||||
|
||||
logging.info(f"[MULTIGPU_DISTORCHV2] Full allocation string: {full_allocation}")
|
||||
|
||||
# Store allocation for the model - EXACTLY like GGUF
|
||||
if hasattr(out[0], 'model'):
|
||||
model_hash = create_safetensor_model_hash(out[0], "override")
|
||||
safetensor_allocation_store[model_hash] = full_allocation
|
||||
safetensor_settings_store[model_hash] = settings_hash # Ensure it's set
|
||||
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
|
||||
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
|
||||
safetensor_allocation_store[model_hash] = full_allocation
|
||||
safetensor_settings_store[model_hash] = settings_hash # Ensure it's set
|
||||
|
||||
return out
|
||||
|
||||
+27
-33
@@ -49,8 +49,7 @@ class WanVideoModelLoader:
|
||||
|
||||
def loadmodel(self, model, base_precision, device, quantization,
|
||||
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None):
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] ========== CUSTOM IMPLEMENTATION ==========")
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] User selected device: {device}")
|
||||
logging.debug(f"[MultiGPU] WanVideoModelLoader: User selected device: {device}")
|
||||
|
||||
selected_device = torch.device(device)
|
||||
|
||||
@@ -62,7 +61,7 @@ class WanVideoModelLoader:
|
||||
loader_module = inspect.getmodule(original_loader)
|
||||
|
||||
if loader_module:
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] Patching WanVideo modules to use {selected_device}")
|
||||
logging.debug(f"[MultiGPU] Patching WanVideo modules to use {selected_device}")
|
||||
|
||||
original_device = getattr(loader_module, 'device', None)
|
||||
original_offload = getattr(loader_module, 'offload_device', None)
|
||||
@@ -72,7 +71,7 @@ class WanVideoModelLoader:
|
||||
setattr(loader_module, 'device', selected_device)
|
||||
if model_offload_override:
|
||||
setattr(loader_module, 'offload_device', model_offload_override)
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] Using model offload override: {model_offload_override}")
|
||||
logging.debug(f"[MultiGPU] Using model offload override: {model_offload_override}")
|
||||
elif device == "cpu":
|
||||
setattr(loader_module, 'offload_device', selected_device)
|
||||
|
||||
@@ -86,9 +85,9 @@ class WanVideoModelLoader:
|
||||
setattr(nodes_module, 'offload_device', nodes_model_offload_override)
|
||||
elif device == "cpu":
|
||||
setattr(nodes_module, 'offload_device', selected_device)
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] Both modules patched successfully")
|
||||
logging.debug(f"[MultiGPU] Both WanVideo modules patched successfully")
|
||||
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] Calling original loader")
|
||||
logging.debug(f"[MultiGPU] Calling original WanVideo loader")
|
||||
result = original_loader.loadmodel(model, base_precision, load_device, quantization,
|
||||
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
|
||||
|
||||
@@ -100,13 +99,13 @@ class WanVideoModelLoader:
|
||||
block_swap_override = getattr(loader_module, '_block_swap_device_override', None)
|
||||
if block_swap_override:
|
||||
transformer.offload_device = block_swap_override
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] Patched transformer for block swap to use: {block_swap_override}")
|
||||
logging.debug(f"[MultiGPU] Patched WanVideo transformer for block swap to use: {block_swap_override}")
|
||||
|
||||
logging.info(f"[MultiGPU WanVideoModelLoader] Model loaded on {selected_device}")
|
||||
logging.info(f"[MultiGPU] WanVideo model loaded on {selected_device}")
|
||||
|
||||
return result
|
||||
else:
|
||||
logging.error(f"[MultiGPU WanVideoModelLoader] Could not patch modules, falling back")
|
||||
logging.error(f"[MultiGPU] Could not patch WanVideo modules, falling back")
|
||||
return original_loader.loadmodel(model, base_precision, load_device, quantization,
|
||||
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
|
||||
|
||||
@@ -137,7 +136,7 @@ class WanVideoVAELoader:
|
||||
DESCRIPTION = "Loads Wan VAE model with explicit device selection"
|
||||
|
||||
def loadmodel(self, model_name, device, precision="bf16", compile_args=None):
|
||||
logging.info(f"[MultiGPU WanVideoVAELoader] User selected device: {device}")
|
||||
logging.debug(f"[MultiGPU] WanVideoVAELoader: User selected device: {device}")
|
||||
|
||||
from nodes import NODE_CLASS_MAPPINGS
|
||||
original_loader = NODE_CLASS_MAPPINGS["WanVideoVAELoader"]()
|
||||
@@ -146,7 +145,7 @@ class WanVideoVAELoader:
|
||||
|
||||
if loader_module:
|
||||
selected_device = torch.device(device)
|
||||
logging.info(f"[MultiGPU WanVideoVAELoader] Patching modules to use {selected_device}")
|
||||
logging.debug(f"[MultiGPU] Patching WanVideo VAE modules to use {selected_device}")
|
||||
|
||||
setattr(loader_module, 'offload_device', selected_device)
|
||||
setattr(loader_module, 'device', selected_device)
|
||||
@@ -159,10 +158,10 @@ class WanVideoVAELoader:
|
||||
|
||||
result = original_loader.loadmodel(model_name, precision, compile_args)
|
||||
|
||||
logging.info(f"[MultiGPU WanVideoVAELoader] VAE loaded on {selected_device}")
|
||||
logging.info(f"[MultiGPU] WanVideo VAE loaded on {selected_device}")
|
||||
return result
|
||||
else:
|
||||
logging.error(f"[MultiGPU WanVideoVAELoader] Could not patch modules")
|
||||
logging.error(f"[MultiGPU] Could not patch WanVideo VAE modules")
|
||||
return original_loader.loadmodel(model_name, precision, compile_args)
|
||||
|
||||
|
||||
@@ -193,8 +192,7 @@ class LoadWanVideoT5TextEncoder:
|
||||
DESCRIPTION = "Loads Wan text_encoder model from 'ComfyUI/models/text_encoders'"
|
||||
|
||||
def loadmodel(self, model_name, precision, device, quantization="disabled"):
|
||||
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
|
||||
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] User selected device: {device}")
|
||||
logging.debug(f"[MultiGPU] LoadWanVideoT5TextEncoder: User selected device: {device}")
|
||||
|
||||
selected_device = torch.device(device)
|
||||
load_device = "offload_device" if device == "cpu" else "main_device"
|
||||
@@ -205,7 +203,7 @@ class LoadWanVideoT5TextEncoder:
|
||||
loader_module = inspect.getmodule(original_loader)
|
||||
|
||||
if loader_module:
|
||||
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Patching WanVideo modules to use {selected_device}")
|
||||
logging.debug(f"[MultiGPU] Patching WanVideo T5 modules to use {selected_device}")
|
||||
|
||||
setattr(loader_module, 'device', selected_device)
|
||||
if device == "cpu":
|
||||
@@ -220,11 +218,11 @@ class LoadWanVideoT5TextEncoder:
|
||||
|
||||
result = original_loader.loadmodel(model_name, precision, load_device, quantization)
|
||||
|
||||
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Text encoder loaded on {selected_device}")
|
||||
logging.info(f"[MultiGPU] WanVideo T5 Text encoder loaded on {selected_device}")
|
||||
|
||||
return result
|
||||
else:
|
||||
logging.error(f"[MultiGPU LoadWanVideoT5TextEncoder] Could not patch modules, falling back")
|
||||
logging.error(f"[MultiGPU] Could not patch WanVideo T5 modules, falling back")
|
||||
return original_loader.loadmodel(model_name, precision, load_device, quantization)
|
||||
|
||||
class WanVideoTextEncode:
|
||||
@@ -255,7 +253,7 @@ class WanVideoTextEncode:
|
||||
|
||||
def process(self, positive_prompt, negative_prompt, device, t5=None, force_offload=True,
|
||||
model_to_offload=None, use_disk_cache=False):
|
||||
logging.info(f"[MultiGPU WanVideoTextEncode] User selected device: {device}")
|
||||
logging.debug(f"[MultiGPU] WanVideoTextEncode: User selected device: {device}")
|
||||
|
||||
original_device = "gpu" if device != "cpu" else "cpu"
|
||||
|
||||
@@ -266,7 +264,7 @@ class WanVideoTextEncode:
|
||||
|
||||
if encoder_module:
|
||||
selected_device = torch.device(device)
|
||||
logging.info(f"[MultiGPU WanVideoTextEncode] Patching module to use {selected_device}")
|
||||
logging.debug(f"[MultiGPU] Patching WanVideo TextEncode module to use {selected_device}")
|
||||
setattr(encoder_module, 'device', selected_device)
|
||||
|
||||
model_loading_name = encoder_module.__name__.replace('.nodes', '.nodes_model_loading')
|
||||
@@ -278,7 +276,7 @@ class WanVideoTextEncode:
|
||||
force_offload=force_offload, model_to_offload=model_to_offload,
|
||||
use_disk_cache=use_disk_cache, device=original_device)
|
||||
|
||||
logging.info(f"[MultiGPU WanVideoTextEncode] Encoding completed on {selected_device}")
|
||||
logging.info(f"[MultiGPU] WanVideo TextEncode completed on {selected_device}")
|
||||
return result
|
||||
else:
|
||||
return original_encoder.process(positive_prompt, negative_prompt, t5=t5,
|
||||
@@ -308,8 +306,7 @@ class LoadWanVideoClipTextEncoder:
|
||||
DESCRIPTION = "Loads Wan CLIP text encoder model from 'ComfyUI/models/clip_vision'"
|
||||
|
||||
def loadmodel(self, model_name, precision, device):
|
||||
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
|
||||
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] User selected device: {device}")
|
||||
logging.debug(f"[MultiGPU] LoadWanVideoClipTextEncoder: User selected device: {device}")
|
||||
|
||||
selected_device = torch.device(device)
|
||||
load_device = "offload_device" if device == "cpu" else "main_device"
|
||||
@@ -320,7 +317,7 @@ class LoadWanVideoClipTextEncoder:
|
||||
loader_module = inspect.getmodule(original_loader)
|
||||
|
||||
if loader_module:
|
||||
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] Patching WanVideo modules to use {selected_device}")
|
||||
logging.debug(f"[MultiGPU] Patching WanVideo CLIP modules to use {selected_device}")
|
||||
|
||||
setattr(loader_module, 'device', selected_device)
|
||||
if device == "cpu":
|
||||
@@ -335,11 +332,11 @@ class LoadWanVideoClipTextEncoder:
|
||||
|
||||
result = original_loader.loadmodel(model_name, precision, load_device)
|
||||
|
||||
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] CLIP encoder loaded on {selected_device}")
|
||||
logging.info(f"[MultiGPU] WanVideo CLIP encoder loaded on {selected_device}")
|
||||
|
||||
return result
|
||||
else:
|
||||
logging.error(f"[MultiGPU LoadWanVideoClipTextEncoder] Could not patch modules, falling back")
|
||||
logging.error(f"[MultiGPU] Could not patch WanVideo CLIP modules, falling back")
|
||||
return original_loader.loadmodel(model_name, precision, load_device)
|
||||
|
||||
class WanVideoModelLoader_2:
|
||||
@@ -376,7 +373,7 @@ class WanVideoSampler:
|
||||
|
||||
def process(self, model, **kwargs):
|
||||
model_device = model.load_device
|
||||
logging.info(f"[MultiGPU WanVideoSampler] Processing on device: {model_device}")
|
||||
logging.info(f"[MultiGPU] WanVideoSampler: Processing on device: {model_device}")
|
||||
|
||||
for module_name in sys.modules.keys():
|
||||
if 'WanVideoWrapper' in module_name and hasattr(sys.modules[module_name], 'device'):
|
||||
@@ -419,10 +416,7 @@ class WanVideoBlockSwap:
|
||||
|
||||
def setargs(self, blocks_to_swap, swap_device, model_offload_device, offload_img_emb, offload_txt_emb,
|
||||
use_non_blocking=False, vace_blocks_to_swap=0):
|
||||
logging.info(f"[MultiGPU WanVideoBlockSwap] ========== CONFIGURATION ==========")
|
||||
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected swap device: {swap_device}")
|
||||
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected model offload device: {model_offload_device}")
|
||||
logging.info(f"[MultiGPU WanVideoBlockSwap] Blocks to swap: {blocks_to_swap}")
|
||||
logging.debug(f"[MultiGPU] WanVideoBlockSwap: swap_device={swap_device}, model_offload_device={model_offload_device}, blocks_to_swap={blocks_to_swap}")
|
||||
|
||||
selected_swap_device = torch.device(swap_device)
|
||||
selected_offload_device = torch.device(model_offload_device)
|
||||
@@ -433,7 +427,7 @@ class WanVideoBlockSwap:
|
||||
setattr(module, 'offload_device', selected_offload_device)
|
||||
setattr(module, '_block_swap_device_override', selected_swap_device)
|
||||
setattr(module, '_model_offload_device_override', selected_offload_device)
|
||||
logging.info(f"[MultiGPU WanVideoBlockSwap] Patched {module_name} for offload to {selected_offload_device} and swap to {selected_swap_device}")
|
||||
logging.debug(f"[MultiGPU] Patched {module_name} for offload to {selected_offload_device} and swap to {selected_swap_device}")
|
||||
|
||||
if 'WanVideoWrapper' in module_name and module_name.endswith('.nodes'):
|
||||
module = sys.modules[module_name]
|
||||
@@ -451,6 +445,6 @@ class WanVideoBlockSwap:
|
||||
"model_offload_device": model_offload_device,
|
||||
}
|
||||
|
||||
logging.info(f"[MultiGPU WanVideoBlockSwap] Block swap configuration complete")
|
||||
logging.info(f"[MultiGPU] WanVideoBlockSwap configuration complete")
|
||||
|
||||
return (block_swap_args,)
|
||||
|
||||
Reference in New Issue
Block a user