Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f68fe920b8 | ||
|
|
65c1c1b6cd | ||
|
|
b40f26167c | ||
|
|
ed53581359 | ||
|
|
71ac9ffe54 | ||
|
|
ffba05907d | ||
|
|
d4dd5e747d | ||
|
|
e2faedaaa6 | ||
|
|
5775ff0f99 | ||
|
|
ff937756a3 | ||
|
|
5848cef05f | ||
|
|
f5b902b8b0 | ||
|
|
19825fa9fa | ||
|
|
021fc7b70f | ||
|
|
b0ac880a4f | ||
|
|
6930e0f13a |
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user