Fixed BlockSwap on MPS backend

This commit is contained in:
lihaoyun6
2025-08-06 23:54:57 +08:00
parent 3b530dc983
commit 3ae18bc845
2 changed files with 31 additions and 12 deletions
+6 -4
View File
@@ -5,6 +5,7 @@
import os
import time
import torch
import platform
from typing import Tuple, Dict, Any
from src.utils.downloads import download_weight, get_base_cache_dir
@@ -155,7 +156,7 @@ class SeedVR2:
cleanup_blockswap(self.runner, keep_state_for_cache=True)
# Clear all caches
if self.runner:
if self.runner:
clear_all_caches(self.runner, debugger)
else:
@@ -335,7 +336,7 @@ class SeedVR2BlockSwap:
"BOOLEAN",
{
"default": True,
"tooltip": "Use non-blocking GPU transfers for better performance.",
"tooltip": "Use non-blocking GPU transfers for better performance.\n(This will always False on macOS to prevent Nan tensors)",
},
),
"offload_io_components": (
@@ -399,11 +400,12 @@ The actual memory savings depend on your specific model architecture and will be
cache_model,
enable_debug,
):
Use_non_blocking = False if platform.system() == "Darwin" else use_non_blocking
if blocks_to_swap > 0 or offload_io_components:
configs = []
if blocks_to_swap > 0:
configs.append(f"{blocks_to_swap} blocks")
if use_non_blocking:
if Use_non_blocking:
configs.append("non blocking")
if offload_io_components:
configs.append("I/O components")
@@ -413,7 +415,7 @@ The actual memory savings depend on your specific model architecture and will be
return (
{
"blocks_to_swap": blocks_to_swap,
"use_non_blocking": use_non_blocking,
"use_non_blocking": Use_non_blocking,
"offload_io_components": offload_io_components,
"cache_model": cache_model,
"enable_debug": enable_debug,
+25 -8
View File
@@ -19,6 +19,7 @@ import weakref
import psutil
import gc
import platform
import psutil
from typing import Dict, Any, List, Tuple, Optional, Union
from src.optimization.memory_manager import get_vram_usage
@@ -105,7 +106,7 @@ class BlockSwapDebugger:
"""Log current memory state for debugging."""
if self.enabled:
# GPU Memory
if torch.cuda.is_available():
if torch.cuda.is_available() or torch.mps.is_available():
allocated_gb, reserved_gb, peak_gb = get_vram_usage()
vram_info = f"VRAM: {allocated_gb:.2f}/{reserved_gb:.2f}GB (peak: {peak_gb:.2f}GB)"
self.vram_history.append(allocated_gb)
@@ -180,6 +181,8 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any]) -> None:
# Determine devices
device = "cuda" if torch.cuda.is_available() else "cpu"
if platform.system() == "Darwin":
device = "mps"
offload_device = str(mm.unet_offload_device())
use_non_blocking = block_swap_config.get("use_non_blocking", True)
@@ -366,7 +369,7 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
self.to(model.main_device, non_blocking=model.use_non_blocking)
# Synchronize if needed
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking:
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking and platform.system() != "Darwin":
torch.cuda.synchronize()
# Execute forward pass with OOM protection
@@ -385,8 +388,13 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
)
# Only clear cache under memory pressure
if torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
mm.soft_empty_cache()
if platform.system() == "Darwin":
mem = psutil.virtual_memory()
if torch.mps.current_allocated_memory() > mem.total * 0.9:
mm.soft_empty_cache()
else:
if torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
mm.soft_empty_cache()
else:
output = original_forward(*args, **kwargs)
@@ -437,7 +445,10 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
# Synchronize if not using non-blocking transfers
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking:
torch.cuda.synchronize()
if platform.system() == "Darwin":
torch.mps.synchronize()
else:
torch.cuda.synchronize()
# Execute forward pass
output = self._original_forward(*args, **kwargs)
@@ -455,8 +466,13 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
)
# Only clear cache under memory pressure
if torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
mm.soft_empty_cache()
if platform.system() == "Darwin":
mem = psutil.virtual_memory()
if torch.mps.current_allocated_memory() > mem.total * 0.9:
mm.soft_empty_cache()
else:
if torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
mm.soft_empty_cache()
return output
@@ -493,7 +509,7 @@ def _patch_rope_for_blockswap(model, debugger: BlockSwapDebugger) -> None:
debugger.log(f"RoPE issue for {name}: {e}")
# Get current device from parameters
current_device = "cuda"
current_device = "mps" if platform.system() == "Darwin" else "cuda"
if list(self.parameters()):
current_device = next(self.parameters()).device
@@ -740,3 +756,4 @@ def cleanup_blockswap(runner, keep_state_for_cache: bool = False) -> None: