More garbage

This commit is contained in:
John Pollock
2025-08-10 22:34:05 -05:00
parent aa49ad1139
commit 02b02629ba
+18
View File
@@ -8,6 +8,9 @@ import logging
import copy
from collections import defaultdict
import comfy.model_management as mm
import comfy.model_detection as model_detection
import comfy.utils
import folder_paths
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, reserved_swap_gb, all_blocks):
@@ -93,6 +96,11 @@ def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
logging.error("[DisTorch SafeTensor] Could not find a valid model to patch for block swapping.")
return
state_dict = model_patcher.model.state_dict()
unet_key_prefix = model_detection.unet_prefix_from_state_dict(state_dict)
unet_config = model_detection.detect_unet_config(state_dict, unet_key_prefix)
logging.info(f"[DisTorch SafeTensor] Detected unet_config: {unet_config}")
all_blocks = []
# 1. Standard UNet Structure
if hasattr(model_to_patch, 'input_blocks') and hasattr(model_to_patch, 'middle_block') and hasattr(model_to_patch, 'output_blocks'):
@@ -134,6 +142,7 @@ def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
for i, block in enumerate(all_blocks):
# Determine target device for this block
target_device = compute_device if i < blocks_on_compute else swap_device
logging.info(f"[DisTorch SafeTensor] Moving block {i} to {target_device}")
block.to(target_device)
# Patch the forward method only if the block is on the swap device
@@ -142,11 +151,20 @@ def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
def create_patched_forward(original_f, b, block_index, cd, sd):
def patched_forward(*args, **kwargs):
logging.info(f"[DisTorch SafeTensor] Forward pass for block {block_index} (start)")
logging.info(f"[DisTorch SafeTensor] - Block type: {type(b).__name__}")
try:
current_device = next(b.parameters()).device
except StopIteration:
current_device = "cpu" # No parameters, assume cpu
logging.info(f"[DisTorch SafeTensor] - Current device: {current_device}")
logging.info(f"[DisTorch SafeTensor] - Target device: {cd}")
logging.info(f"[DisTorch SafeTensor] Swapping block {block_index} to {cd} for computation.")
b.to(cd, non_blocking=True)
result = original_f(*args, **kwargs)
logging.info(f"[DisTorch SafeTensor] Swapping block {block_index} back to {sd}.")
b.to(sd, non_blocking=True)
logging.info(f"[DisTorch SafeTensor] Forward pass for block {block_index} (end)")
return result
return patched_forward