Add FLUX support with new safetensor v2 implementation

- Add distorch_safetensor.py module with safetensor allocation and model hashing utilities
- Update all DisTorch2 nodes to use new override_class_with_distorch_safetensor_v2
- Add FLUX-specific DisTorch2 nodes for checkpoint and UNET loading
- Import new safetensor functions for VRAM allocation analysis and model patching
This commit is contained in:
John Pollock
2025-08-12 20:16:29 -05:00
parent 298b4b829b
commit d5dc678c04
2 changed files with 443 additions and 12 deletions
+31 -12
View File
@@ -192,6 +192,16 @@ from .block_swap import (
override_class_with_distorch_bs
)
# Import from distorch_safetensor.py for FLUX support
from .distorch_safetensor import (
safetensor_allocation_store,
create_safetensor_model_hash,
register_patched_safetensor_modelpatcher,
analyze_safetensor_loading,
calculate_safetensor_vvram_allocation,
override_class_with_distorch_safetensor_v2
)
# Initialize NODE_CLASS_MAPPINGS
NODE_CLASS_MAPPINGS = {
"DeviceSelectorMultiGPU": DeviceSelectorMultiGPU,
@@ -215,22 +225,31 @@ if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
# DisTorch 2 SafeTensor nodes
NODE_CLASS_MAPPINGS["UNETLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
NODE_CLASS_MAPPINGS["VAELoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"])
NODE_CLASS_MAPPINGS["CLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
NODE_CLASS_MAPPINGS["DualCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
# DisTorch 2 SafeTensor nodes for FLUX and other safetensor models
NODE_CLASS_MAPPINGS["UNETLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
NODE_CLASS_MAPPINGS["VAELoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"])
NODE_CLASS_MAPPINGS["CLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
NODE_CLASS_MAPPINGS["DualCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
if "TripleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["TripleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
NODE_CLASS_MAPPINGS["TripleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
if "QuadrupleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"])
NODE_CLASS_MAPPINGS["CLIPVisionLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"])
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
NODE_CLASS_MAPPINGS["ControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"])
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"])
NODE_CLASS_MAPPINGS["CLIPVisionLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"])
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
NODE_CLASS_MAPPINGS["ControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"])
if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffusersLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"])
NODE_CLASS_MAPPINGS["DiffusersLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"])
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
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"])
# ComfyUI-LTXVideo
if check_module_exists("ComfyUI-LTXVideo") or check_module_exists("comfyui-ltxvideo"):
+412
View File
@@ -0,0 +1,412 @@
"""
DisTorch Safetensor Memory Management Module
Contains all safetensor related code for distributed memory management
Following EXACT patterns from distorch.py for GGUF
"""
import sys
import torch
import logging
import hashlib
import copy
from collections import defaultdict
import comfy.model_management as mm
import comfy.model_patcher
# Global store for safetensor model allocations - EXACTLY like GGUF
safetensor_allocation_store = {}
def create_safetensor_model_hash(model, caller):
"""Create a unique hash for a safetensor model to track allocations - EXACTLY like GGUF"""
if hasattr(model, 'model'):
# For ModelPatcher objects
actual_model = model.model
model_type = type(actual_model).__name__
# Use ComfyUI's model_size if available
if hasattr(model, 'model_size'):
model_size = model.model_size()
else:
model_size = sum(p.numel() * p.element_size() for p in actual_model.parameters())
if hasattr(model, 'model_state_dict'):
first_layers = str(list(model.model_state_dict().keys())[:3])
else:
first_layers = str(list(actual_model.state_dict().keys())[:3])
else:
# Direct model
model_type = type(model).__name__
model_size = sum(p.numel() * p.element_size() for p in model.parameters())
first_layers = str(list(model.state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
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]}...")
return final_hash
def register_patched_safetensor_modelpatcher():
"""Register and patch the ModelPatcher for distributed safetensor loading"""
# Patch ComfyUI's ModelPatcher
if not hasattr(comfy.model_patcher.ModelPatcher, '_distorch_patched'):
original_partially_load = comfy.model_patcher.ModelPatcher.partially_load
def new_partially_load(self, device_to, extra_memory=0, full_load=False, **kwargs):
"""Override to use our static device assignments"""
global safetensor_allocation_store
# Check if we have a device allocation for this model
debug_hash = create_safetensor_model_hash(self, "partial_load")
allocations = safetensor_allocation_store.get(debug_hash)
if allocations:
logging.info(f"[DISTORCH_SAFETENSOR] Using static allocation for model {debug_hash[:8]}")
# Parse allocation string and apply static assignment
device_assignments = analyze_safetensor_loading(self, allocations)
# Apply our static assignments instead of ComfyUI's dynamic ones
for block_name, target_device in device_assignments['block_assignments'].items():
# Find the module by name
parts = block_name.split('.')
module = self.model
for part in parts:
if hasattr(module, part):
module = getattr(module, part)
else:
break
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}")
module.to(target_device)
# Mark for ComfyUI's cast system if not already marked
if hasattr(module, 'comfy_cast_weights'):
module.comfy_cast_weights = True
# Return 0 to indicate no additional memory used on compute device
return 0
else:
# Fall back to original behavior - only pass valid args
return original_partially_load(self, device_to, extra_memory, **kwargs)
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")
def analyze_safetensor_loading(model_patcher, allocations_str):
"""
Analyze and distribute safetensor model blocks across devices
IDENTICAL LOGGING FORMAT TO analyze_ggml_loading
"""
DEVICE_RATIOS_DISTORCH = {}
device_table = {}
distorch_alloc = allocations_str
virtual_vram_gb = 0.0
# Parse allocation string EXACTLY like GGML
if '#' in allocations_str:
distorch_alloc, virtual_vram_str = allocations_str.split('#')
if not distorch_alloc:
distorch_alloc = calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str)
# EXACT SAME FORMATTING AS GGML
eq_line = "=" * 47
dash_line = "-" * 47
fmt_assign = "{:<12}{:>10}{:>14}{:>10}"
# Parse device allocations
for allocation in distorch_alloc.split(';'):
if ',' not in allocation:
continue
dev_name, fraction = allocation.split(',')
fraction = float(fraction)
total_mem_bytes = mm.get_total_memory(torch.device(dev_name))
alloc_gb = (total_mem_bytes * fraction) / (1024**3)
DEVICE_RATIOS_DISTORCH[dev_name] = alloc_gb
device_table[dev_name] = {
"fraction": fraction,
"total_gb": total_mem_bytes / (1024**3),
"alloc_gb": alloc_gb
}
# 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(eq_line)
logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
logging.info(dash_line)
sorted_devices = sorted(device_table.keys(), key=lambda d: (d == "cpu", d))
for dev in sorted_devices:
frac = device_table[dev]["fraction"]
tot_gb = device_table[dev]["total_gb"]
alloc_gb = device_table[dev]["alloc_gb"]
logging.info(fmt_assign.format(dev,f"{int(frac * 100)}%",f"{tot_gb:.2f}",f"{alloc_gb:.2f}"))
logging.info(dash_line)
# Analyze model blocks using ComfyUI's structure
block_summary = {}
block_list = []
memory_by_type = defaultdict(int)
total_memory = 0
# 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
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:
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()
memory_by_type[block_type] += block_memory
total_memory += block_memory
block_list.append((name, module, block_type))
# Log layer distribution - IDENTICAL FORMAT TO GGML
logging.info(" DisTorch Safetensor Layer Distribution")
logging.info(dash_line)
fmt_layer = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_layer.format("Layer Type", "Layers", "Memory (MB)", "% Total"))
logging.info(dash_line)
for layer_type, count in block_summary.items():
mem_mb = memory_by_type[layer_type] / (1024 * 1024)
mem_percent = (memory_by_type[layer_type] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_layer.format(layer_type[:12], str(count), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
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)
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
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
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
# Log final assignments - IDENTICAL FORMAT TO GGML
logging.info(" DisTorch Safetensor Final Device/Layer Assignments")
logging.info(dash_line)
logging.info(fmt_assign.format("Device", "Layers", "Memory (MB)", "% Total"))
logging.info(dash_line)
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
device_memories[device] = device_memory
total_assigned_memory += device_memory
sorted_assignments = sorted(device_assignments.keys(), key=lambda d: (d == "cpu", d))
for dev in sorted_assignments:
blocks = device_assignments[dev]
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(dash_line)
return {
"device_assignments": device_assignments,
"block_assignments": block_assignments
}
def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
"""Calculate virtual VRAM allocation string for distributed safetensor loading"""
recipient_device, vram_amount, donors = virtual_vram_str.split(';')
virtual_vram_gb = float(vram_amount)
# EXACT SAME FORMATTING AS GGML
eq_line = "=" * 47
dash_line = "-" * 47
fmt_assign = "{:<8} {:<6} {:>11} {:>9} {:>9}"
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(eq_line)
logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virt(GB)"))
logging.info(dash_line)
# Calculate recipient VRAM
recipient_vram = mm.get_total_memory(torch.device(recipient_device)) / (1024**3)
recipient_virtual = recipient_vram + virtual_vram_gb
logging.info(fmt_assign.format(recipient_device, 'recip', f"{recipient_vram:.2f}GB",f"{recipient_virtual:.2f}GB", f"+{virtual_vram_gb:.2f}GB"))
# Handle donor devices
ram_donors = [d for d in donors.split(',') if d != 'cpu']
remaining_vram_needed = virtual_vram_gb
donor_device_info = {}
donor_allocations = {}
for donor in ram_donors:
donor_vram = mm.get_total_memory(torch.device(donor)) / (1024**3)
max_donor_capacity = donor_vram * 0.9 # Use 90% max
donation = min(remaining_vram_needed, max_donor_capacity)
donor_virtual = donor_vram - donation
remaining_vram_needed -= donation
donor_allocations[donor] = donation
donor_device_info[donor] = (donor_vram, donor_virtual)
logging.info(fmt_assign.format(donor, 'donor', f"{donor_vram:.2f}GB", f"{donor_virtual:.2f}GB", f"-{donation:.2f}GB"))
# CPU gets the rest
system_dram_gb = mm.get_total_memory(torch.device('cpu')) / (1024**3)
cpu_donation = remaining_vram_needed
cpu_virtual = system_dram_gb - cpu_donation
donor_allocations['cpu'] = cpu_donation
logging.info(fmt_assign.format('cpu', 'donor', f"{system_dram_gb:.2f}GB", f"{cpu_virtual:.2f}GB", f"-{cpu_donation:.2f}GB"))
logging.info(dash_line)
# Calculate model size
model = model_patcher.model if hasattr(model_patcher, 'model') else model_patcher
total_memory = 0
for name, module in model.named_modules():
if hasattr(module, "weight"):
if module.weight is not None:
total_memory += module.weight.numel() * module.weight.element_size()
if hasattr(module, "bias") and module.bias is not None:
total_memory += module.bias.numel() * module.bias.element_size()
model_size_gb = total_memory / (1024**3)
new_model_size_gb = max(0, model_size_gb - virtual_vram_gb)
logging.info(fmt_assign.format('model', 'model', f"{model_size_gb:.2f}GB",f"{new_model_size_gb:.2f}GB", f"-{virtual_vram_gb:.2f}GB"))
# Warning if model too large
if model_size_gb > (recipient_vram * 0.9):
on_recipient = recipient_vram * 0.9
on_virtuals = model_size_gb - on_recipient
logging.info(f"\nWarning: Model size is greater than 90% of recipient VRAM. {on_virtuals:.2f} GB of Layers Offloaded Automatically to Virtual VRAM.\n")
else:
on_recipient = model_size_gb
on_virtuals = 0
new_on_recipient = max(0, on_recipient - virtual_vram_gb)
# Build allocation string - EXACTLY like GGML
allocation_parts = []
recipient_percent = new_on_recipient / recipient_vram
allocation_parts.append(f"{recipient_device},{recipient_percent:.4f}")
for donor in ram_donors:
donor_vram = donor_device_info[donor][0]
donor_percent = donor_allocations[donor] / donor_vram
allocation_parts.append(f"{donor},{donor_percent:.4f}")
cpu_percent = donor_allocations['cpu'] / system_dram_gb
allocation_parts.append(f"cpu,{cpu_percent:.4f}")
allocation_string = ";".join(allocation_parts)
fmt_mem = "{:<20}{:>20}"
logging.info(fmt_mem.format("\nAllocation String", allocation_string))
return allocation_string
def override_class_with_distorch_safetensor_v2(cls):
"""DisTorch 2.0 wrapper for safetensor models - EXACTLY like GGUF wrapper"""
from .nodes import get_device_list
from . import current_device
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"
TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (DisTorch2)"
def override(self, *args, compute_device=None, virtual_vram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", **kwargs):
from . import set_current_device
if compute_device is not None:
set_current_device(compute_device)
# 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 = ""
if virtual_vram_gb > 0:
vram_string = f"{compute_device};{virtual_vram_gb};{donor_device}"
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}")
# 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
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
return out
return NodeOverrideDisTorchSafetensorV2