Merge pull request #248 from AInVFX/main
Fix: torch.mps AttributeError on Windows
This commit is contained in:
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "seedvr2_videoupscaler"
|
||||
description = "SeedVR2 official ComfyUI integration: ByteDance-Seed's one-step diffusion-based video/image upscaling with memory-efficient inference"
|
||||
version = "2.5.1"
|
||||
version = "2.5.2"
|
||||
authors = [
|
||||
{name = "numz"},
|
||||
{name = "adrientoupet"}
|
||||
|
||||
@@ -47,7 +47,7 @@ def get_device() -> torch.device:
|
||||
"""
|
||||
Get current rank device.
|
||||
"""
|
||||
if torch.mps.is_available():
|
||||
if hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
return torch.device("mps")
|
||||
return torch.device("cuda", get_local_rank())
|
||||
|
||||
|
||||
@@ -31,7 +31,7 @@ class AreaResize:
|
||||
self.max_area = max_area
|
||||
self.downsample_only = downsample_only
|
||||
self.interpolation = interpolation
|
||||
if torch.mps.is_available():
|
||||
if hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
self.interpolation = InterpolationMode.BILINEAR
|
||||
|
||||
def __call__(self, image: Union[torch.Tensor, Image.Image]):
|
||||
|
||||
@@ -26,7 +26,7 @@ def NaResize(
|
||||
max_resolution: int = 0,
|
||||
interpolation: InterpolationMode = InterpolationMode.BICUBIC,
|
||||
):
|
||||
Interpolation = InterpolationMode.BILINEAR if torch.mps.is_available() else interpolation
|
||||
Interpolation = InterpolationMode.BILINEAR if (hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available()) else interpolation
|
||||
if mode == "area":
|
||||
return AreaResize(
|
||||
max_area=resolution**2,
|
||||
|
||||
@@ -30,7 +30,7 @@ class SideResize:
|
||||
self.max_size = max_size
|
||||
self.downsample_only = downsample_only
|
||||
self.interpolation = interpolation
|
||||
if torch.mps.is_available():
|
||||
if hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
self.interpolation = InterpolationMode.BILINEAR
|
||||
|
||||
def __call__(self, image: Union[torch.Tensor, Image.Image]):
|
||||
|
||||
@@ -207,7 +207,7 @@ class FP8CompatibleDiT(torch.nn.Module):
|
||||
self._convert_rope_freqs(target_dtype=self.compute_dtype)
|
||||
self.debug.end_timer("_convert_rope_freqs", "RoPE freqs conversion")
|
||||
|
||||
if torch.mps.is_available():
|
||||
if hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
self.debug.log(f"Also converting NaDiT parameters/buffers for MPS backend", category="setup", force=True)
|
||||
self.debug.start_timer("_force_nadit_precision")
|
||||
self._force_nadit_precision(target_dtype=self.compute_dtype)
|
||||
@@ -510,7 +510,7 @@ class FP8CompatibleDiT(torch.nn.Module):
|
||||
k = k.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
|
||||
v = v.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
|
||||
|
||||
if torch.mps.is_available():
|
||||
if hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
dropout_p=0.0,
|
||||
|
||||
@@ -80,7 +80,7 @@ def get_basic_vram_info(device: Optional[torch.device] = None) -> Dict[str, Any]
|
||||
elif not isinstance(device, torch.device):
|
||||
device = torch.device(device)
|
||||
free_memory, total_memory = torch.cuda.mem_get_info(device)
|
||||
elif torch.mps.is_available():
|
||||
elif hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
# MPS doesn't support per-device queries or mem_get_info
|
||||
# Use system memory as proxy
|
||||
mem = psutil.virtual_memory()
|
||||
@@ -100,7 +100,7 @@ def get_basic_vram_info(device: Optional[torch.device] = None) -> Dict[str, Any]
|
||||
# Initial VRAM check at module load
|
||||
vram_info = get_basic_vram_info(device=None)
|
||||
if "error" not in vram_info:
|
||||
backend = "MPS" if torch.mps.is_available() else "CUDA"
|
||||
backend = "MPS" if (hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available()) else "CUDA"
|
||||
print(f"📊 Initial {backend} memory: {vram_info['free_gb']:.2f}GB free / {vram_info['total_gb']:.2f}GB total")
|
||||
else:
|
||||
print(f"⚠️ Memory check failed: {vram_info['error']} - No available backend!")
|
||||
@@ -129,7 +129,7 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug
|
||||
reserved = torch.cuda.memory_reserved(device) / (1024**3)
|
||||
max_allocated = torch.cuda.max_memory_allocated(device) / (1024**3)
|
||||
return allocated, reserved, max_allocated
|
||||
elif torch.mps.is_available():
|
||||
elif hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
# MPS doesn't support per-device queries - uses global memory tracking
|
||||
allocated = torch.mps.current_allocated_memory() / (1024**3)
|
||||
reserved = torch.mps.driver_allocated_memory() / (1024**3)
|
||||
@@ -235,11 +235,11 @@ def clear_memory(debug: Optional['Debug'] = None, deep: bool = False, force: boo
|
||||
if free_ratio < 0.05:
|
||||
should_clear = True
|
||||
if debug:
|
||||
backend = "MPS" if torch.mps.is_available() else "VRAM"
|
||||
backend = "MPS" if (hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available()) else "VRAM"
|
||||
debug.log(f"{backend} pressure: {mem_info['free_gb']:.2f}GB free of {mem_info['total_gb']:.2f}GB", category="memory")
|
||||
|
||||
# For non-MPS systems, also check system RAM separately
|
||||
if not should_clear and not torch.mps.is_available():
|
||||
if not should_clear and not (hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available()):
|
||||
mem = psutil.virtual_memory()
|
||||
if mem.available < mem.total * 0.05:
|
||||
should_clear = True
|
||||
@@ -265,7 +265,7 @@ def clear_memory(debug: Optional['Debug'] = None, deep: bool = False, force: boo
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
elif torch.mps.is_available():
|
||||
elif hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
|
||||
if debug:
|
||||
@@ -302,7 +302,7 @@ def clear_memory(debug: Optional['Debug'] = None, deep: bool = False, force: boo
|
||||
handle = _os_memory_lib.GetCurrentProcess()
|
||||
_os_memory_lib.SetProcessWorkingSetSize(handle, -1, -1)
|
||||
|
||||
elif torch.mps.is_available():
|
||||
elif hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available():
|
||||
# macOS with MPS
|
||||
import ctypes # Import only when needed
|
||||
import ctypes.util
|
||||
|
||||
+3
-3
@@ -333,7 +333,7 @@ class Debug:
|
||||
}
|
||||
|
||||
# VRAM metrics
|
||||
if torch.cuda.is_available() or torch.mps.is_available():
|
||||
if torch.cuda.is_available() or (hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available()):
|
||||
metrics['vram_allocated'], metrics['vram_reserved'], current_global_peak = get_vram_usage(device=None, debug=self)
|
||||
|
||||
# Calculate peak since last log_memory_state
|
||||
@@ -346,7 +346,7 @@ class Debug:
|
||||
metrics['vram_free'] = vram_info["free_gb"]
|
||||
metrics['vram_total'] = vram_info["total_gb"]
|
||||
|
||||
backend = "MPS" if torch.mps.is_available() else "VRAM"
|
||||
backend = "MPS" if (hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available()) else "VRAM"
|
||||
metrics['summary_vram'] = (f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / "
|
||||
f"{metrics['vram_reserved']:.2f}GB reserved / "
|
||||
f"Peak: {metrics['vram_peak_since_last']:.2f}GB / "
|
||||
@@ -369,7 +369,7 @@ class Debug:
|
||||
metrics['summary_ram'] = ""
|
||||
|
||||
# Update VRAM history for tracking
|
||||
if torch.cuda.is_available() or torch.mps.is_available():
|
||||
if torch.cuda.is_available() or (hasattr(torch, 'mps') and callable(getattr(torch.mps, 'is_available', None)) and torch.mps.is_available()):
|
||||
self.vram_history.append(metrics['vram_allocated'])
|
||||
|
||||
return metrics
|
||||
|
||||
Reference in New Issue
Block a user