refactor: Enhance block swap memory analysis report

The memory analysis function, `analyze_safetensor_distorch`, has been improved to provide a more accurate and detailed report.

Instead of estimating the number of swapped blocks based on an average size, the function now receives the actual list of blocks being swapped. It generates a per-block table detailing each block's ID, type, size, and its final assignment (COMPUTE or SWAP).

This provides users with a precise breakdown of the memory offload, reflecting the actual state of the model rather than a theoretical calculation.

Additionally, the unused `log_memory_usage` helper function has been removed.
This commit is contained in:
John Pollock
2025-08-11 19:19:58 -05:00
parent bde78780da
commit 04b5bb0a6f
+39 -49
View File
@@ -12,30 +12,6 @@ import torch.nn as nn
from .model_sig import get_model_type
def log_memory_usage(device, stage=""):
"""Logs the memory usage of a given device."""
if not isinstance(device, torch.device):
device = torch.device(device)
if device.type == 'cuda':
stats = torch.cuda.memory_stats(device)
total_mem = mm.get_total_memory(device)
allocated = stats['allocated_bytes.all.current']
reserved = stats['reserved_bytes.all.current']
logging.info(
f"[MemLog] {stage} - {device}: "
f"Allocated: {allocated / 1024**2:.2f}MB, "
f"Reserved: {reserved / 1024**2:.2f}MB, "
f"Total: {total_mem / 1024**3:.2f}GB"
)
elif device.type == 'cpu':
# Basic CPU memory logging (less detailed than CUDA)
# This requires psutil, which might not be a dependency.
# For now, we'll just log that it's a CPU.
logging.info(f"[MemLog] {stage} - {device}: CPU memory logging is not as detailed.")
class FluxBlockSwapManager:
"""
Manages block-swapping for FLUX models using the original, fast `forward` patching method.
@@ -68,7 +44,7 @@ class FluxBlockSwapManager:
self.patched_blocks = {}
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, all_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
@@ -78,7 +54,6 @@ def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram
logging.info(" DisTorch SafeTensor Memory Analysis")
logging.info(eq_line)
# Device Allocation Table
fmt_assign = "{:<12}{:>15}{:>15}{:>15}"
logging.info(fmt_assign.format("Device", "Role", "Total Mem (GB)", "Config (GB)"))
logging.info(dash_line)
@@ -90,7 +65,6 @@ def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram
logging.info(fmt_assign.format(swap_device, "Swap", f"{swap_total_gb:.2f}", f"Offload: {virtual_vram_gb:.2f}"))
logging.info(dash_line)
# Block Analysis Table
block_summary = defaultdict(lambda: {'count': 0, 'memory': 0})
total_memory = 0
@@ -115,20 +89,28 @@ def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram
logging.info(fmt_layer.format(block_type, str(data['count']), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
logging.info(dash_line)
# Final Assignment Table
model_size_gb = total_memory / (1024**3)
block_size_gb = model_size_gb / len(all_blocks) if all_blocks else 0
blocks_to_offload = int(virtual_vram_gb / block_size_gb) if block_size_gb > 0 else 0
blocks_on_compute = len(all_blocks) - blocks_to_offload
logging.info(" DisTorch Final Block Assignments")
logging.info(dash_line)
fmt_final = "{:<20}{:>15}"
logging.info(fmt_final.format("Total Model Size (GB):", f"{model_size_gb:.2f}"))
logging.info(fmt_final.format("Average Block Size (MB):", f"{block_size_gb * 1024:.2f}" if all_blocks else "N/A"))
fmt_final = "{:<5} {:<25} {:>15} {:>15}"
logging.info(fmt_final.format("ID", "Block Type", "Size (MB)", "Assignment"))
logging.info(dash_line)
logging.info(fmt_final.format("Blocks on Compute:", f"{blocks_on_compute}"))
logging.info(fmt_final.format("Blocks on Swap:", f"{blocks_to_offload}"))
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)
@@ -155,16 +137,27 @@ def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for FLUX model.")
return
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks)
log_memory_usage(compute_device, "Before Swap")
log_memory_usage(swap_device, "Before Swap")
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
block_size_gb = model_size_gb / len(all_blocks) if all_blocks else 0
blocks_to_offload_count = int(virtual_vram_gb / block_size_gb) if block_size_gb > 0 else 0
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 = all_blocks[-blocks_to_offload_count:] if blocks_to_offload_count > 0 else []
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.")
@@ -176,9 +169,6 @@ def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
log_memory_usage(compute_device, "After Swap")
log_memory_usage(swap_device, "After Swap")
logging.info("[BlockSwap] FLUX block swap setup complete.")
else:
logging.warning(f"[BlockSwap] Model type '{model_type}' is not yet supported for block swapping.")