Fixed BlockSwap on MPS backend
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user