Compare commits

...
16 Commits
Author SHA1 Message Date
Adrien Toupet f68fe920b8 Merge pull request #358 from AInVFX/main
v2.5.15: MPS compatibility fixes, autocast device type, VRAM tracking, triton 3.0+ compatibility
2025-12-03 13:16:46 -05:00
Adrien Toupet 65c1c1b6cd Release v2.5.15: MPS fixes, autocast device type, accurate VRAM tracking, triton 3.0 compatibility 2025-12-03 13:14:51 -05:00
Adrien Toupet b40f26167c Fix MPS compatibility: disable antialias for MPS tensors, fix bfloat16 arange (#354) 2025-12-03 13:09:59 -05:00
Adrien Toupet ed53581359 Fix triton.ops compatibility for bitsandbytes 0.45+ / triton 3.0+
Fixes #340 - Installation error with PyTorch 2.7+cu126 and triton_windows

Add compatibility shim for missing triton.ops.matmul_perf_model module.
Reverts local VAE types approach
2025-12-03 12:43:35 -05:00
Adrien Toupet 71ac9ffe54 fix: use max_memory_reserved for accurate VRAM peak tracking 2025-12-03 11:51:30 -05:00
Adrien Toupet ffba05907d Fix autocast device_type error by using .type attribute instead of str() #350 2025-12-03 11:32:39 -05:00
Adrien Toupet d4dd5e747d Merge pull request #344 from AInVFX/main
v2.5.14: MPS device fix, VRAM swap detection, enforce physical VRAM limit
2025-12-01 00:33:53 -05:00
Adrien Toupet e2faedaaa6 Release v2.5.14 - MPS device fix, VRAM swap detection, enforce physical VRAM limit 2025-12-01 00:30:49 -05:00
Adrien Toupet 5775ff0f99 Enforce VRAM limit to physical capacity - OOM instead of silent swap 2025-12-01 00:25:26 -05:00
Adrien Toupet ff937756a3 Add VRAM swap detection - show GPU+swap breakdown in peak stats, warn when swap detected 2025-11-30 23:47:20 -05:00
Adrien Toupet 5848cef05f fix(mps): normalize device strings to prevent unnecessary tensor movements
- Add _device_str() helper to normalize MPS variants (mps:0 → MPS)
- Fix device comparison: mps:0 and mps now correctly identified as same device
- Consistent MPS logging across all memory management functions
2025-11-30 21:03:26 -05:00
Adrien Toupet f5b902b8b0 Merge pull request #341 from AInVFX/main
v2.5.13: Fix triton import error, OOM on long video float32 conversion, macOS CLI watermark
2025-11-30 09:04:04 -05:00
Adrien Toupet 19825fa9fa Release v2.5.13: Fix triton import, OOM on long videos, macOS watermark 2025-11-30 09:01:53 -05:00
Adrien Toupet 021fc7b70f Fix triton.ops import error by using local VAE types
Fixes #340 - Installation error with PyTorch 2.7+cu126 and triton_windows

Replace diffusers.models.autoencoders.vae imports with local implementations
of DecoderOutput and DiagonalGaussianDistribution to avoid triggering the
bitsandbytes -> triton.ops import chain that fails on newer triton versions.
2025-11-30 08:27:01 -05:00
Adrien Toupet b0ac880a4f Fix OOM crash on float32 conversion for long videos. Gracefully fallback to native dtype if insufficient memory. Fixes #299 2025-11-30 08:03:07 -05:00
Adrien Toupet 6930e0f13a Fix CLI MPS watermark error on macOS (fixes #336) 2025-11-30 07:37:51 -05:00
14 changed files with 144 additions and 35 deletions
+19
View File
@@ -36,6 +36,25 @@ We're actively working on improvements and new features. To stay informed:
## 🚀 Updates
**2025.12.03 - Version 2.5.15**
- **🍎 Fix: MPS compatibility** - Disable antialias for MPS tensors and fix bfloat16 arange issues
- **⚡ Fix: Autocast device type** - Use proper device type attribute to prevent autocast errors
- **📊 Memory: Accurate VRAM tracking** - Use max_memory_reserved for more precise peak reporting
- **🔧 Fix: Triton compatibility** - Add shim for bitsandbytes 0.45+ / triton 3.0+ (fixes PyTorch 2.7 installation errors)
**2025.12.01 - Version 2.5.14**
- **🍎 Fix: MPS device comparison** - Normalize device strings to prevent unnecessary tensor movements
- **📊 Memory: VRAM swap detection** - Peak stats now show GPU+swap breakdown when overflow occurs, with warning when swap detected
- **🛡️ Memory: Enforce physical VRAM limit** - PyTorch now OOMs instead of silently swapping to shared memory (prevents extreme slowdowns on Windows)
**2025.11.30 - Version 2.5.13**
- **🔧 Fix: PyTorch 2.7+ triton import error** - Resolved installation crash caused by triton.ops import chain on newer triton versions
- **💾 Fix: OOM on float32 conversion for long videos** - Graceful fallback to native dtype when insufficient memory for float32 conversion
- **🍎 Fix: CLI watermark error on macOS** - Resolved MPS-related watermark processing crash on Apple Silicon
**2025.11.28 - Version 2.5.12**
- **🐛 Fix: Color artifacts regression** - Reverted in-place tensor operations in video transform pipeline that caused color artifacts on some images
+1
View File
@@ -3,6 +3,7 @@ ComfyUI-SeedVR2_VideoUpscaler
Official SeedVR2 integration for ComfyUI
"""
from .src.optimization.compatibility import ensure_triton_compat # noqa: F401
from .src.interfaces import comfy_entrypoint, SeedVR2Extension
__all__ = ["comfy_entrypoint", "SeedVR2Extension"]
+8 -2
View File
@@ -64,8 +64,14 @@ os.environ['PYTHONPATH'] = script_dir + ':' + os.environ.get('PYTHONPATH', '')
if mp.get_start_method(allow_none=True) != 'spawn':
mp.set_start_method('spawn', force=True)
# Configure VRAM management and validate CUDA devices before heavy imports
if platform.system() != "Darwin":
# Configure platform-specific memory management before heavy imports
# Must be set BEFORE import torch
if platform.system() == "Darwin":
# MPS allocator requires: low_watermark <= high_watermark
# Setting both to 0.0 disables PyTorch memory limits, letting macOS manage memory
os.environ.setdefault("PYTORCH_MPS_HIGH_WATERMARK_RATIO", "0.0")
os.environ.setdefault("PYTORCH_MPS_LOW_WATERMARK_RATIO", "0.0")
else:
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
# Pre-parse CUDA device argument for validation and environment setup
+1 -1
View File
@@ -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.12"
version = "2.5.15"
authors = [
{name = "numz"},
{name = "adrientoupet"}
@@ -36,7 +36,7 @@ class UniformTrailingSamplingTimesteps(SamplingTimesteps):
dtype: torch.dtype = torch.float32,
):
# Create trailing timesteps with specified dtype
timesteps = torch.arange(1.0, 0.0, -1.0 / steps, device=device, dtype=dtype)
timesteps = torch.arange(1.0, 0.0, -1.0 / steps, device='cpu').to(device=device, dtype=dtype)
# Shift timesteps.
timesteps = shift * timesteps / (1 + (shift - 1) * timesteps)
+1 -1
View File
@@ -711,7 +711,7 @@ def upscale_all_batches(
debug.start_timer(f"dit_inference_{upscale_idx+1}")
with torch.no_grad():
if dit_dtype != ctx['compute_dtype']:
with torch.autocast(str(ctx['dit_device']), ctx['compute_dtype'], enabled=True):
with torch.autocast(ctx['dit_device'].type, ctx['compute_dtype'], enabled=True):
upscaled_latents = runner.inference(
noises=noises,
conditions=conditions,
+2 -2
View File
@@ -155,7 +155,7 @@ class VideoDiffusionInfer():
# Use autocast if VAE dtype differs from input dtype
if vae_dtype != sample.dtype:
with torch.autocast(str(device), sample.dtype, enabled=True):
with torch.autocast(device.type, sample.dtype, enabled=True):
if use_sample:
latent = self.vae.encode(sample, tiled=self.encode_tiled, tile_size=self.encode_tile_size,
tile_overlap=self.encode_tile_overlap).latent
@@ -231,7 +231,7 @@ class VideoDiffusionInfer():
# Use autocast if VAE dtype differs from latent dtype
if vae_dtype != latent.dtype:
with torch.autocast(str(device), latent.dtype, enabled=True):
with torch.autocast(device.type, latent.dtype, enabled=True):
sample = self.vae.decode(
latent,
tiled=self.decode_tiled, tile_size=self.decode_tile_size,
+2
View File
@@ -50,10 +50,12 @@ class AreaResize:
resized_height, resized_width = round(height * scale), round(width * scale)
antialias = not (isinstance(image, torch.Tensor) and image.device.type == 'mps')
return TVF.resize(
image,
size=(resized_height, resized_width),
interpolation=self.interpolation,
antialias=antialias,
)
+4 -3
View File
@@ -56,8 +56,9 @@ class SideResize:
else:
size = self.size
# Resize to shortest edge
resized = TVF.resize(image, size, self.interpolation)
# Resize to shortest edge (disable antialias only for MPS tensors - not supported)
antialias = not (isinstance(image, torch.Tensor) and image.device.type == 'mps')
resized = TVF.resize(image, size, self.interpolation, antialias=antialias)
# Apply max_size constraint if specified
if self.max_size > 0:
@@ -69,6 +70,6 @@ class SideResize:
if max(h, w) > self.max_size:
scale = self.max_size / max(h, w)
new_h, new_w = round(h * scale), round(w * scale)
resized = TVF.resize(resized, (new_h, new_w), self.interpolation)
resized = TVF.resize(resized, (new_h, new_w), self.interpolation, antialias=antialias)
return resized
+9 -3
View File
@@ -509,15 +509,21 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
)
sample = ctx['final_video']
debug.log("", category="none", force=True)
# Ensure CPU tensor in float32 for maximum ComfyUI compatibility
if torch.is_tensor(sample):
if sample.is_cuda or sample.is_mps:
sample = sample.cpu()
if sample.dtype != torch.float32:
sample = sample.to(torch.float32)
src_dtype = sample.dtype
try:
sample = sample.to(torch.float32)
debug.log(f"Converted output from {src_dtype} to float32", category="precision")
except Exception as e:
debug.log(f"Could not convert to float32: {e}. Output is {src_dtype}, compatibility with other nodes not guaranteed",
level="WARNING", category="precision", force=True)
debug.log("", category="none", force=True)
debug.log("Upscaling completed successfully!", category="success", force=True)
debug.end_timer("generation", "Video generation")
+31
View File
@@ -5,6 +5,37 @@ Contains FP8/FP16 compatibility layers and wrappers for different model architec
Extracted from: seedvr2.py (lines 1045-1630)
"""
# Triton compatibility shim for bitsandbytes 0.45+ with triton 3.0+
# Must be called before any diffusers import
import sys
def ensure_triton_compat():
"""Create minimal triton.ops stubs only if missing, to allow bitsandbytes import."""
if 'triton.ops.matmul_perf_model' in sys.modules:
return
try:
from triton.ops.matmul_perf_model import early_config_prune # noqa: F401
return
except (ImportError, ModuleNotFoundError, AttributeError):
pass
import types
if 'triton.ops' not in sys.modules:
sys.modules['triton.ops'] = types.ModuleType('triton.ops')
matmul_perf = types.ModuleType('triton.ops.matmul_perf_model')
matmul_perf.early_config_prune = lambda configs, *a, **kw: configs
matmul_perf.estimate_matmul_time = lambda *a, **kw: 0.0
sys.modules['triton.ops'].matmul_perf_model = matmul_perf
sys.modules['triton.ops.matmul_perf_model'] = matmul_perf
# Run immediately on import
ensure_triton_compat()
import torch
import types
import os
+39 -17
View File
@@ -13,6 +13,12 @@ import psutil
from typing import Tuple, Dict, Any, Optional, List, Union
def _device_str(device: Union[torch.device, str]) -> str:
"""Normalized uppercase device string for comparison and logging. MPS variants → 'MPS'."""
s = str(device).upper()
return 'MPS' if s.startswith('MPS') else s
def get_device_list(include_none: bool = False, include_cpu: bool = False) -> List[str]:
"""
Get list of available compute devices for SeedVR2
@@ -106,6 +112,22 @@ else:
print(f"⚠️ Memory check failed: {vram_info['error']} - No available backend!")
def _enforce_vram_limit() -> None:
"""
Enforce VRAM limit to physical capacity to prevent silent swap to system RAM.
Called once at module load. No-op on MPS or unsupported platforms.
"""
if not torch.cuda.is_available():
return
try:
for i in range(torch.cuda.device_count()):
torch.cuda.set_per_process_memory_fraction(1.0, i)
except Exception:
pass
_enforce_vram_limit()
def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float]:
"""
Get current VRAM usage metrics for monitoring.
@@ -116,7 +138,7 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug
debug: Optional debug instance for logging
Returns:
tuple: (allocated_gb, reserved_gb, max_allocated_gb)
tuple: (allocated_gb, reserved_gb, max_reserved_gb)
Returns (0, 0, 0) if no GPU available
"""
try:
@@ -127,8 +149,8 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug
device = torch.device(device)
allocated = torch.cuda.memory_allocated(device) / (1024**3)
reserved = torch.cuda.memory_reserved(device) / (1024**3)
max_allocated = torch.cuda.max_memory_allocated(device) / (1024**3)
return allocated, reserved, max_allocated
max_reserved = torch.cuda.max_memory_reserved(device) / (1024**3)
return allocated, reserved, max_reserved
elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
# MPS doesn't support per-device queries - uses global memory tracking
allocated = torch.mps.current_allocated_memory() / (1024**3)
@@ -591,7 +613,7 @@ def manage_tensor(
target_dtype = dtype if dtype is not None else current_dtype
# Check if movement is actually needed
needs_device_move = current_device != target_device
needs_device_move = _device_str(current_device) != _device_str(target_device)
needs_dtype_change = dtype is not None and current_dtype != target_dtype
if not needs_device_move and not needs_dtype_change:
@@ -609,8 +631,8 @@ def manage_tensor(
# Log the movement
if debug:
current_device_str = str(current_device).upper()
target_device_str = str(target_device).upper()
current_device_str = _device_str(current_device)
target_device_str = _device_str(target_device)
dtype_info = ""
if needs_dtype_change:
@@ -681,8 +703,8 @@ def manage_model_device(model: torch.nn.Module, target_device: torch.device, mod
# Extract device type for comparison (both are torch.device objects)
target_type = target_device.type
current_device_upper = str(current_device).upper()
target_device_upper = str(target_device).upper()
current_device_upper = _device_str(current_device)
target_device_upper = _device_str(target_device)
# Compare normalized device types
if current_device_upper == target_device_upper and not is_blockswap_model:
@@ -737,10 +759,10 @@ def _handle_blockswap_model_movement(runner: Any, model: torch.nn.Module,
actual_source_device = param.device
break
source_device_desc = str(actual_source_device).upper() if actual_source_device else str(target_device).upper()
source_device_desc = _device_str(actual_source_device) if actual_source_device else _device_str(target_device)
if debug:
debug.log(f"Moving {model_name} from {source_device_desc} to {str(target_device).upper()} ({reason or 'model caching'})", category="general")
debug.log(f"Moving {model_name} from {source_device_desc} to {_device_str(target_device)} ({reason or 'model caching'})", category="general")
# Enable bypass to allow movement
set_blockswap_bypass(runner=runner, bypass=True, debug=debug)
@@ -755,7 +777,7 @@ def _handle_blockswap_model_movement(runner: Any, model: torch.nn.Module,
model.zero_grad(set_to_none=True)
if debug:
debug.end_timer(timer_name, f"BlockSwap model offloaded to {str(target_device).upper()}")
debug.end_timer(timer_name, f"BlockSwap model offloaded to {_device_str(target_device)}")
return True
@@ -775,10 +797,10 @@ def _handle_blockswap_model_movement(runner: Any, model: torch.nn.Module,
actual_current_device = param.device
break
current_device_desc = str(actual_current_device).upper() if actual_current_device else "OFFLOAD"
current_device_desc = _device_str(actual_current_device) if actual_current_device else "OFFLOAD"
if debug:
debug.log(f"Moving {model_name} from {current_device_desc} to {str(target_device).upper()} ({reason or 'inference requirement'})", category="general")
debug.log(f"Moving {model_name} from {current_device_desc} to {_device_str(target_device)} ({reason or 'inference requirement'})", category="general")
timer_name = f"{model_name.lower()}_to_gpu"
if debug:
@@ -818,7 +840,7 @@ def _handle_blockswap_model_movement(runner: Any, model: torch.nn.Module,
blocks_on_gpu = model._block_swap_config.get('total_blocks', 32) - model._block_swap_config.get('blocks_swapped', 16)
total_blocks = model._block_swap_config.get('total_blocks', 32)
main_device = model._block_swap_config.get('main_device', 'GPU')
debug.log(f"BlockSwap blocks restored to configured devices ({blocks_on_gpu}/{total_blocks} blocks on {str(main_device).upper()})", category="success")
debug.log(f"BlockSwap blocks restored to configured devices ({blocks_on_gpu}/{total_blocks} blocks on {_device_str(main_device)})", category="success")
else:
debug.log("BlockSwap blocks restored to configured devices", category="success")
@@ -865,8 +887,8 @@ def _standard_model_movement(model: torch.nn.Module, current_device: torch.devic
# Log the movement with full device strings
if debug:
current_device_str = str(current_device).upper()
target_device_str = str(target_device).upper()
current_device_str = _device_str(current_device)
target_device_str = _device_str(target_device)
debug.log(f"Moving {model_name} from {current_device_str} to {target_device_str} ({reason})", category="general")
# Start timer based on direction
@@ -891,7 +913,7 @@ def _standard_model_movement(model: torch.nn.Module, current_device: torch.devic
# End timer
if debug:
debug.end_timer(timer_name, f"{model_name} moved to {str(target_device).upper()}")
debug.end_timer(timer_name, f"{model_name} moved to {_device_str(target_device)}")
return True
+1 -1
View File
@@ -4,7 +4,7 @@ Only includes constants actually used in the codebase
"""
# Version information
__version__ = "2.5.12"
__version__ = "2.5.15"
import os
import warnings
+25 -4
View File
@@ -14,6 +14,14 @@ from ..optimization.memory_manager import get_vram_usage, get_basic_vram_info, g
from ..utils.constants import __version__
def _format_peak_with_swap(peak_gb: float, total_vram_gb: float) -> str:
"""Format peak memory, showing swap breakdown if overflow occurred."""
if total_vram_gb > 0 and peak_gb > total_vram_gb:
swap_gb = peak_gb - total_vram_gb
return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {swap_gb:.2f}GB swap)"
return f"{peak_gb:.2f}GB"
class Debug:
"""
Unified debug logging for generation pipeline and BlockSwap monitoring
@@ -307,7 +315,12 @@ class Debug:
if show_diff and self.memory_checkpoints:
self._log_memory_diff(current_metrics=memory_info, force=force)
# Log detailed analysis if requested
# Warn if swap detected (peak > physical VRAM)
if memory_info['vram_total'] > 0 and memory_info['vram_peak_since_last'] > memory_info['vram_total']:
self.log("VRAM swap detected - severe slowdown expected. Consider optimizing (e.g., reduce resolution, batch_size, enable BlockSwap, VAE tiling...).",
level="WARNING", category="memory", force=True)
# Log detailed analysis if requested
if detailed_tensors and tensor_stats.get('details'):
self._log_detailed_tensor_analysis(details=tensor_stats['details'], force=force)
@@ -361,9 +374,10 @@ class Debug:
metrics['vram_total'] = vram_info["total_gb"]
backend = "MPS" if (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()) else "VRAM"
peak_str = _format_peak_with_swap(metrics['vram_peak_since_last'], metrics['vram_total'])
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 / "
f"Peak: {peak_str} / "
f"{metrics['vram_free']:.2f}GB free / "
f"{metrics['vram_total']:.2f}GB total")
else:
@@ -525,6 +539,13 @@ class Debug:
is_mps = hasattr(torch.backends, 'mps') and torch.backends.mps.is_available() and not torch.cuda.is_available()
# Get total VRAM for swap detection (reuse existing function)
total_vram_gb = 0.0
if not is_mps:
vram_info = get_basic_vram_info(device=None)
if "error" not in vram_info:
total_vram_gb = vram_info["total_gb"]
self.log("", category="none", force=force)
self.log("────────────────────────", category="none", force=force)
self.log("Peak memory by phase:", category="memory", force=force)
@@ -539,7 +560,7 @@ class Debug:
if is_mps:
self.log(f" Phase {phase_num} ({phase_name}): {vram:.2f}GB", category="memory", force=force)
else:
self.log(f" Phase {phase_num} ({phase_name}): VRAM {vram:.2f}GB | RAM {ram:.2f}GB", category="memory", force=force)
self.log(f" Phase {phase_num} ({phase_name}): {_format_peak_with_swap(vram, total_vram_gb)} | RAM {ram:.2f}GB", category="memory", force=force)
if is_mps:
overall = max(self.phase_vram_peaks.values()) if self.phase_vram_peaks else 0
@@ -547,7 +568,7 @@ class Debug:
else:
overall_vram = max(self.phase_vram_peaks.values()) if self.phase_vram_peaks else 0
overall_ram = max(self.phase_ram_peaks.values()) if self.phase_ram_peaks else 0
self.log(f"Overall peak: VRAM {overall_vram:.2f}GB | RAM {overall_ram:.2f}GB", category="memory", force=force)
self.log(f"Overall peak: {_format_peak_with_swap(overall_vram, total_vram_gb)} | RAM {overall_ram:.2f}GB", category="memory", force=force)
@torch._dynamo.disable # Skip tracing to avoid time.time() warnings
def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None: