Add .vscode/settings.json to .gitignore to exclude IDE-specific settings, DisTorch work in progress. Clip also loaded and distributed. Much WIP.

This commit is contained in:
John Pollock
2025-01-21 02:50:10 -06:00
parent d8d122f397
commit 5eb03a220a
2 changed files with 91 additions and 104 deletions
+1
View File
@@ -1 +1,2 @@
__pycache__/
.vscode/settings.json
+90 -104
View File
@@ -9,11 +9,10 @@ import importlib.util
import logging
import folder_paths
from collections import defaultdict
import builtins
current_device = comfy.model_management.get_torch_device()
current_offload_device = comfy.model_management.get_torch_device()
distorch_allocations = {}
def get_torch_device_patched():
device = None
@@ -23,6 +22,14 @@ def get_torch_device_patched():
device = torch.device(current_device)
return device
def text_encoder_device_patched():
device = None
if (not torch.cuda.is_available() or comfy.model_management.cpu_state == comfy.model_management.CPUState.CPU or "cpu" in str(current_device).lower()):
device = torch.device("cpu")
else:
device = torch.device(current_device)
return device
def unet_offload_device_patched():
device = None
if (not torch.cuda.is_available() or comfy.model_management.cpu_state == comfy.model_management.CPUState.CPU or "cpu" in str(current_offload_device).lower()):
@@ -31,8 +38,67 @@ def unet_offload_device_patched():
device = torch.device(current_offload_device)
return device
def text_encoder_offload_device_patched():
device = None
if (not torch.cuda.is_available() or comfy.model_management.cpu_state == comfy.model_management.CPUState.CPU or "cpu" in str(current_offload_device).lower()):
device = torch.device("cpu")
else:
device = torch.device(current_offload_device)
return device
comfy.model_management.get_torch_device = get_torch_device_patched
comfy.model_management.unet_offload_device = unet_offload_device_patched
comfy.model_management.text_encoder_device = text_encoder_device_patched
comfy.model_management.text_encoder_offload_device = text_encoder_offload_device_patched
def register_patched_ggufmodelpatcher():
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]
module = sys.modules[original_loader.__module__]
if not hasattr(module.GGUFModelPatcher, '_patched'):
original_load = module.GGUFModelPatcher.load
logging.info("MultiGPU: GGUFDisTorch - GGUF ModelPatcher not yet patched, applying patch")
def new_load(self, *args, force_patch_weights=False, **kwargs):
super(module.GGUFModelPatcher, self).load(*args, force_patch_weights=True, **kwargs)
linked = []
module_count = 0
for n, m in self.model.named_modules():
module_count += 1
if hasattr(m, "weight"):
device = getattr(m.weight, "device", None)
#logging.info(f"MultiGPU: GGUFDisTorch - Weight Module {n} on device {device}, offload_device is {self.offload_device}")
if device is not None:
linked.append((n, m))
continue
if hasattr(m, "bias"):
device = getattr(m.bias, "device", None)
#logging.info(f"MultiGPU: GGUFDisTorch - Bias Module {n} on device {device}, offload_device is {self.offload_device}")
if device is not None:
linked.append((n, m))
continue
logging.info(f"MultiGPU: GGUFDisTorch - Found {len(linked)} linked modules out of {module_count} total modules")
if linked:
logging.info(f"MultiGPU: GGUFDisTorch - Found {len(linked)} linked modules, computing reallocation")
device_assignments = analyze_ggml_loading(self.model)['device_assignments']
for device, layers in device_assignments.items():
target_device = torch.device(device)
logging.info(f"MultiGPU: GGUFDisTorch - Moving {len(layers)} layers to {device}")
for n, m, _ in layers:
m.to(self.load_device).to(target_device)
self.mmap_released = True
logging.info("MultiGPU: GGUFDisTorch - self.mmap_released = True")
module.GGUFModelPatcher.load = new_load
module.GGUFModelPatcher._patched = True
logging.info("MultiGPU: GGUFDisTorch - Successfully patched GGUF ModelPatcher")
else:
logging.info("MultiGPU: GGUFDisTorch - GGUF ModelPatcher already patched")
def analyze_ggml_loading(model):
@@ -259,6 +325,7 @@ def override_class_with_distorch(cls):
CATEGORY = "multigpu"
FUNCTION = "override"
def override(self, *args, **kwargs):
global current_device
global distorch_allocations
@@ -267,9 +334,11 @@ def override_class_with_distorch(cls):
distorch_compute_device = kwargs.get("compute_device", None)
if distorch_compute_device is not None:
current_device = distorch_compute_device
register_patched_ggufmodelpatcher()
for key, value in list(kwargs.items()):
if key not in {"unet_name"}:
if key not in {"unet_name", "clip_name1", "clip_name2", "clip_name2", "type"}:
logging.info(f"MultiGPU: Removing {key} from kwargs")
logging.info(f"MultiGPU: Value: {value}")
distorch_allocations[key] = kwargs.pop(key)
@@ -294,7 +363,7 @@ def check_module_exists(module_path):
def register_module(target_nodes):
from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS
for node in target_nodes:
NODE_CLASS_MAPPINGS[f"{node}MultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS[node])
return
@@ -371,53 +440,8 @@ def register_UnetLoaderGGUFDisTorchMultiGPU():
def load_unet(self, unet_name, dequant_dtype=None, patch_dtype=None, patch_on_device=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]
module = sys.modules[original_loader.__module__]
if not hasattr(module.GGUFModelPatcher, '_patched'):
original_load = module.GGUFModelPatcher.load
logging.info("MultiGPU: GGUFDisTorch - GGUF ModelPatcher not yet patched, applying patch")
def new_load(self, *args, force_patch_weights=False, **kwargs):
super(module.GGUFModelPatcher, self).load(*args, force_patch_weights=True, **kwargs)
if not self.mmap_released:
linked = []
if kwargs.get("lowvram_model_memory", 0) > 0:
module_count = 0
for n, m in self.model.named_modules():
module_count += 1
if hasattr(m, "weight"):
device = getattr(m.weight, "device", None) # logging.info(f"MultiGPU: GGUFDisTorch - Weight Module {n} on device {device}, offload_device is {self.offload_device}")
if device is not None:
linked.append((n, m))
continue
if hasattr(m, "bias"):
device = getattr(m.bias, "device", None) # logging.info(f"MultiGPU: GGUFDisTorch - Bias Module {n} on device {device}, offload_device is {self.offload_device}")
if device is not None:
linked.append((n, m))
continue
if linked:
logging.info(f"MultiGPU: GGUFDisTorch - Found {len(linked)} linked modules, computing reallocation")
device_assignments = analyze_ggml_loading(self.model)['device_assignments']
for device, layers in device_assignments.items():
target_device = torch.device(device)
logging.info(f"MultiGPU: GGUFDisTorch - Moving {len(layers)} layers to {device}")
for n, m, _ in layers:
m.to(self.load_device).to(target_device)
self.mmap_released = True
module.GGUFModelPatcher.load = new_load
module.GGUFModelPatcher._patched = True
logging.info("MultiGPU: GGUFDisTorch - Successfully patched GGUF ModelPatcher")
else:
logging.info("MultiGPU: GGUFDisTorch - GGUF ModelPatcher already patched")
logging.info("MultiGPU: GGUFDisTorch - Calling original GGUF loader")
loader_instance = original_loader()
return loader_instance.load_unet(unet_name, dequant_dtype, patch_dtype, patch_on_device)
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]()
return original_loader.load_unet(unet_name, dequant_dtype, patch_dtype, patch_on_device)
# Create the MultiGPU version of the base class
UnetLoaderGGUFDisTorchMultiGPU = override_class_with_distorch(UnetLoaderGGUFDisTorch)
@@ -461,8 +485,6 @@ def register_CLIPLoaderGGUFMultiGPU():
def load_clip(self, clip_name, type="stable_diffusion"):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["CLIPLoaderGGUF"]()
return original_loader.load_clip(clip_name, type)
@@ -471,6 +493,10 @@ def register_CLIPLoaderGGUFMultiGPU():
NODE_CLASS_MAPPINGS["CLIPLoaderGGUFMultiGPU"] = CLIPLoaderGGUFMultiGPU
logging.info(f"MultiGPU: Registered CLIPLoaderGGUFMultiGPU")
CLIPLoaderGGUFDisTorchMultiGPU = override_class_with_distorch(CLIPLoaderGGUF)
NODE_CLASS_MAPPINGS["CLIPLoaderGGUFDisTorchMultiGPU"] = CLIPLoaderGGUFDisTorchMultiGPU
logging.info(f"MultiGPU: Registered CLIPLoaderGGUFDisTorchMultiGPU")
# Now create the advanced version that inherits from the MultiGPU base class
class DualCLIPLoaderGGUF(CLIPLoaderGGUF):
@@ -490,12 +516,18 @@ def register_CLIPLoaderGGUFMultiGPU():
def load_clip(self, clip_name1, clip_name2, type):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUF"]()
return original_loader.load_clip(clip_name1, clip_name2, type)
clip = original_loader.load_clip(clip_name1, clip_name2, type)
clip[0].patcher.load(force_patch_weights=True)
return clip
# Create the MultiGPU version of the advanced class
DualCLIPLoaderGGUFMultiGPU = override_class(DualCLIPLoaderGGUF)
NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFMultiGPU"] = DualCLIPLoaderGGUFMultiGPU
logging.info(f"MultiGPU: Registered DualCLIPLoaderGGUFMultiGPU")
DualCLIPLoaderGGUFDisTorchMultiGPU = override_class_with_distorch(DualCLIPLoaderGGUF)
NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFDisTorchMultiGPU"] = DualCLIPLoaderGGUFDisTorchMultiGPU
logging.info(f"MultiGPU: Registered DualCLIPLoaderGGUFDisTorchMultiGPU")
class TripleCLIPLoaderGGUF(CLIPLoaderGGUF):
@classmethod
def INPUT_TYPES(s):
@@ -519,6 +551,10 @@ def register_CLIPLoaderGGUFMultiGPU():
NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFMultiGPU"] = TripleCLIPLoaderGGUFMultiGPU
logging.info(f"MultiGPU: Registered TripleCLIPLoaderGGUFMultiGPU")
TripleCLIPLoaderGGUFDisTorchMultiGPU = override_class_with_distorch(TripleCLIPLoaderGGUF)
NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFDisTorchMultiGPU"] = TripleCLIPLoaderGGUFDisTorchMultiGPU
logging.info(f"MultiGPU: Registered TripleCLIPLoaderGGUFDisTorchMultiGPU")
def register_LTXVLoaderMultiGPU():
global NODE_CLASS_MAPPINGS
@@ -639,28 +675,6 @@ def register_CheckpointLoaderNF4():
CATEGORY = "loaders"
def load_checkpoint(self, ckpt_name):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["CheckpointLoaderNF4"]()
return original_loader.load_checkpoint(ckpt_name)
NODE_CLASS_MAPPINGS["CheckpointLoaderNF4MultiGPU"] = override_class(CheckpointLoaderNF4)
logging.info(f"MultiGPU: Registered CheckpointLoaderNF4MultiGPU")
def register_CheckpointLoaderNF4():
global NODE_CLASS_MAPPINGS
class CheckpointLoaderNF4:
@classmethod
def INPUT_TYPES(s):
return {"required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
}}
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
FUNCTION = "load_checkpoint"
CATEGORY = "loaders"
def load_checkpoint(self, ckpt_name):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["CheckpointLoaderNF4"]()
@@ -720,34 +734,6 @@ def register_MMAudioModelLoaderMultiGPU():
NODE_CLASS_MAPPINGS["MMAudioModelLoaderMultiGPU"] = override_class(MMAudioModelLoader)
logging.info(f"MultiGPU: Registered MMAudioModelLoaderMultiGPU")
def register_MMAudioModelLoaderMultiGPU():
global NODE_CLASS_MAPPINGS
class MMAudioModelLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mmaudio_model": (folder_paths.get_filename_list("mmaudio"), {"tooltip": "These models are loaded from the 'ComfyUI/models/mmaudio' -folder",}),
"base_precision": (["fp16", "fp32", "bf16"], {"default": "fp16"}),
},
}
RETURN_TYPES = ("MMAUDIO_MODEL",)
RETURN_NAMES = ("mmaudio_model", )
FUNCTION = "loadmodel"
CATEGORY = "MMAudio"
def loadmodel(self, mmaudio_model, base_precision):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["MMAudioModelLoader"]()
return original_loader.loadmodel(mmaudio_model, base_precision)
NODE_CLASS_MAPPINGS["MMAudioModelLoaderMultiGPU"] = override_class(MMAudioModelLoader)
logging.info(f"MultiGPU: Registered MMAudioModelLoaderMultiGPU")
def register_MMAudioFeatureUtilsLoaderMultiGPU():
global NODE_CLASS_MAPPINGS