refactor(WIP): Memory Management Overhaul

This is a work-in-progress commit that consolidates a series of changes to fix memory leaks, optimize VRAM/RAM usage, improve performance, and enhance code maintainability.

*   **BlockSwap Pinned Memory:** Disabled `use_non_blocking=True` for CPU-to-GPU transfers to resolve a memory leak where pinned memory was not being released.
*   **Logging-Induced Leaks:** Modified `log_memory_state()` to avoid holding references to tensors during analysis and added a history limit to the checkpoint system to prevent unbounded memory growth.
*   **Incomplete Model Cleanup:** Ensured models are completely deleted and their tensor storage is released when `cache_model=False`.
*   **Lingering Tensors:** Fixed an issue where a scalar tensor from sampling timesteps and text embeddings remained on the GPU between batches when `preserve_vram` is active.

*   **Centralized Cleanup Functions:** Introduced `clear_memory()` to replace `clear_vram_cache()` and all manual `torch.cuda.empty_cache()` calls, providing consistent VRAM/RAM cleanup logic. The function features a `full` parameter to distinguish between a fast, GPU-only cache clear (~1-5ms) for frequent operations and a full cleanup with garbage collection (~10-50ms) for critical stages.
*   **Direct-to-CPU Model Loading:** Modified DiT/VAE weight loading to load directly onto the CPU when `preserve_vram` or `BlockSwap` is active, avoiding unnecessary VRAM spikes during model preparation.
*   **VAE Device Management:** Created the `manage_vae_device()` helper function to centralize the logic for moving the VAE between the CPU and GPU, reducing code duplication. This also fixed a bug that incorrectly kept the VAE on the GPU when `preserve_vram` was active.
*   **CPU Offloading:** Implemented logic to move text embeddings and sampling timesteps to the CPU after each batch when `preserve_vram` is active, reducing idle VRAM usage.

*   **VAE Decode Performance:** Replaced proactive, frequent memory clearing during VAE decode with a reactive Out-of-Memory (OOM) handling system. This fixed a significant performance regression and eliminated the need for the `keep_vae_loaded_during_decode` flag.
*   **Reduced Overhead:** Removed redundant `gc.collect()` calls from multiple locations to decrease unnecessary processing overhead.

*   **Interval-Based VRAM Tracking:** Modified the logging system to reset peak VRAM statistics after each `log_memory_state()` call, enabling accurate tracking of peak memory usage for specific processing intervals (e.g., encode, inference, decode).
*   **Accurate RAM Monitoring:** Added the `get_ram_usage()` function for correct process-specific RAM tracking.
*   **Efficient Log Refactoring:** Refactored `log_memory_state()` into modular helper methods, optimizing tensor analysis into a single-pass `gc` iteration to improve both performance and maintainability.
*   **Log Clarity:** Refined memory state and debug logging to remove redundant snapshots and add new ones for critical operations like model loading, weight loading, VAE encoding, and decoding. Standardized log message conventions.
*   **Per-Batch Timers:** Implemented timer namespacing to ensure that performance timers for each batch are logged correctly without overwriting one another.

*   **Error Handling:** Added `try/except` blocks to key memory and device management functions to handle edge cases and improve robustness.
*   **Code Cleanup:** Removed deprecated code and outdated comments throughout the related modules.
*   **Documentation:** Updated comments and function docstrings to reflect the new memory management architecture.
This commit is contained in:
Adrien Toupet
2025-08-22 09:20:41 -04:00
parent 0373a81d2b
commit 9b7c113681
12 changed files with 1032 additions and 753 deletions
+2 -2
View File
@@ -54,7 +54,7 @@ if MODULES_AVAILABLE['downloads']:
if MODULES_AVAILABLE['memory_manager']:
from src.optimization.memory_manager import (
get_vram_usage,
clear_vram_cache,
clear_memory,
reset_vram_peak,
preinitialize_rope_cache,
)
@@ -122,7 +122,7 @@ __all__ = [
'download_weight',
# Memory Management
'get_vram_usage', 'clear_vram_cache', 'reset_vram_peak',
'get_vram_usage', 'clear_memory', 'reset_vram_peak',
'preinitialize_rope_cache',
# Performance & Video Processing
+4 -7
View File
@@ -21,6 +21,7 @@ from typing import Callable
import torch
from einops import rearrange
from torch.nn import functional as F
from src.optimization.memory_manager import clear_memory
#from ....models.dit_v2 import na
@@ -71,13 +72,9 @@ class EulerSampler(Sampler):
# Nettoyer les tenseurs temporaires
del pred
if torch.mps.is_available():
if torch.mps.is_available():
torch.mps.empty_cache()
else:
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
# Use debug if available from the sampler
debug = getattr(self, 'debug', None)
clear_memory(debug=debug, full=False, force=True)
i += 1
progress.update()
+164 -112
View File
@@ -26,7 +26,7 @@ from src.common.distributed import get_device
# Import required modules
from src.optimization.memory_manager import reset_vram_peak, clear_all_caches
from src.optimization.memory_manager import manage_vae_device, clear_all_caches, clear_memory
from src.optimization.performance import (
optimized_video_rearrange, optimized_single_video_rearrange,
optimized_sample_to_image_format, temporal_latent_blending
@@ -114,6 +114,10 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
cond_noise_scale = 0.0
def _add_noise(x, aug_noise):
# Early return if no noise is being added
if cond_noise_scale == 0.0:
return x
# Use adaptive optimal dtype
t = (
torch.tensor([1000.0], device=device, dtype=dtype)
@@ -122,6 +126,10 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
shape = torch.tensor(x.shape[1:], device=device)[None]
t = runner.timestep_transform(t, shape)
x = runner.schedule.forward(x, aug_noise, t)
# Explicit cleanup of intermediate tensors
del t, shape
return x
# Generate conditions with memory optimization
@@ -137,17 +145,33 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
# Use adaptive autocast for optimal performance
with torch.no_grad():
# Restore timesteps to GPU if they were offloaded
if preserve_vram and hasattr(runner, 'sampling_timesteps') and hasattr(runner.sampling_timesteps, 'timesteps'):
if not runner.sampling_timesteps.timesteps.is_cuda:
debug.log("Restoring timesteps tensor to GPU (preserve_vram)", category="memory")
debug.start_timer("timesteps_to_gpu")
runner.sampling_timesteps.timesteps = runner.sampling_timesteps.timesteps.to(device, non_blocking=True)
debug.end_timer("timesteps_to_gpu", "Sampling timesteps restored to GPU")
with torch.autocast(str(get_device()), autocast_dtype, enabled=True):
video_tensors = runner.inference(
noises=noises,
conditions=conditions,
preserve_vram=preserve_vram # Memory offload optimization
and not use_blockswap, # Disable dit_offload if BlockSwap active
preserve_vram=preserve_vram, # Memory offload optimization
temporal_overlap=temporal_overlap,
use_blockswap=use_blockswap,
**text_embeds_dict,
)
# Clean up diffusion timesteps from GPU if preserve_vram is enabled
if preserve_vram:
if hasattr(runner, 'sampling_timesteps') and hasattr(runner.sampling_timesteps, 'timesteps'):
if runner.sampling_timesteps.timesteps.is_cuda:
debug.log("Moving timesteps tensor to CPU (preserve_vram)", category="memory")
debug.start_timer("timesteps_to_cpu")
runner.sampling_timesteps.timesteps = runner.sampling_timesteps.timesteps.cpu()
debug.end_timer("timesteps_to_cpu", "Sampling timesteps offloaded to CPU")
# Process samples with advanced optimization
samples = optimized_video_rearrange(video_tensors)
#last_latents = samples[-temporal_overlap:] if temporal_overlap > 0 else samples[-1:]
@@ -283,8 +307,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
# Set random seed
set_seed(seed)
debug.log_memory_state("Model configuration - Memory")
debug.end_timer("model_config", "Model configuration completed", show_breakdown=True)
debug.log_memory_state("After model configuration", detailed_tensors=False)
# ───────────────────────────────────────────────────────────────
# Step 2: Input Preparation & Transformation Setup
@@ -315,8 +339,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
text_neg_embeds = torch.load(os.path.join(script_directory, 'neg_emb.pt')).to(device, dtype=compute_dtype)
text_embeds = {"texts_pos": [text_pos_embeds], "texts_neg": [text_neg_embeds]}
debug.log_memory_state("Input preparation - Memory ")
debug.end_timer("input_prep", "Input preparation completed", show_breakdown=True)
debug.log_memory_state("After input preparation", detailed_tensors=False)
# ───────────────────────────────────────────────────────────────
# Step 3: Batch Processing
@@ -325,8 +349,6 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
debug.start_timer("batch_processing")
# Standard processing (non-TileVAE) continues below
# Memory optimization
reset_vram_peak(debug)
# Calculate processing parameters
step = batch_size - temporal_overlap
@@ -364,103 +386,124 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
break # Not enough new frames, stop
batch_number = (batch_idx // step + 1) if step > 0 else 1
debug.start_timer(f"batch_{batch_number}")
current_frames = end_idx - start_idx
debug.log("", category="none")
debug.log(f"━━━ Batch {batch_number}/{total_batches}: frames {start_idx}-{end_idx-1} ━━━", category="generation", force=True)
debug.log_memory_state(f"Before batch {batch_number} processing", detailed_tensors=False)
# Use timer context for this batch - all timers within will be namespaced
with debug.timer_context(f"batch_{batch_number}"):
debug.start_timer("batch") # This becomes "batch_1_batch" internally
# Process current batch
video = images[start_idx:end_idx]
debug.log(f"Video compute dtype: {compute_dtype}", category="generation")
# Use adaptive computation dtype
video = video.permute(0, 3, 1, 2).to(device, dtype=compute_dtype)
# Apply video transformations with memory optimization
transformed_video = video_transform(video)
del video
#video = video.to("cpu")
#del video
ori_lengths = [transformed_video.size(1)]
# Handle correct format: frames % 4 == 1
t = transformed_video.size(1)
debug.log(f"Sequence of {t} frames", category="video", force=True)
if len(images) >= 5 and t % 4 != 1:
debug.log(f"Transformed video shape before cut: {transformed_video.shape}", category="video")
transformed_video = cut_videos(transformed_video)
debug.log(f"Transformed video shape: {transformed_video.shape}", category="video")
# Context-aware temporal strategy
# First batch: standard complete diffusion
debug.start_timer("vae_to_gpu")
runner.vae = runner.vae.to(device)
debug.end_timer("vae_to_gpu", "VAE to GPU")
debug.start_timer("vae_encode")
debug.log(f"VAE encoding precision: {autocast_dtype}", category="vae")
with torch.autocast(str(device), autocast_dtype, enabled=True):
cond_latents = runner.vae_encode([transformed_video])
debug.end_timer("vae_encode", "VAE encoding")
#tps = time.time()
#transformed_video = transformed_video.to("cpu")
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
debug.log(f"Cond latents shape: {cond_latents[0].shape}", category="info")
# Normal generation
samples = generation_step(runner, text_embeds, preserve_vram, cond_latents=cond_latents, temporal_overlap=temporal_overlap, debug=debug)
#del cond_latents
del cond_latents
# Post-process samples
sample = samples[0]
del samples
#del samples
if ori_lengths[0] < sample.shape[0]:
sample = sample[:ori_lengths[0]]
#if temporal_overlap > 0 and not is_first_batch and sample.shape[0] > effective_batch_size - temporal_overlap:
# sample = sample[temporal_overlap:] # Remove overlap frames from output
# Apply color correction if available
debug.start_timer("video_to_device")
transformed_video = transformed_video.to(device)
debug.end_timer("video_to_device", "Transformed video to device")
input_video = [optimized_single_video_rearrange(transformed_video)]
del transformed_video
#transformed_video = transformed_video.to("cpu")
#del transformed_video
sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)], debug)
del input_video
# Move text embeddings back to GPU if they were offloaded
if preserve_vram or (block_swap_config and block_swap_config.get("offload_io_components", False)):
if text_pos_embeds.device.type == "cpu":
reason = "preserve_vram" if preserve_vram else "BlockSwap I/O offload"
debug.log(f"Restoring text embeddings to GPU ({reason} active)", category="memory")
debug.start_timer("text_embeddings_to_gpu")
text_pos_embeds = text_pos_embeds.to(device, dtype=compute_dtype)
text_neg_embeds = text_neg_embeds.to(device, dtype=compute_dtype)
text_embeds["texts_pos"][0] = text_pos_embeds
text_embeds["texts_neg"][0] = text_neg_embeds
debug.end_timer("text_embeddings_to_gpu", "Text embeddings restored to GPU")
# Process current batch
video = images[start_idx:end_idx]
debug.log(f"Video compute dtype: {compute_dtype}", category="precision")
# Use adaptive computation dtype
video = video.permute(0, 3, 1, 2).to(device, dtype=compute_dtype)
# Apply video transformations with memory optimization
transformed_video = video_transform(video)
del video
#video = video.to("cpu")
#del video
ori_lengths = [transformed_video.size(1)]
# Handle correct format: frames % 4 == 1
t = transformed_video.size(1)
debug.log(f"Sequence of {t} frames", category="video", force=True)
if len(images) >= 5 and t % 4 != 1:
debug.log(f"Transformed video shape before cut: {transformed_video.shape}", category="video")
transformed_video = cut_videos(transformed_video)
debug.log(f"Transformed video shape: {transformed_video.shape}", category="video")
# Context-aware temporal strategy
# First batch: standard complete diffusion
# Convert to final image format
sample = optimized_sample_to_image_format(sample)
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
sample_cpu = sample.to(torch.float16).to("cpu")
del sample
batch_samples.append(sample_cpu)
#del sample
# Aggressive cleanup after each batch
# tps = time.time()
# Progress callback - batch start
if progress_callback:
progress_callback(batch_count+1, total_batches, current_frames, "Processing batch...")
#transformed_video = transformed_video.to("cpu")
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
# Clean VRAM after each batch when preserve_vram is active
if preserve_vram:
# Only offload the VAE when we are not keeping it resident
offload = not getattr(runner, 'keep_vae_in_vram', False)
clear_all_caches(runner, debug, offload_vae=offload)
#del transformed_video
#clear_vram_cache()
# Log memory state at the end of each batch
debug.log_memory_state(f"Batch {batch_number} - Memory")
debug.end_timer(f"batch_{batch_number}", f"Batch {batch_number} processed", show_breakdown=True)
# Move VAE to GPU if needed for encoding
manage_vae_device(runner, str(device), preserve_vram=False, debug=debug)
debug.log(f"VAE encoding precision: {autocast_dtype}", category="precision")
debug.log("Encoding video to latents...", category="vae")
debug.start_timer("vae_encoding")
with torch.autocast(str(device), autocast_dtype, enabled=True):
cond_latents = runner.vae_encode([transformed_video])
debug.end_timer("vae_encoding", "VAE encoding")
debug.log(f"Cond latents shape: {cond_latents[0].shape}", category="info")
# Move VAE back to CPU after encoding if preserve_vram is enabled
if preserve_vram:
manage_vae_device(runner, 'cpu', preserve_vram=preserve_vram, debug=debug)
debug.log_memory_state("After VAE encode", detailed_tensors=False)
# Normal generation
samples = generation_step(runner, text_embeds, preserve_vram, cond_latents=cond_latents, temporal_overlap=temporal_overlap, debug=debug)
#del cond_latents
del cond_latents
# Post-process samples
sample = samples[0]
del samples
#del samples
if ori_lengths[0] < sample.shape[0]:
sample = sample[:ori_lengths[0]]
#if temporal_overlap > 0 and not is_first_batch and sample.shape[0] > effective_batch_size - temporal_overlap:
# sample = sample[temporal_overlap:] # Remove overlap frames from output
# Apply color correction if available
debug.start_timer("video_to_device")
transformed_video = transformed_video.to(device)
debug.end_timer("video_to_device", "Transformed video to device")
input_video = [optimized_single_video_rearrange(transformed_video)]
del transformed_video
#transformed_video = transformed_video.to("cpu")
#del transformed_video
sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)], debug)
del input_video
# Convert to final image format
sample = optimized_sample_to_image_format(sample)
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
sample_cpu = sample.to(torch.float16).to("cpu")
del sample
batch_samples.append(sample_cpu)
# Aggressive cleanup after each batch
# tps = time.time()
# Progress callback - batch start
if progress_callback:
progress_callback(batch_count+1, total_batches, current_frames, "Processing batch...")
#transformed_video = transformed_video.to("cpu")
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
# Clean VRAM after each batch when preserve_vram is active
if preserve_vram or (block_swap_config and block_swap_config.get("offload_io_components", False)):
# Move text embeddings to CPU when using memory-saving features
reason = "preserve_vram" if preserve_vram else "BlockSwap I/O offload"
debug.log(f"Moving text embeddings to CPU ({reason})", category="memory")
debug.start_timer("text_embeddings_to_cpu")
text_pos_embeds = text_pos_embeds.to("cpu")
text_neg_embeds = text_neg_embeds.to("cpu")
text_embeds["texts_pos"][0] = text_pos_embeds
text_embeds["texts_neg"][0] = text_neg_embeds
debug.end_timer("text_embeddings_to_cpu", "Text embeddings moved to CPU")
# Log memory state at the end of each batch
debug.end_timer("batch", f"Batch {batch_number} processed", show_breakdown=True)
debug.log_memory_state(f"After batch {batch_number} processing", detailed_tensors=True)
finally:
debug.log("", category="none")
@@ -468,22 +511,24 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
debug.start_timer("generation_cleanup")
# Final cleanup of embeddings
debug.log("Moving text embeddings to CPU (final cleanup)", category="memory")
text_pos_embeds = text_pos_embeds.to("cpu")
text_neg_embeds = text_neg_embeds.to("cpu")
# Move DiT to CPU
debug.log("Moving DiT to CPU (final cleanup)", category="memory")
debug.start_timer("dit_to_cpu_cleanup")
runner.dit.to("cpu")
if not getattr(runner, 'keep_vae_in_vram', False):
runner.vae.to("cpu")
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
#del text_pos_embeds, text_neg_embeds
#clear_vram_cache()
debug.end_timer("dit_to_cpu_cleanup", "DiT moved to CPU (final cleanup)")
# Move VAE to CPU
manage_vae_device(runner, 'cpu', preserve_vram=False, debug=debug, reason="final cleanup")
clear_memory(debug=debug, full=True, force=True)
# Log final memory state
debug.log_memory_state("Generation cleanup - Memory")
debug.end_timer("generation_cleanup", "Batch generation cleanup")
debug.log_memory_state("After batch generation cleanup", detailed_tensors=False)
debug.end_timer("batch_processing", "Batch processing completed", show_breakdown=True)
@@ -535,6 +580,15 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
# Clean up merged batch memory
del batch_group, merged_result
# Clean up batch_samples list completely
for batch in batch_samples:
if torch.is_tensor(batch):
if batch.is_cuda:
batch.cpu()
del batch
batch_samples.clear()
del batch_samples
debug.log(f"Memory pre-allocation completed for output tensor: {final_video_images.shape}", category="success")
debug.log("Pre-allocation ensures contiguous memory for final video output", category="info")
@@ -565,11 +619,9 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
debug.log(f" Most swapped: Block {swap_summary['most_swapped_block']} "
f"({swap_summary['most_swapped_count']} times)", category="blockswap")
debug.log_memory_state("Post-processing - Memory")
debug.end_timer("post_processing", "Post-processing completed", show_breakdown=True)
debug.log_memory_state("After post-processing", detailed_tensors=False)
# Cleanup batch_samples
#del batch_samples
return final_video_images
@@ -661,4 +713,4 @@ def calculate_optimal_batch_params(total_frames, batch_size, temporal_overlap):
'best_batch': best_batch,
'padding_waste': padding_waste,
'is_optimal': batch_size in optimal_batches
}
}
+28 -82
View File
@@ -18,7 +18,7 @@ import torch
from einops import rearrange
from omegaconf import DictConfig, ListConfig
from torch import Tensor
from src.optimization.memory_manager import clear_vram_cache
from src.optimization.memory_manager import clear_memory, manage_vae_device
from src.common.diffusion import (
classifier_free_guidance_dispatcher,
@@ -68,7 +68,7 @@ def optimized_channels_to_second(tensor):
return tensor.permute(*dims)
class VideoDiffusionInfer():
def __init__(self, config: DictConfig, debug=None, vae_tiling_enabled: bool = False,
def __init__(self, config: DictConfig, debug=None, vae_tiling_enabled: bool = False,
vae_tile_size: Tuple[int, int] = (512, 512), vae_tile_overlap: Tuple[int, int] = (64, 64)):
# Check if debug instance is available
if debug is None:
@@ -78,9 +78,6 @@ class VideoDiffusionInfer():
self.vae_tiling_enabled = vae_tiling_enabled
self.vae_tile_size = vae_tile_size
self.vae_tile_overlap = vae_tile_overlap
# Keep the VAE on the GPU between decode calls
self.keep_vae_in_vram: bool = False
def get_condition(self, latent: Tensor, latent_blur: Tensor, task: str) -> Tensor:
t, h, w, c = latent.shape
@@ -187,7 +184,6 @@ class VideoDiffusionInfer():
"""🚀 VAE decode optimisé - décodage direct sans chunking, compatible avec autocast externe"""
samples = []
if len(latents) > 0:
#t = time.time()
device = get_device()
dtype = getattr(torch, self.config.vae.dtype)
scale = self.config.vae.scaling_factor
@@ -209,9 +205,7 @@ class VideoDiffusionInfer():
self.debug.log(f"Using VAE Tiled Decoding (Tile: {self.vae_tile_size}, Overlap: {self.vae_tile_overlap})", category="vae", force=True)
self.debug.log(f"Latents batch shape: {latents[0].shape}", category="info")
self.debug.start_timer("vae_decode")
# If the user wants to keep the VAE resident, do not let the VAE free its buffers
internal_preserve_vram = preserve_vram and not getattr(self, 'keep_vae_in_vram', False)
for i, latent in enumerate(latents):
effective_dtype = target_dtype if target_dtype is not None else dtype
latent = latent.to(device, effective_dtype, non_blocking=True)
@@ -220,7 +214,7 @@ class VideoDiffusionInfer():
latent = latent.squeeze(2)
sample = self.vae.decode(
latent, preserve_vram=internal_preserve_vram,
latent, preserve_vram=preserve_vram,
tiled=use_tiling, tile_size=self.vae_tile_size,
tile_overlap=self.vae_tile_overlap).sample
@@ -229,8 +223,6 @@ class VideoDiffusionInfer():
samples.append(sample)
self.debug.end_timer("vae_decode", "VAE decode completed")
if self.config.vae.grouping:
samples = na.unpack(samples, indices)
else:
@@ -271,30 +263,6 @@ class VideoDiffusionInfer():
timesteps = timesteps * self.schedule.T
return timesteps
def get_vram_usage(self):
"""Obtenir l'utilisation VRAM actuelle (allouée et réservée)"""
if torch.mps.is_available():
allocated = torch.mps.current_allocated_memory() / (1024**3)
reserved = torch.mps.driver_allocated_memory() / (1024**3)
max_allocated = 0
return allocated, reserved, max_allocated
if torch.cuda.is_available():
allocated = torch.cuda.memory_allocated() / (1024**3)
reserved = torch.cuda.memory_reserved() / (1024**3)
max_allocated = torch.cuda.max_memory_allocated() / (1024**3)
return allocated, reserved, max_allocated
return 0, 0, 0
def get_vram_peak(self):
"""Obtenir le pic VRAM depuis le dernier reset"""
if torch.cuda.is_available():
return torch.cuda.max_memory_allocated() / (1024**3)
return 0
def reset_vram_peak(self):
"""Reset le compteur de pic VRAM"""
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
@torch.no_grad()
def inference(
@@ -314,9 +282,6 @@ class VideoDiffusionInfer():
# Return if empty.
if batch_size == 0:
return []
# Monitoring VRAM initial et reset des pics
#self.reset_vram_peak()
# Set cfg scale
if cfg_scale is None:
@@ -369,12 +334,9 @@ class VideoDiffusionInfer():
if preserve_vram:
if conditions[0].shape[0] > 1:
self.debug.start_timer("vae_to_cpu")
self.vae = self.vae.to("cpu")
self.debug.end_timer("vae_to_cpu", "VAE to CPU")
# Before sampling, check if BlockSwap is active
if not use_blockswap and not hasattr(self, "_blockswap_active"):
self.debug.log("Moving DiT to GPU (inference requirement)", category="memory")
self.debug.start_timer("dit_to_gpu")
self.dit = self.dit.to(get_device())
self.debug.end_timer("dit_to_gpu", "DiT to GPU")
@@ -413,59 +375,43 @@ class VideoDiffusionInfer():
)
self.debug.end_timer("dit_inference", "DiT inference completed")
self.debug.log_memory_state("After inference upscale", detailed_tensors=False)
latents = na.unflatten(latents, latents_shapes)
#self.debug.log(f"UNFLATTEN time: {time.time() - t} seconds", category="timing")
# 🎯 Pré-calcul des dtypes (une seule fois)
# Pre-calculate dtypes (only once for efficiency)
vae_dtype = getattr(torch, self.config.vae.dtype)
decode_dtype = torch.float16 if (vae_dtype == torch.float16 or target_dtype == torch.float16) else vae_dtype
self.debug.log(f"VAE decode precision: {decode_dtype}", category="precision")
if preserve_vram:
if preserve_vram and not hasattr(self, "_blockswap_active"):
self.debug.log("Moving DiT back to CPU (preserve_vram mode)", category="memory")
self.debug.start_timer("dit_to_cpu")
self.dit = self.dit.to("cpu")
latents_cond = latents_cond.to("cpu")
latents_shapes = latents_shapes.to("cpu")
self.debug.end_timer("dit_to_cpu", "DiT moved to CPU (preserve_vram)")
if latents[0].shape[0] > 1:
clear_vram_cache(self.debug)
self.debug.end_timer("dit_to_cpu", "DiT moved to CPU")
clear_memory(debug=self.debug, full=True, force=True)
if latents[0].shape[0] > 1:
self.debug.start_timer("vae_to_gpu")
self.vae = self.vae.to(get_device())
self.debug.end_timer("vae_to_gpu", "VAE moved to GPU")
#with torch.autocast("cuda", decode_dtype, enabled=True):
# Move VAE to GPU if needed for decoding
manage_vae_device(self, str(get_device()), preserve_vram=False, debug=self.debug)
self.debug.log(f"VAE decode precision: {decode_dtype}", category="precision")
self.debug.log("Decoding latents to video...", category="vae")
self.debug.start_timer("vae_decode")
samples = self.vae_decode(latents, target_dtype=decode_dtype, preserve_vram=preserve_vram)
self.debug.end_timer("vae_decode", "VAE decode completed")
self.debug.log(f"Samples shape: {samples[0].shape}", category="vae")
#self.debug.log(f"🔄 ULTRA-FAST VAE DECODE time: {time.time() - t} seconds", category="timing")
#t = time.time()
#self.dit.to(get_device())
#self.vae.to("cpu")
#self.debug.log(f"🔄 Dit to GPU time: {time.time() - t} seconds", category="timing")
#t = time.time()
# 🚀 CORRECTION CRITIQUE: Conversion batch Float16 pour ComfyUI (plus rapide)
# Move VAE back to CPU after decoding if preserve_vram is enabled
if preserve_vram:
manage_vae_device(self, 'cpu', preserve_vram=preserve_vram, debug=self.debug)
self.debug.log_memory_state("After VAE decode", detailed_tensors=False)
# Converting batch Float16 for ComfyUI (faster)
if samples and len(samples) > 0 and samples[0].dtype != torch.float16:
self.debug.log(f"Converting {len(samples)} samples from {samples[0].dtype} to Float16", category="precision")
samples = [sample.to(torch.float16, non_blocking=True) for sample in samples]
#self.debug.log(f"🚀 Conversion batch Float16 time: {time.time() - t} seconds", category="timing")
# 🚀 OPTIMISATION: Nettoyage final minimal
#t = time.time()
#if dit_offload:
# self.vae.to("cpu")
# torch.cuda.empty_cache()
# self.dit.to(get_device())
#else:
# Garder VAE sur GPU pour les prochains appels
#torch.cuda.empty_cache()
#self.debug.log(f"🔄 FINAL CLEANUP time: {time.time() - t} seconds", category="timing")
return samples
+53 -41
View File
@@ -29,7 +29,7 @@ except ImportError:
print("⚠️ SafeTensors not available, recommended install: pip install safetensors")
SAFETENSORS_AVAILABLE = False
from src.optimization.memory_manager import get_basic_vram_info, clear_vram_cache
from src.optimization.memory_manager import get_basic_vram_info, clear_memory
from src.optimization.compatibility import FP8CompatibleDiT
from src.optimization.memory_manager import preinitialize_rope_cache, clear_rope_lru_caches
from src.common.config import load_config, create_object
@@ -183,12 +183,13 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=None,
runner = configure_dit_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug, block_swap_config)
debug.end_timer("dit_model_infer", "DiT model configured")
debug.log_memory_state("After DiT model configuration", detailed_tensors=False)
debug.start_timer("vae_model_infer")
checkpoint_path = os.path.join(base_cache_dir, f'./{config.vae.checkpoint}')
runner = configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug)
runner = configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug, block_swap_config)
debug.end_timer("vae_model_infer", "VAE model configured")
debug.log_memory_state("After VAE model configuration", detailed_tensors=False)
debug.start_timer("vae_memory_limit")
if hasattr(runner.vae, "set_memory_limit"):
@@ -199,7 +200,10 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=None,
blockswap_active = (
block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0
)
# Clear memory after model setup if using memory-saving features
if preserve_vram:
clear_memory(debug=debug, full=False, force=True)
# Pre-initialize RoPE cache for optimal performance if BlockSwap is NOT active
if not blockswap_active:
debug.start_timer("rope_cache_preinit")
@@ -211,7 +215,6 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=None,
# Apply BlockSwap if configured
if blockswap_active:
apply_block_swap_to_dit(runner, block_swap_config, debug)
#clear_vram_cache()
# Store debug instance on runner for consistent access
runner.debug = debug
@@ -325,14 +328,15 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
runner.dit.set_gradient_checkpointing(config.dit.gradient_checkpoint)
# Detect and log model format
debug.log(f"Loading model_weight: {model_weight}", category="model", force=True)
# Determine loading device and reason
state_loading_device = "cpu" if (preserve_vram or blockswap_active) else device
reason = f" ({('preserve_vram' if preserve_vram else 'BlockSwap active')})" if state_loading_device == "cpu" else ""
debug.log(f"Loading DiT weights {model_weight} to {state_loading_device}{reason}", category="model", force=True)
debug.start_timer("dit_load_state_dict")
state_loading_device = "cpu" if "7b" in model_weight and vram_info['total_gb'] < 25 else device
state = load_quantized_state_dict(checkpoint, state_loading_device, keep_native_fp8=True)
debug.end_timer("dit_load_state_dict", "DiT state dict loaded")
debug.start_timer("dit_load")
runner.dit.load_state_dict(state, strict=True, assign=True)
@@ -340,8 +344,7 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
del state
debug.end_timer("dit_load", "DiT load")
#state.to("cpu")
#runner.dit = runner.dit.to(device)
debug.log_memory_state("After DiT weights loaded", detailed_tensors=False)
# Apply universal compatibility wrapper to ALL models
# This ensures RoPE compatibility and optimal performance across all architectures
@@ -351,12 +354,12 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
runner.dit = FP8CompatibleDiT(runner.dit, skip_conversion=False, debug=debug)
debug.end_timer("FP8CompatibleDiT", "FP8/RoPE compatibility wrapper applied to DiT model")
# Move DiT to CPU to prevent VRAM leaks (especially for 3B model with complex RoPE)
# Move DiT to CPU to prevent VRAM leaks when preserve_vram is enabled
if preserve_vram and not blockswap_active:
debug.log("Moving DiT model to CPU (preserve_vram enabled)", category="memory")
debug.log("Moving DiT model to CPU (preserve_vram)", category="memory")
runner.dit = runner.dit.to("cpu")
if "7b" in model_weight:
clear_vram_cache(debug)
# Clear VRAM after moving models to CPU
clear_memory(debug=debug, full=True, force=True)
else:
if state_loading_device == "cpu" and not blockswap_active:
runner.dit.to(device)
@@ -370,25 +373,37 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
def configure_vae_model_inference(runner, device, checkpoint_path, config,
preserve_vram=False, model_weight=None,
vram_info=None, debug=None):
vram_info=None, debug=None, block_swap_config=None):
"""
Configure VAE model for inference without distributed decorators
Args:
runner: VideoDiffusionInfer instance
config: Model configuration
device (str): Target device
checkpoint_path (str): Path to VAE checkpoint
config: Model configuration
preserve_vram (bool): Whether to preserve VRAM by keeping model on CPU
model_weight (str): Model weight identifier
vram_info (dict): VRAM information dictionary
debug: Debug instance for logging
block_swap_config (dict): BlockSwap configuration dictionary
Features:
- Dynamic path resolution for VAE checkpoints
- SafeTensors and PyTorch format support
- FP8 and FP16 VAE handling
- Causal slicing configuration
- BlockSwap-aware device placement
"""
# Check if debug instance is available
if debug is None:
raise ValueError("Debug instance must be provided to configure_vae_model_inference")
# Check if BlockSwap is active
blockswap_active = (
block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0
)
# Create vae model
if torch.mps.is_available():
config.vae.dtype = "float16"
@@ -397,19 +412,18 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config,
dtype = getattr(torch, config.vae.dtype)
debug.start_timer("vae_model_create")
loading_device = "cpu" if preserve_vram else device
# VAE should be on CPU when preserve_vram is True or BlockSwap is active
loading_device = "cpu" if (preserve_vram or blockswap_active) else device
with torch.device(device):
with torch.device(loading_device):
runner.vae = create_object(config.vae.model)
debug.end_timer("vae_model_create", f"VAE model created on {device} with dtype {dtype}")
debug.end_timer("vae_model_create", f"VAE model created on {loading_device} with dtype {dtype}")
debug.start_timer("model_requires_grad")
runner.vae.requires_grad_(False).eval()
debug.end_timer("model_requires_grad", f"VAE model set to eval mode (gradients disabled)")
# t = time.time()
#runner.vae.to(device=loading_device, dtype=dtype)
#debug.log(f"🔄 CONFIG VAE : TO CPU TIME: {time.time() - t} seconds device: {device} dtype: {dtype}", category="timing")
# Resolve VAE checkpoint path dynamically
'''
checkpoint_path = config.vae.checkpoint
@@ -432,9 +446,12 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config,
raise FileNotFoundError(f"VAE checkpoint not found. Tried paths: {possible_paths}")
'''
# Load VAE with format detection
# Determine loading device and reason
state_loading_device = "cpu" if (preserve_vram or blockswap_active) else device
reason = f" ({('preserve_vram' if preserve_vram else 'BlockSwap active')})" if state_loading_device == "cpu" else ""
debug.log(f"Loading VAE SafeTensors to {state_loading_device}{reason}: {checkpoint_path}", category="vae", force=True)
debug.start_timer("vae_load")
state_loading_device = "cpu" if "7b" in model_weight and vram_info['total_gb'] < 25 else device
debug.log(f"Loading VAE SafeTensors: {checkpoint_path}", category="vae", force=True)
# Use optimized loading for all SafeTensors formats
if "fp8_e4m3fn" in checkpoint_path:
state = load_quantized_state_dict(checkpoint_path, state_loading_device, keep_native_fp8=True)
@@ -445,29 +462,24 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config,
debug.end_timer("vae_load", "VAE loaded")
debug.start_timer("vae_load_state_dict")
runner.vae.load_state_dict(state)
if torch.mps.is_available():
runner.vae = runner.vae.to(dtype=getattr(torch, config.vae.dtype))
if state_loading_device == "cpu":
runner.vae.to(device)
# Apply correct dtype after loading weights
vae_dtype = getattr(torch, config.vae.dtype)
runner.vae = runner.vae.to(dtype=vae_dtype)
if 'state' in locals():
del state
debug.end_timer("vae_load_state_dict", "VAE state dict loaded")
debug.log_memory_state("After VAE weights loaded", detailed_tensors=False)
# Set causal slicing if available
debug.start_timer("vae_set_causal_slicing")
if hasattr(runner.vae, "set_causal_slicing") and hasattr(config.vae, "slicing"):
debug.start_timer("vae_set_causal_slicing")
debug.log("Configuring VAE causal slicing for temporal processing", category="vae")
runner.vae.set_causal_slicing(**config.vae.slicing)
debug.end_timer("vae_set_causal_slicing", "VAE causal slicing configured")
debug.end_timer("vae_set_causal_slicing", "VAE causal slicing configured")
# Attach debug to VAE
runner.vae.debug = debug
# Propagate debug to all modules efficiently
for module in runner.vae.modules():
module.debug = debug
return runner
#runner.vae.to("cpu")
return runner
+68 -74
View File
@@ -14,14 +14,14 @@ from src.utils.constants import get_script_directory
from src.utils.debug import Debug
from src.core.model_manager import configure_runner
from src.core.generation import generation_loop
from src.optimization.memory_manager import fast_model_cleanup, fast_ram_cleanup, get_vram_usage
from src.optimization.blockswap import cleanup_blockswap
from src.optimization.memory_manager import (
clear_rope_lru_caches,
fast_model_cleanup,
fast_ram_cleanup,
complete_model_deletion,
clear_memory,
clear_all_caches,
get_device_list
get_device_list,
reset_vram_peak
)
# Import ComfyUI progress reporting
@@ -139,11 +139,15 @@ class SeedVR2:
self.debug = Debug(enabled=enable_debug)
else:
self.debug.enabled = enable_debug
self.debug.start_timer("total_execution")
self.debug.log("\n─── Model Preparation ───", category="none")
# Reset PyTorch's global peak memory stats for clean generation metrics
reset_vram_peak(self.debug)
self.debug.log_memory_state("Before model preparation", detailed_tensors=False)
self.debug.start_timer("model_preparation")
self.debug.log_memory_state("Execution start")
self.debug.log(f"Preparing model: {model}", category="model", force=True)
# Check if download succeeded
@@ -160,16 +164,15 @@ class SeedVR2:
preserve_vram, keep_vae_loaded, temporal_overlap,
cache_model, device, block_swap_config)
except Exception as e:
self.cleanup(force_ram_cleanup=True, cache_model=cache_model, debug=self.debug)
self.cleanup(cache_model=cache_model, debug=self.debug)
raise e
def cleanup(self, force_ram_cleanup: bool = True, cache_model: bool = False, debug=None):
def cleanup(self, cache_model: bool = False, debug=None):
"""
Comprehensive cleanup with memory tracking
Args:
force_ram_cleanup (bool): Whether to perform aggressive RAM cleanup
cache_model (bool): Whether to keep the model in RAM
debug: Optional Debug instance for logging
"""
@@ -186,7 +189,6 @@ class SeedVR2:
cleanup_type = "partial" if should_keep_model else "full"
debug.log(f"Starting {cleanup_type} cleanup", category="cleanup")
debug.log_memory_state(f"Before {cleanup_type} cleanup")
# Perform partial or full cleanup based on model caching
if should_keep_model:
@@ -194,13 +196,12 @@ class SeedVR2:
if hasattr(self.runner, "_blockswap_active") and self.runner._blockswap_active:
cleanup_blockswap(self.runner, keep_state_for_cache=True)
if self.runner:
offload = not getattr(self.runner, 'keep_vae_in_vram', False)
clear_all_caches(self.runner, debug, offload_vae=offload)
clear_all_caches(self.runner, debug, offload_vae=True)
debug.log("Models kept in RAM for next run", category="store")
else:
# Full cleanup - existing implementation
# Full cleanup
debug.log("Performing full cleanup", category="cleanup")
if self.runner:
@@ -218,49 +219,32 @@ class SeedVR2:
del value
self.runner.cache.cache.clear()
# Clear DiT model
# Clear DiT model completely
if hasattr(self.runner, 'dit') and self.runner.dit is not None:
# Handle FP8CompatibleDiT wrapper
if hasattr(self.runner.dit, 'dit_model'):
# Clean inner model first
clear_rope_lru_caches(self.runner.dit.dit_model)
# Ensure RoPE modules are on CPU
for name, module in self.runner.dit.dit_model.named_modules():
if hasattr(module, 'rope') and hasattr(module.rope, 'to'):
module.rope = module.rope.to('cpu')
if hasattr(module.rope, 'freqs'):
module.rope.freqs = module.rope.freqs.to('cpu')
fast_model_cleanup(self.runner.dit.dit_model)
# Aggressively clear the wrapper too
self.runner.dit.dit_model = None
# Delete the wrapper's __dict__ to break any circular refs
self.runner.dit.__dict__.clear()
# Break all references
self.runner.dit.dit_model = None
if hasattr(self.runner.dit, 'debug'):
self.runner.dit.debug = None
else:
# Direct model cleanup
clear_rope_lru_caches(self.runner.dit)
# Ensure RoPE modules are on CPU
for name, module in self.runner.dit.named_modules():
if hasattr(module, 'rope') and hasattr(module.rope, 'to'):
module.rope = module.rope.to('cpu')
if hasattr(module.rope, 'freqs'):
module.rope.freqs = module.rope.freqs.to('cpu')
fast_model_cleanup(self.runner.dit)
del self.runner.dit
self.runner.dit = None
try:
# Handle FP8CompatibleDiT wrapper
if hasattr(self.runner.dit, 'dit_model'):
clear_rope_lru_caches(self.runner.dit.dit_model)
complete_model_deletion(self.runner.dit.dit_model)
self.runner.dit.dit_model = None
else:
clear_rope_lru_caches(self.runner.dit)
# Delete the entire dit (wrapper or direct)
complete_model_deletion(self.runner.dit)
except Exception as e:
debug.log(f"Warning during DiT cleanup: {e}", category="warning")
finally:
self.runner.dit = None
# Clear VAE model
# Clear VAE model completely
if hasattr(self.runner, 'vae') and self.runner.vae is not None:
#from src.optimization.memory_manager import fast_model_cleanup
fast_model_cleanup(self.runner.vae)
# Clear VAE's internal dict
self.runner.vae.__dict__.clear()
del self.runner.vae
self.runner.vae = None
try:
complete_model_deletion(self.runner.vae)
except Exception as e:
debug.log(f"Warning during VAE cleanup: {e}", category="warning")
finally:
self.runner.vae = None
# Clear other components
for component in ['sampler', 'sampling_timesteps', 'schedule', 'config']:
@@ -285,9 +269,8 @@ class SeedVR2:
self.current_model_name = ""
# Fast RAM cleanup
if force_ram_cleanup:
fast_ram_cleanup()
# Final memory cleanup
clear_memory(debug=debug, full=True, force=True)
def _internal_execute(self, images, model, seed, new_resolution, cfg_scale, batch_size,
@@ -306,7 +289,6 @@ class SeedVR2:
if model_changed and self.runner is not None:
debug.log(f"Model changed from {current_model} to {model}, clearing cache...", category="cache")
self.cleanup(
force_ram_cleanup=True,
cache_model=False, # Don't keep old model
debug=debug,
)
@@ -324,12 +306,8 @@ class SeedVR2:
vae_tile_overlap=(vae_tile_overlap, vae_tile_overlap),
cached_runner=self.runner if cache_model else None
)
# Set whether to keep the VAE in VRAM for this run
self.runner.keep_vae_in_vram = bool(keep_vae_loaded)
self.current_model_name = model
debug.log_memory_state("Model preparation completed")
debug.end_timer("model_preparation", "Model preparation", force=True, show_breakdown=True)
@@ -355,19 +333,25 @@ class SeedVR2:
if swap_summary and swap_summary.get('total_swaps', 0) > 0:
total_time = swap_summary.get('block_total_ms', 0) + swap_summary.get('io_total_ms', 0)
debug.log(f"BlockSwap overhead: {total_time:.1f}ms across {swap_summary['total_swaps']} swaps", category="blockswap")
# Log memory usage summary
allocated, reserved, peak = get_vram_usage()
debug.log(f"Final VRAM usage - Allocated: {allocated:.2f}GB, Peak: {peak:.2f}GB", category="memory")
debug.log_memory_state("Video generation - Memory")
debug.end_timer("generation_loop", "Video generation completed", show_breakdown=True)
debug.log_memory_state("After video generation", detailed_tensors=False)
debug.log("\n─── Final Cleanup ───", category="none")
debug.start_timer("final_cleanup")
self.cleanup(force_ram_cleanup=True, cache_model=cache_model, debug=debug)
debug.log_memory_state("Final cleanup - Memory", detailed_tensors=False)
# Perform cleanup (this already calls clear_memory internally)
self.cleanup(cache_model=cache_model, debug=debug)
# Ensure sample is on CPU (ComfyUI expects CPU tensors)
if torch.is_tensor(sample) and sample.is_cuda:
sample = sample.cpu()
# Log final memory state after ALL cleanup is done
debug.end_timer("final_cleanup", "Final cleanup completed", show_breakdown=True)
# Cleanup
debug.log_memory_state("After final cleanup", detailed_tensors=True)
# Final timing summary
debug.log("\n─────────", category="none")
child_times = {
"Model preparation": debug.timer_durations.get("model_preparation", 0),
@@ -376,9 +360,10 @@ class SeedVR2:
}
debug.end_timer("total_execution", "Total execution", show_breakdown=True, custom_children=child_times)
debug.log("─────────", category="none")
# Clear history for next run
# Clear history for next run (do this last, after all logging)
debug.clear_history()
return (sample,)
def _progress_callback(self, batch_idx, total_batches, current_batch_frames, message=""):
@@ -408,12 +393,21 @@ class SeedVR2:
def __del__(self):
"""Destructor"""
try:
debug = self.debug
self.cleanup(force_ram_cleanup=True, cache_model=False, debug=debug)
# Store debug reference
debug = self.debug if hasattr(self, 'debug') else None
# Full cleanup
if hasattr(self, 'cleanup'):
self.cleanup(cache_model=False, debug=debug)
# Clear all remaining references
for attr in ['runner', 'text_pos_embeds', 'text_neg_embeds',
'current_model_name', 'debug', 'last_batch_time']:
if hasattr(self, attr):
delattr(self, attr)
except:
pass
class SeedVR2BlockSwap:
"""Configure block swapping to reduce VRAM usage"""
@@ -52,6 +52,7 @@ from .types import (
_memory_device_t,
_receptive_field_t,
)
from src.optimization.memory_manager import clear_memory
logger = get_logger(__name__) # pylint: disable=invalid-name
@@ -131,22 +132,32 @@ class Upsample3D(Upsample2D):
)
else:
hidden_states = [hidden_states]
# ADD BY NUMZ
if preserve_vram:
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
for i in range(len(hidden_states)):
hidden_states[i] = self.upscale_conv(hidden_states[i])
hidden_states[i] = rearrange(
hidden_states[i],
"b (x y z c) f h w -> b c (f z) (h x) (w y)",
x=self.spatial_ratio,
y=self.spatial_ratio,
z=self.temporal_ratio,
)
# OOM recovery attempt
try:
hidden_states[i] = self.upscale_conv(hidden_states[i])
hidden_states[i] = rearrange(
hidden_states[i],
"b (x y z c) f h w -> b c (f z) (h x) (w y)",
x=self.spatial_ratio,
y=self.spatial_ratio,
z=self.temporal_ratio,
)
except Exception as e:
debug = getattr(self, 'debug', None)
if debug:
debug.log("OOM recovery: Upsample3D upscale_conv", category="warning", force=True)
clear_memory(debug=debug, full=True, force=True)
time.sleep(1)
hidden_states[i] = self.upscale_conv(hidden_states[i])
hidden_states[i] = rearrange(
hidden_states[i],
"b (x y z c) f h w -> b c (f z) (h x) (w y)",
x=self.spatial_ratio,
y=self.spatial_ratio,
z=self.temporal_ratio,
)
# [Overridden] For causal temporal conv
if self.temporal_up and memory_state != MemoryState.ACTIVE:
@@ -154,18 +165,24 @@ class Upsample3D(Upsample2D):
if not self.slicing:
hidden_states = hidden_states[0]
# ADD BY NUMZ
if preserve_vram:
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
if self.use_conv:
if self.name == "conv":
hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
else:
hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
# OOM recovery attempt
try:
if self.name == "conv":
hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
else:
hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
except Exception as e:
debug = getattr(self, 'debug', None)
if debug:
debug.log("OOM recovery: Upsample3D conv", category="warning", force=True)
clear_memory(debug=debug, full=True, force=True)
time.sleep(1)
if self.name == "conv":
hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
else:
hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
if not self.slicing:
return hidden_states
@@ -313,19 +330,16 @@ class ResnetBlock3D(ResnetBlock2D):
hidden_states = input_tensor
hidden_states = causal_norm_wrapper(self.norm1, hidden_states, preserve_vram=preserve_vram)
# ADD BY NUMZ
# OOM recovery attempt
try:
hidden_states = self.nonlinearity(hidden_states)
except Exception as e:
if hasattr(self, 'debug') and self.debug:
self.debug.log("OOM second chance: ResnetBlock3D", category="warning", force=True)
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
time.sleep(1)
hidden_states = self.nonlinearity(hidden_states)
debug = getattr(self, 'debug', None)
if debug:
debug.log("OOM recovery: ResnetBlock3D", category="warning", force=True)
clear_memory(debug=debug, full=True, force=True)
time.sleep(1)
hidden_states = self.nonlinearity(hidden_states)
if self.upsample is not None:
# upsample_nearest_nhwc fails with large batch sizes.
@@ -27,6 +27,7 @@ from .context_parallel_lib import cache_send_recv, get_cache_size
from .global_config import get_norm_limit
from .types import MemoryState, _inflation_mode_t, _memory_device_t
from ....common.half_precision_fixes import safe_pad_operation
from src.optimization.memory_manager import clear_memory
# Single GPU inference - no distributed processing needed
#print("Warning: Using single GPU inference mode - distributed features disabled in causal_inflation_lib")
@@ -118,12 +119,11 @@ class InflatedCausalConv3d(Conv3d):
x = list(x.split(split_sizes, dim=split_dim))
if prev_cache is not None:
prev_cache = list(prev_cache.split(split_sizes, dim=split_dim))
if preserve_vram:
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
# Memory cleanup when preserve_vram is true
'''if preserve_vram:
# Use debug if passed through the module
debug = getattr(self, 'debug', None)
clear_memory(debug=debug, full=False, force=True)'''
# Loop Fwd.
cache = None
for idx in range(len(x)):
@@ -167,26 +167,15 @@ class InflatedCausalConv3d(Conv3d):
# Update cache.
cache = next_cache
# ADD BY NUMZ
if preserve_vram:
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
#print("empty cache 1")
#time.sleep(2)
# OOM recovery attempt
try:
output = torch.cat(x, split_dim)
except Exception as e:
if hasattr(self, 'debug') and self.debug:
self.debug.log("OOM Second Chance", category="warning", force=True)
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
time.sleep(2)
debug = getattr(self, 'debug', None)
if debug:
debug.log("OOM recovery: Concatenating conv splits", category="warning", force=True)
clear_memory(debug=debug, full=True, force=True)
time.sleep(1)
output = torch.cat(x, split_dim)
return output
@@ -363,38 +352,26 @@ def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor, preserve_vram: b
weights = norm_layer.weight.chunk(num_chunks, dim=0)
biases = norm_layer.bias.chunk(num_chunks, dim=0)
for i, (w, b) in enumerate(zip(weights, biases)):
# OOM recovery attempt
try:
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
except Exception as e:
if hasattr(norm_layer, 'debug') and norm_layer.debug:
norm_layer.debug.log("OOM Second Chance: Group Norm", category="warning", force=True)
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
time.sleep(2)
debug = getattr(norm_layer, 'debug', None)
if debug:
debug.log("OOM recovery: Group Norm chunk", category="warning", force=True)
clear_memory(debug=debug, full=True, force=True)
time.sleep(1)
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
x[i] = x[i].to(input_dtype)
# ADD BY NUMZ
if preserve_vram:
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
# ADD BY NUMZ
# OOM recovery attempt
try:
x = torch.cat(x, dim=1)
except Exception as e:
if hasattr(norm_layer, 'debug') and norm_layer.debug:
norm_layer.debug.log("OOM Second Chance: Cat", category="warning", force=True)
if torch.mps.is_available():
torch.mps.empty_cache()
else:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
time.sleep(2)
debug = getattr(norm_layer, 'debug', None)
if debug:
debug.log("OOM recovery: Concatenating norm chunks", category="warning", force=True)
clear_memory(debug=debug, full=True, force=True)
time.sleep(1)
x = torch.cat(x, dim=1)
else:
x = norm_layer(x)
+3 -2
View File
@@ -5,8 +5,9 @@ Contains memory management, performance optimizations, and compatibility layers
'''
# Memory management functions
from .memory_manager import (
get_basic_vram_info,
get_vram_usage,
clear_vram_cache,
clear_memory,
reset_vram_peak,
preinitialize_rope_cache,
)
@@ -27,7 +28,7 @@ from .compatibility import (
__all__ = [
# Memory management
"get_vram_usage",
"clear_vram_cache",
"clear_memory",
"reset_vram_peak",
"preinitialize_rope_cache",
+38 -52
View File
@@ -20,7 +20,7 @@ import gc
import psutil
from typing import Dict, Any, List, Tuple, Optional, Union
from src.optimization.memory_manager import get_vram_usage
from src.optimization.memory_manager import clear_memory
from src.optimization.compatibility import call_rope_with_stability
from src.common.distributed import get_device
@@ -72,7 +72,6 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) ->
raise ValueError("Debug instance must be provided to apply_block_swap_to_dit")
debug.start_timer("apply_blockswap")
debug.log_memory_state("Before BlockSwap")
# Get the actual model (handle FP8CompatibleDiT wrapper)
model = runner.dit
@@ -152,8 +151,8 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) ->
_protect_model_from_move(model, runner, debug)
debug.log("BlockSwap configuration complete", category="success")
debug.log_memory_state("After BlockSwap")
debug.end_timer("apply_blockswap", "BlockSwap configuration applied")
debug.log_memory_state("After BlockSwap", detailed_tensors=False)
def _configure_io_components(model, device: str, offload_device: str,
@@ -166,7 +165,8 @@ def _configure_io_components(model, device: str, offload_device: str,
for name, param in model.named_parameters():
if "block" not in name:
target_device = offload_device if offload_io_components else device
param.data = param.data.to(target_device, non_blocking=use_non_blocking)
# Never use non_blocking for initial setup to avoid pinned memory
param.data = param.data.to(target_device, non_blocking=False)
status = "(offloaded)" if offload_io_components else ""
debug.log(f" {name} → {target_device} {status}", category="blockswap")
@@ -192,6 +192,7 @@ def _configure_blocks(model, device: str, offload_device: str,
total_main_memory = 0.0
# Move blocks based on swap configuration
# NEVER use non_blocking for initial CPU offload to avoid pinned memory
for b, block in enumerate(model.blocks):
block_memory = get_module_memory_mb(block)
@@ -199,7 +200,7 @@ def _configure_blocks(model, device: str, offload_device: str,
block.to(device)
total_main_memory += block_memory
else:
block.to(offload_device, non_blocking=use_non_blocking)
block.to(offload_device, non_blocking=False)
total_offload_memory += block_memory
# Ensure all buffers match their containing module's device
@@ -207,13 +208,13 @@ def _configure_blocks(model, device: str, offload_device: str,
target_device = device if b > model.blocks_to_swap else offload_device
for name, buffer in block.named_buffers():
if buffer.device != torch.device(target_device):
buffer.data = buffer.data.to(target_device)
buffer.data = buffer.data.to(target_device, non_blocking=False)
# Clean up memory
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
gc.collect()
# Only force clear memory if we actually moved blocks
if model.blocks_to_swap > 0:
clear_memory(debug=debug, full=True, force=True)
else:
clear_memory(debug=debug, full=True, force=False)
return {
"offload_memory": total_offload_memory,
@@ -278,17 +279,18 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
target_device = torch.device(model.main_device)
if current_device != target_device:
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 and torch.cuda.is_available():
torch.cuda.synchronize(get_device())
# CPU->GPU: never use non_blocking (prevents pinned memory allocation)
self.to(model.main_device, non_blocking=False)
# Execute forward pass with OOM protection
output = original_forward(*args, **kwargs)
# Move back to offload device
self.to(model.offload_device, non_blocking=model.use_non_blocking)
# Only use non_blocking for GPU->GPU transfers (following WanVideo pattern)
if model.use_non_blocking and model.offload_device != "cpu":
self.to(model.offload_device, non_blocking=True)
else:
self.to(model.offload_device, non_blocking=False)
# Log timing if debug is available
if debug and t_start is not None:
@@ -299,13 +301,7 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
)
# Only clear cache under memory pressure
if torch.mps.is_available():
mem = psutil.virtual_memory()
if torch.mps.current_allocated_memory() > mem.total * 0.9:
torch.mps.empty_cache()
if torch.cuda.is_available() and torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
clear_memory(debug=debug, full=True, force=False)
else:
output = original_forward(*args, **kwargs)
@@ -353,20 +349,18 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
# Move to GPU for computation if needed
if current_device != target_device:
self.to(model.main_device)
# Synchronize if not using non-blocking transfers
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking:
if torch.mps.is_available():
torch.mps.synchronize()
else:
torch.cuda.synchronize(get_device())
# CPU->GPU: never use non_blocking
self.to(model.main_device, non_blocking=False)
# Execute forward pass
output = self._original_forward(*args, **kwargs)
# Move back to offload device
self.to(model.offload_device, non_blocking=model.use_non_blocking)
# Only use non_blocking for GPU->GPU transfers
if model.use_non_blocking and model.offload_device != "cpu":
self.to(model.offload_device, non_blocking=True)
else:
self.to(model.offload_device, non_blocking=False)
# Log timing if debug is available
if debug and t_start is not None:
@@ -377,13 +371,7 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
)
# Only clear cache under memory pressure
if torch.mps.is_available():
mem = psutil.virtual_memory()
if torch.mps.current_allocated_memory() > mem.total * 0.9:
torch.mps.empty_cache()
if torch.cuda.is_available() and torch.cuda.memory_allocated(get_device()) > torch.cuda.get_device_properties(get_device()).total_memory * 0.9:
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
clear_memory(debug=debug, full=True, force=False)
return output
@@ -666,6 +654,16 @@ def cleanup_blockswap(runner, keep_state_for_cache: bool = False) -> None:
# Move model to CPU to free VRAM
if not keep_state_for_cache:
model.to("cpu")
# Ensure all sub-modules are also on CPU
for module in model.modules():
if hasattr(module, '_parameters'):
for param in module._parameters.values():
if param is not None and param.is_cuda:
param.data = param.data.cpu()
if hasattr(module, '_buffers'):
for buffer in module._buffers.values():
if buffer is not None and buffer.is_cuda:
buffer.data = buffer.data.cpu()
debug.log("Moved model to CPU", category="store")
# Clean up runner attributes
@@ -684,15 +682,3 @@ def cleanup_blockswap(runner, keep_state_for_cache: bool = False) -> None:
# Clear local debug reference
debug = None
# Force garbage collection (multiple passes for thorough cleanup)
gc.collect(2) # Full collection including oldest generation
gc.collect()
gc.collect()
# Final memory cleanup
if torch.mps.is_available():
torch.mps.empty_cache()
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
+338 -105
View File
@@ -4,13 +4,13 @@ Handles VRAM usage, cache management, and memory optimization
Extracted from: seedvr2.py (lines 373-405, 607-626, 1016-1044)
"""
import os
\
import torch
import gc
import sys
import time
import psutil
from typing import Tuple, Optional
from typing import Tuple, Optional, Dict, Any
from src.common.cache import Cache
from src.models.dit_v2.rope import RotaryEmbeddingBase
from src.common.distributed import get_device
@@ -38,74 +38,283 @@ def get_device_list():
return devs[1:]
return devs
def get_basic_vram_info():
if torch.mps.is_available():
mem = psutil.virtual_memory()
free_memory = mem.total - mem.used
total_memory = mem.total
else:
"""🔍 Méthode basique avec PyTorch natif"""
if not torch.cuda.is_available():
return {"error": "CUDA not available"}
# Mémoire libre et totale (en bytes)
free_memory, total_memory = torch.cuda.mem_get_info(get_device())
def get_basic_vram_info() -> Dict[str, Any]:
"""
Get basic VRAM availability info (free and total memory).
Used for capacity planning and initial checks.
# Conversion en GB
free_gb = free_memory / (1024**3)
total_gb = total_memory / (1024**3)
return {
"free_gb": free_gb,
"total_gb": total_gb
}
Returns:
dict: {"free_gb": float, "total_gb": float} or {"error": str}
"""
try:
if torch.cuda.is_available():
device = get_device()
free_memory, total_memory = torch.cuda.mem_get_info(device)
elif torch.mps.is_available():
mem = psutil.virtual_memory()
free_memory = mem.total - mem.used
total_memory = mem.total
else:
return {"error": "No GPU backend available (CUDA/MPS)"}
return {
"free_gb": free_memory / (1024**3),
"total_gb": total_memory / (1024**3)
}
except Exception as e:
return {"error": f"Failed to get memory info: {str(e)}"}
# Initial VRAM check at module load
vram_info = get_basic_vram_info()
if "error" not in vram_info:
print(f"📊 Initial VRAM status: {vram_info['free_gb']:.2f}GB free / {vram_info['total_gb']:.2f}GB total")
backend = "MPS" if 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"⚠️ VRAM check: {vram_info['error']} - No available backend!")
print(f"⚠️ Memory check failed: {vram_info['error']} - No available backend!")
def get_vram_usage() -> Tuple[float, float, float]:
"""
Get current VRAM usage (allocated, reserved, peak)
Get current VRAM usage metrics for monitoring.
Used for tracking memory consumption during processing.
Returns:
tuple: (allocated_gb, reserved_gb, max_allocated_gb)
Returns (0, 0, 0) if CUDA not available
Returns (0, 0, 0) if no GPU available
"""
if torch.mps.is_available():
allocated = torch.mps.current_allocated_memory() / (1024**3)
reserved = torch.mps.driver_allocated_memory() / (1024**3)
max_allocated = 0
return allocated, reserved, max_allocated
if torch.cuda.is_available():
allocated = torch.cuda.memory_allocated(get_device()) / (1024**3)
reserved = torch.cuda.memory_reserved(get_device()) / (1024**3)
max_allocated = torch.cuda.max_memory_allocated(get_device()) / (1024**3)
return allocated, reserved, max_allocated
return 0, 0, 0
try:
if torch.cuda.is_available():
device = get_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
elif torch.mps.is_available():
allocated = torch.mps.current_allocated_memory() / (1024**3)
reserved = torch.mps.driver_allocated_memory() / (1024**3)
max_allocated = allocated # MPS doesn't track peak separately
return allocated, reserved, max_allocated
except Exception:
pass
return 0.0, 0.0, 0.0
def clear_vram_cache(debug) -> None:
"""Clear VRAM cache and run garbage collection"""
def get_ram_usage() -> Tuple[float, float, float, float]:
"""
Get current RAM usage metrics for the current process.
Provides accurate tracking of process-specific memory consumption.
Returns:
tuple: (process_gb, available_gb, total_gb, used_by_others_gb)
Returns (0, 0, 0, 0) if psutil not available
"""
try:
if not psutil:
return 0.0, 0.0, 0.0, 0.0
# Get current process memory
process = psutil.Process()
process_memory = process.memory_info()
process_gb = process_memory.rss / (1024**3)
debug.log("Clearing VRAM cache...", category="cleanup")
if torch.mps.is_available():
torch.mps.empty_cache()
# Get system memory
sys_memory = psutil.virtual_memory()
total_gb = sys_memory.total / (1024**3)
available_gb = sys_memory.available / (1024**3)
# Calculate memory used by other processes
# This is the CORRECT calculation:
total_used_gb = total_gb - available_gb # Total memory used by ALL processes
used_by_others_gb = max(0, total_used_gb - process_gb) # Subtract current process
return process_gb, available_gb, total_gb, used_by_others_gb
except Exception:
return 0.0, 0.0, 0.0, 0.0
# Global cache for OS libraries (initialized once)
_os_memory_lib = None
def clear_memory(debug=None, full=False, force=True) -> None:
"""
Clear memory caches with two-tier approach for optimal performance.
Args:
debug: Debug instance for logging (optional)
force: If True, always clear. If False, only clear when <15% free
full: If True, perform full cleanup including GC and OS operations.
If False (default), only perform minimal GPU cache clearing.
Two-tier approach:
- Minimal mode (full=False): GPU cache operations (~1-5ms)
Used for frequent calls during batch processing
- Full mode (full=True): Complete cleanup with GC and OS operations (~10-50ms)
Used at key points like model switches or final cleanup
"""
global _os_memory_lib
# Check if we should clear based on memory pressure
if not force:
should_clear = False
# Use existing function for memory info
mem_info = get_basic_vram_info()
if "error" not in mem_info:
# Check VRAM/MPS memory pressure (15% free threshold)
free_ratio = mem_info["free_gb"] / mem_info["total_gb"]
if free_ratio < 0.15:
should_clear = True
if debug:
backend = "MPS" if torch.mps.is_available() else "VRAM"
debug.log(f"{backend} pressure: {mem_info['free_gb']:.1f}GB free of {mem_info['total_gb']:.1f}GB", category="memory")
# For non-MPS systems, also check system RAM separately
if not should_clear and not torch.mps.is_available():
mem = psutil.virtual_memory()
if mem.available < mem.total * 0.15:
should_clear = True
if debug:
debug.log(f"RAM pressure: {mem.available/(1024**3):.1f}GB free of {mem.total/(1024**3):.1f}GB", category="memory")
if not should_clear:
return
# Determine cleanup level
cleanup_mode = "full" if full else "minimal"
if debug:
debug.log(f"Clearing memory caches ({cleanup_mode})...", category="cleanup")
# ===== MINIMAL OPERATIONS (Always performed) =====
# Step 1: Clear GPU caches - Fast operations (~1-5ms)
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
gc.collect()
elif torch.mps.is_available():
torch.mps.empty_cache()
# ===== FULL OPERATIONS (Only when full=True) =====
if full:
# Step 2: Clear PyTorch internal caches
if hasattr(torch, '_C'):
try:
torch._C._clear_cache()
except:
pass
# Step 3: Full garbage collection (expensive ~5-20ms)
gc.collect(2)
# Step 4: Return memory to OS (platform-specific, ~5-30ms)
try:
if sys.platform == 'linux':
# Linux: malloc_trim
import ctypes # Import only when needed
if _os_memory_lib is None:
_os_memory_lib = ctypes.CDLL("libc.so.6")
_os_memory_lib.malloc_trim(0)
elif sys.platform == 'win32':
# Windows: Trim working set
import ctypes # Import only when needed
if _os_memory_lib is None:
_os_memory_lib = ctypes.windll.kernel32
handle = _os_memory_lib.GetCurrentProcess()
_os_memory_lib.SetProcessWorkingSetSize(handle, -1, -1)
elif torch.mps.is_available():
# macOS with MPS
import ctypes # Import only when needed
import ctypes.util
if _os_memory_lib is None:
libc_path = ctypes.util.find_library('c')
if libc_path:
_os_memory_lib = ctypes.CDLL(libc_path)
if _os_memory_lib:
_os_memory_lib.sync()
except:
# OS-specific memory operations are optional
pass
def manage_vae_device(runner, target_device: str, preserve_vram: bool = False,
debug=None, reason: str = None) -> bool:
"""
Manage VAE device placement with intelligent movement and logging.
Args:
runner: Runner instance containing the VAE
target_device: Target device ('cuda:0', 'cpu', etc.)
preserve_vram: Whether preserve_vram mode is active
debug: Debug instance for logging
reason: Optional custom reason for the movement
Returns:
bool: True if VAE was moved, False if already on target device
"""
if not hasattr(runner, 'vae') or runner.vae is None:
return False
# Get current VAE device
current_device = next(runner.vae.parameters()).device if hasattr(runner.vae, 'parameters') else None
if current_device is None:
return False
# Normalize device strings for comparison
target_type = target_device.split(':')[0] if ':' in target_device else target_device
current_type = str(current_device.type)
# Skip if already on target device
if current_type == target_type:
return False
# Determine reason for movement
if reason:
reason = reason
elif preserve_vram:
reason = "preserve_vram"
else:
reason = "inference requirement"
# Start timer based on direction
timer_name = "vae_to_gpu" if target_type != 'cpu' else "vae_to_cpu"
if debug:
debug.start_timer(timer_name)
# Log the movement
if debug:
if target_type == 'cpu':
debug.log(f"Moving VAE to CPU ({reason})", category="memory")
else:
debug.log(f"Moving VAE from {current_type} to {target_device} ({reason})", category="memory")
# Move VAE
runner.vae = runner.vae.to(target_device)
# End timer
if debug:
if target_type == 'cpu':
debug.end_timer(timer_name, "VAE moved to CPU")
else:
debug.end_timer(timer_name, "VAE moved to GPU")
return True
def reset_vram_peak(debug) -> None:
"""
Reset VRAM peak counter for new tracking
Reset VRAM peak memory statistics for fresh tracking.
"""
debug.log("Resetting VRAM peak memory statistics", category="memory")
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats(get_device())
try:
if torch.cuda.is_available():
device = get_device()
torch.cuda.reset_peak_memory_stats(device)
# MPS doesn't support peak memory reset
except Exception as e:
debug.log(f"Failed to reset peak memory stats: {e}", category="warning")
def preinitialize_rope_cache(runner, debug) -> None:
"""
@@ -171,8 +380,8 @@ def preinitialize_rope_cache(runner, debug) -> None:
except Exception as e:
debug.log(f"Failed for {cache_key}: {e}", level="WARNING", category="cache")
# Return empty tensors as fallback
clear_memory(debug=debug, full=True, force=True)
time.sleep(1)
clear_vram_cache(debug)
return torch.zeros(1, 64)
@@ -200,57 +409,92 @@ def clear_rope_lru_caches(model) -> int:
"""Clear ALL LRU caches from RoPE modules"""
cleared_count = 0
for name, module in model.named_modules():
if hasattr(module, 'get_axial_freqs') and hasattr(module.get_axial_freqs, 'cache_clear'):
module.get_axial_freqs.cache_clear()
cleared_count += 1
if model is None:
return 0
try:
for name, module in model.named_modules():
if hasattr(module, 'get_axial_freqs') and hasattr(module.get_axial_freqs, 'cache_clear'):
module.get_axial_freqs.cache_clear()
cleared_count += 1
except AttributeError:
# Model structure already damaged, skip
pass
return cleared_count
def fast_model_cleanup(model):
"""Fast model cleanup without logs"""
def complete_model_deletion(model):
"""Completely delete a model and free all its memory"""
if model is None:
return
# Move to CPU
model.to("cpu")
# Clear parameters and buffers recursively
def clear_recursive(m):
for child in m.children():
clear_recursive(child)
for param in m.parameters():
if param is not None:
param.data = param.data.cpu()
param.grad = None
for buffer in m.buffers():
if buffer is not None:
buffer.data = buffer.data.cpu()
clear_recursive(model)
def fast_ram_cleanup():
"""Fast RAM cleanup without excessive logging"""
# Garbage collection
gc.collect()
# Clear MPS cache
if torch.mps.is_available():
torch.mps.empty_cache()
# Clear CUDA cache
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
torch.cuda.reset_peak_memory_stats(get_device())
# Clear PyTorch internal caches
try:
torch._C._clear_cache()
except:
# Move to CPU first
model.to("cpu")
# Clear parameters and buffers recursively and release storage
def clear_recursive(m):
# Process children first
for child in m.children():
clear_recursive(child)
# Clear parameters and release storage
if hasattr(m, '_parameters'):
for param_name, param in list(m._parameters.items()):
if param is not None:
param.data = param.data.cpu()
param.grad = None
# Release underlying storage
if param.data.numel() > 0:
param.data.set_()
# Clear buffers and release storage
if hasattr(m, '_buffers'):
for buffer_name, buffer in list(m._buffers.items()):
if buffer is not None:
buffer.data = buffer.data.cpu()
# Release underlying storage
if buffer.data.numel() > 0:
buffer.data.set_()
clear_recursive(model)
# Clear all module dicts but keep the structure
if hasattr(model, 'modules'):
for module in model.modules():
# Clear custom attributes but keep PyTorch internals
if hasattr(module, '__dict__'):
keys_to_delete = []
for key in module.__dict__.keys():
# Keep PyTorch internal attributes
if not key.startswith('_') or key.startswith('_original_'):
keys_to_delete.append(key)
for key in keys_to_delete:
try:
delattr(module, key)
except:
pass
# Now clear the model's dict
if hasattr(model, '__dict__'):
# Clear everything except PyTorch internals
keys_to_delete = []
for key in model.__dict__.keys():
if not key in ['_modules', '_parameters', '_buffers', 'training']:
keys_to_delete.append(key)
for key in keys_to_delete:
try:
delattr(model, key)
except:
pass
except AttributeError:
# Model already partially cleaned, that's OK
pass
# Final cleanup - now we can clear everything
if hasattr(model, '__dict__'):
model.__dict__.clear()
def clear_all_caches(runner, debug, offload_vae=False) -> int:
"""
@@ -374,9 +618,7 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
# Handle VAE offloading if requested
if offload_vae and hasattr(runner, 'vae') and runner.vae is not None:
debug.log("Moving VAE to CPU and clearing intermediate tensors", category="cleanup")
# Clear any intermediate tensors/buffers in VAE
# Clear intermediate tensors BEFORE moving to CPU (more efficient)
vae_caches_cleared = 0
for module in runner.vae.modules():
# Clear module-specific caches
@@ -388,7 +630,7 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
# Clear any CUDA tensors in module attributes
for attr_name in list(vars(module).keys()):
attr = getattr(module, attr_name, None)
if torch.is_tensor(attr) and attr.is_cuda:
if torch.is_tensor(attr) and (attr.is_cuda or attr.is_mps):
# Move tensor to CPU if it's not a parameter/buffer
if attr_name not in module._parameters and attr_name not in module._buffers:
setattr(module, attr_name, attr.cpu())
@@ -397,21 +639,12 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
if vae_caches_cleared > 0:
debug.log(f"Cleared {vae_caches_cleared} VAE caches", category="success")
# Move entire VAE to CPU (preserves model for reuse)
runner.vae = runner.vae.to('cpu')
debug.log("VAE moved to CPU, intermediate tensors cleared", category="success")
# Now move VAE to CPU using helper
manage_vae_device(runner, 'cpu', preserve_vram=True, debug=debug)
cleaned_items += vae_caches_cleared
# Force garbage collection
gc.collect(2) # Collect all generations
# Clear MPS cache
if torch.mps.is_available():
torch.mps.empty_cache()
# Clear CUDA cache
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
# Final memory cleanup
clear_memory(debug=debug, full=True, force=True)
return cleaned_items
+260 -193
View File
@@ -10,7 +10,8 @@ import torch
import psutil
import gc
from typing import Optional, List, Dict, Any, Tuple, Union, Set
from src.optimization.memory_manager import get_vram_usage, get_basic_vram_info
from src.optimization.memory_manager import get_vram_usage, get_basic_vram_info, get_ram_usage, reset_vram_peak
from contextlib import contextmanager
class Debug:
@@ -57,12 +58,14 @@ class Debug:
self.enabled = enabled
self.timers: Dict[str, float] = {}
self.memory_checkpoints: List[Dict[str, Any]] = []
self.max_checkpoints = 100
self.timer_hierarchy: Dict[str, List[str]] = {}
self.timer_durations: Dict[str, float] = {}
self.timer_messages: Dict[str, str] = {}
self.swap_times: List[Dict[str, Any]] = []
self.vram_history: List[float] = []
self.active_timer_stack: List[str] = []
self.timer_namespace: str = ""
def log(self, message: str, level: str = "INFO", category: str = "general", force: bool = False) -> None:
"""
@@ -93,6 +96,23 @@ class Debug:
print(f"{prefix} {message}")
@contextmanager
def timer_context(self, namespace: str):
"""
Context manager for setting a timer namespace temporarily.
All timers started within this context will be prefixed with the namespace.
Usage:
with debug.timer_context("batch_1"):
debug.start_timer("vae_encode") # Will be "batch_1_vae_encode"
"""
old_namespace = self.timer_namespace
self.timer_namespace = namespace
try:
yield
finally:
self.timer_namespace = old_namespace
def start_timer(self, name: str, force: bool = False) -> None:
"""
Start a named timer
@@ -102,6 +122,10 @@ class Debug:
force: If True, start timer even when debug is disabled
"""
if self.enabled or force:
# Apply namespace if set
if self.timer_namespace:
name = f"{self.timer_namespace}_{name}"
self.timers[name] = time.time()
# Auto-hierarchy: if there's an active timer, this is a child
@@ -132,6 +156,10 @@ class Debug:
Returns:
Duration in seconds (0.0 if timer not found)
"""
# Apply namespace if set
if self.timer_namespace:
name = f"{self.timer_namespace}_{name}"
# Check if timer exists
if name not in self.timers:
return 0.0
@@ -199,216 +227,255 @@ class Debug:
self.log(f" └─ (other operations): {unaccounted:.2f}s", category="timing", force=force)
return duration
def log_memory_state(self, label: str, show_diff: bool = True, show_tensors: bool = True,
detailed_tensors: bool = False) -> None:
"""Log current memory usage with optional diff and tensor count
detailed_tensors: bool = False) -> None:
"""
Log current memory state with minimal overhead.
Args:
label: Description label for this memory checkpoint
show_diff: Show difference from last checkpoint
show_tensors: Show tensor counts
detailed_tensors: Show detailed tensor analysis (shapes, sizes, etc.)
label: Description for this checkpoint
show_diff: Show change from last checkpoint
show_tensors: Include tensor counts
detailed_tensors: Show detailed tensor analysis (use sparingly)
"""
if not self.enabled:
return
# GPU Memory
if torch.cuda.is_available():
vram_allocated, vram_reserved, vram_max_allocated = get_vram_usage()
vram_basic_info = get_basic_vram_info()
if "error" not in vram_basic_info:
vram_free = vram_basic_info["free_gb"]
vram_total = vram_basic_info["total_gb"]
vram_used = vram_total - vram_free
# Clear, concise VRAM format
vram_info = (f"[VRAM] {vram_allocated:.2f}GB allocated / "
f"{vram_reserved:.2f}GB reserved / "
f"{vram_free:.2f}GB free / "
f"{vram_total:.2f}GB total")
self.vram_history.append(vram_allocated)
else:
vram_used = 0
vram_free = 0
vram_info = "VRAM: CPU mode"
elif torch.mps.is_available():
vram_used = 0
vram_free = 0
vram_info = "VRAM: MPS mode"
else:
vram_used = 0
vram_free = 0
vram_info = "VRAM: CPU mode"
# Collect memory metrics efficiently
memory_info = self._collect_memory_metrics()
# RAM Memory - Clear and informative
ram_info = ""
ram_process_gb = 0
if psutil:
try:
# Process-specific memory
process = psutil.Process()
mem_info = process.memory_info()
ram_process_gb = mem_info.rss / (1024**3) # Physical memory used by our process
# System-wide memory
sys_mem = psutil.virtual_memory()
ram_total_gb = sys_mem.total / (1024**3)
ram_available_gb = sys_mem.available / (1024**3)
# Calculate what's used by other processes
ram_others_gb = ram_total_gb - ram_available_gb - ram_process_gb
# Clear format matching user's request
ram_info = (f" --- [RAM] {ram_process_gb:.1f}GB SeedVR2 / "
f"{ram_others_gb:.1f}GB other processes / "
f"{ram_available_gb:.1f}GB free / "
f"{ram_total_gb:.1f}GB total ")
except Exception:
# Fallback to basic info
try:
process = psutil.Process()
ram_process_gb = process.memory_info().rss / (1024**3)
ram_info = f" | RAM: {ram_process_gb:.1f}GB used"
except:
pass
# Format and log basic memory info
log_msg = f"{label}: {memory_info['summary']}"
# Tensor count and detailed analysis
tensor_info = ""
# Add tensor info if requested
if show_tensors:
# Collect all tensors
all_tensors = []
for obj in gc.get_objects():
try:
if torch.is_tensor(obj):
all_tensors.append(obj)
except:
pass
# Separate by device
gpu_tensors = [t for t in all_tensors if t.is_cuda or t.is_mps]
cpu_tensors = [t for t in all_tensors if not (t.is_cuda or t.is_mps)]
tensor_info = f" --- [Tensors] {len(gpu_tensors)} on GPU / {len(all_tensors)} total"
# Detailed tensor analysis
if detailed_tensors and (gpu_tensors or cpu_tensors):
self.log("\n" + "─" * 60, category="memory")
self.log("DETAILED TENSOR ANALYSIS", category="memory")
self.log("─" * 60, category="memory")
# GPU Tensors Analysis
if gpu_tensors:
# Calculate total memory
gpu_memory = sum(t.element_size() * t.nelement() for t in gpu_tensors)
self.log(f"\nGPU Tensors: {len(gpu_tensors)} tensors using {gpu_memory / 1024**3:.2f} GB", category="memory")
# Group by shape for pattern recognition
shape_groups = {}
for t in gpu_tensors:
shape_key = str(list(t.shape))
if shape_key not in shape_groups:
shape_groups[shape_key] = {
'count': 0,
'dtype': str(t.dtype),
'size_mb': t.element_size() * t.nelement() / 1024**2,
'example': t
}
shape_groups[shape_key]['count'] += 1
# Sort by total memory used (count * size)
sorted_shapes = sorted(
shape_groups.items(),
key=lambda x: x[1]['count'] * x[1]['size_mb'],
reverse=True
)
self.log("\nTop GPU tensor patterns (by total memory):", category="memory")
for i, (shape, info) in enumerate(sorted_shapes[:10]):
total_mb = info['count'] * info['size_mb']
self.log(f" {i+1}. Shape {shape} × {info['count']} = {total_mb:.1f} MB total", category="memory")
self.log(f" Each: {info['size_mb']:.1f} MB, dtype: {info['dtype']}", category="memory")
# Show largest individual tensors
self.log("\nLargest individual GPU tensors:", category="memory")
sorted_gpu = sorted(gpu_tensors, key=lambda t: t.element_size() * t.nelement(), reverse=True)
for i, t in enumerate(sorted_gpu[:5]):
size_mb = t.element_size() * t.nelement() / 1024**2
self.log(f" {i+1}. Shape: {list(t.shape)}, Size: {size_mb:.1f} MB, Dtype: {t.dtype}", category="memory")
# Try to identify what it might be
shape = t.shape
if len(shape) == 4 and shape[1] in [320, 640, 1280, 1920]: # UNet features
self.log(f" → Likely UNet feature map", category="memory")
elif len(shape) == 2 and shape[0] == shape[1]: # Square matrix
self.log(f" → Likely attention matrix", category="memory")
elif len(shape) == 2 and shape[1] in [768, 1024, 2048, 4096]: # Embeddings
self.log(f" → Likely embedding/hidden states", category="memory")
# CPU Tensors Analysis (brief)
if cpu_tensors:
cpu_memory = sum(t.element_size() * t.nelement() for t in cpu_tensors)
self.log(f"\nCPU Tensors: {len(cpu_tensors)} tensors using {cpu_memory / 1024**3:.2f} GB", category="memory")
# Just show a few largest
sorted_cpu = sorted(cpu_tensors, key=lambda t: t.element_size() * t.nelement(), reverse=True)
self.log("Largest CPU tensors:", category="memory")
for i, t in enumerate(sorted_cpu[:3]):
size_mb = t.element_size() * t.nelement() / 1024**2
self.log(f" {i+1}. Shape: {list(t.shape)}, Size: {size_mb:.1f} MB", category="memory")
# Try to find model references
self.log("\n" + "─" * 60, category="memory")
# Check for nn.Module instances
modules = [obj for obj in gc.get_objects() if isinstance(obj, torch.nn.Module)]
if modules:
self.log(f"Found {len(modules)} nn.Module instances", category="memory")
# Count by type
module_types = {}
for m in modules:
mtype = type(m).__name__
module_types[mtype] = module_types.get(mtype, 0) + 1
self.log("Module types (top 5):", category="memory")
for mtype, count in sorted(module_types.items(), key=lambda x: x[1], reverse=True)[:5]:
self.log(f" {mtype}: {count}", category="memory")
tensor_stats = self._collect_tensor_stats(detailed=detailed_tensors)
log_msg += tensor_stats['summary']
# Build checkpoint
checkpoint = {
"label": label,
"vram_used_gb": vram_used,
"vram_allocated_gb": vram_allocated if torch.cuda.is_available() else 0,
"vram_reserved_gb": vram_reserved if torch.cuda.is_available() else 0,
"vram_free_gb": vram_free if torch.cuda.is_available() else 0,
"ram_process_gb": ram_process_gb,
"timestamp": time.time()
}
# Log the state
self.log(f"{label}: {vram_info}{ram_info}{tensor_info}", category="memory")
self.log(log_msg, category="memory")
# Show diff from last checkpoint
if show_diff and self.memory_checkpoints:
last = self.memory_checkpoints[-1]
vram_diff = vram_used - last["vram_used_gb"]
ram_diff = ram_process_gb - last.get("ram_process_gb", ram_process_gb)
self._log_memory_diff(memory_info)
# Log detailed analysis if requested
if detailed_tensors and tensor_stats.get('details'):
self._log_detailed_tensor_analysis(tensor_stats['details'])
# Store checkpoint with memory limit
self._store_checkpoint(label, memory_info)
# Reset PyTorch's peak memory stats for next interval
reset_vram_peak(self)
def _collect_memory_metrics(self) -> Dict[str, Any]:
"""Collect current memory metrics efficiently."""
metrics = {
'vram_allocated': 0.0,
'vram_reserved': 0.0,
'vram_free': 0.0,
'vram_total': 0.0,
'vram_peak_since_last': 0.0,
'ram_process': 0.0,
'ram_available': 0.0,
'ram_total': 0.0,
'ram_others': 0.0,
'summary': ""
}
# VRAM metrics
if torch.cuda.is_available() or torch.mps.is_available():
metrics['vram_allocated'], metrics['vram_reserved'], current_global_peak = get_vram_usage()
diffs = []
if abs(vram_diff) > 0.1: # Significant VRAM change
sign = "+" if vram_diff > 0 else ""
diffs.append(f"VRAM {sign}{vram_diff:.2f}GB")
if abs(ram_diff) > 0.1: # Significant RAM change
sign = "+" if ram_diff > 0 else ""
diffs.append(f"RAM {sign}{ram_diff:.2f}GB")
# Calculate peak since last log_memory_state
# This captures the actual peak that occurred between calls
metrics['vram_peak_since_last'] = current_global_peak
if diffs:
self.log(f" Memory changes: {', '.join(diffs)}", category="memory")
vram_info = get_basic_vram_info()
if "error" not in vram_info:
metrics['vram_free'] = vram_info["free_gb"]
metrics['vram_total'] = vram_info["total_gb"]
backend = "MPS" if torch.mps.is_available() else "VRAM"
vram_str = (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"{metrics['vram_free']:.2f}GB free / "
f"{metrics['vram_total']:.2f}GB total")
else:
vram_str = "[CPU mode]"
else:
vram_str = "[CPU mode]"
# RAM metrics using new function
metrics['ram_process'], metrics['ram_available'], metrics['ram_total'], metrics['ram_others'] = get_ram_usage()
if metrics['ram_total'] > 0:
ram_str = (f" --- [RAM] {metrics['ram_process']:.1f}GB process / "
f"{metrics['ram_others']:.1f}GB others / "
f"{metrics['ram_available']:.1f}GB free / "
f"{metrics['ram_total']:.1f}GB total")
else:
ram_str = ""
metrics['summary'] = vram_str + ram_str
# Update VRAM history for tracking
if torch.cuda.is_available() or torch.mps.is_available():
self.vram_history.append(metrics['vram_allocated'])
return metrics
def _collect_tensor_stats(self, detailed: bool = False) -> Dict[str, Any]:
"""Collect tensor statistics with minimal overhead."""
stats = {
'gpu_count': 0,
'cpu_count': 0,
'total_count': 0,
'summary': "",
'details': None
}
if detailed:
stats['details'] = {
'gpu_tensors': [],
'large_cpu_tensors': [],
'shape_patterns': {},
'module_types': {}
}
# Single pass through gc objects
for obj in gc.get_objects():
try:
if torch.is_tensor(obj):
stats['total_count'] += 1
is_gpu = obj.is_cuda or (hasattr(obj, 'is_mps') and obj.is_mps)
if is_gpu:
stats['gpu_count'] += 1
else:
stats['cpu_count'] += 1
# Collect detailed info if requested
if detailed and obj.numel() > 0:
size_mb = obj.element_size() * obj.nelement() / (1024**2)
if is_gpu or size_mb > 10: # Only track GPU tensors or large CPU tensors
tensor_info = {
'shape': tuple(obj.shape),
'dtype': str(obj.dtype),
'size_mb': size_mb,
'requires_grad': obj.requires_grad
}
if is_gpu:
stats['details']['gpu_tensors'].append(tensor_info)
elif size_mb > 10: # Large CPU tensors (>10MB)
stats['details']['large_cpu_tensors'].append(tensor_info)
# Track shape patterns
shape_key = str(tuple(obj.shape))
stats['details']['shape_patterns'][shape_key] = stats['details']['shape_patterns'].get(shape_key, 0) + 1
elif detailed and isinstance(obj, torch.nn.Module):
module_type = type(obj).__name__
stats['details']['module_types'][module_type] = stats['details']['module_types'].get(module_type, 0) + 1
except (ReferenceError, AttributeError):
# Object was deleted or doesn't have expected attributes
pass
stats['summary'] = f" --- [Tensors] {stats['gpu_count']} GPU / {stats['cpu_count']} CPU / {stats['total_count']} total"
return stats
def _log_detailed_tensor_analysis(self, details: Dict[str, Any]) -> None:
"""Log detailed tensor analysis when requested."""
self.log("─" * 60, category="memory")
self.log("DETAILED MEMORY ANALYSIS", category="memory")
self.log("─" * 60, category="memory")
# GPU tensors
if details['gpu_tensors']:
gpu_total_gb = sum(t['size_mb'] for t in details['gpu_tensors']) / 1024
self.log(f"GPU TENSORS: {len(details['gpu_tensors'])} using {gpu_total_gb:.2f}GB", category="memory")
# Show top 5 largest
largest = sorted(details['gpu_tensors'], key=lambda x: x['size_mb'], reverse=True)[:5]
for t in largest:
self.log(f" {t['shape']}: {t['size_mb']:.1f}MB, {t['dtype']}", category="memory")
# Large CPU tensors
if details['large_cpu_tensors']:
cpu_large_gb = sum(t['size_mb'] for t in details['large_cpu_tensors']) / 1024
self.log(f"LARGE CPU TENSORS (>10MB): {len(details['large_cpu_tensors'])} using {cpu_large_gb:.2f}GB", category="memory")
# Show top 3 largest
largest = sorted(details['large_cpu_tensors'], key=lambda x: x['size_mb'], reverse=True)[:3]
for t in largest:
self.log(f" {t['shape']}: {t['size_mb']:.1f}MB, {t['dtype']}", category="memory")
# Common shape patterns
if details['shape_patterns']:
common_shapes = sorted(details['shape_patterns'].items(),
key=lambda x: x[1], reverse=True)[:5]
if len(common_shapes) > 0:
self.log("COMMON TENSOR SHAPES:", category="memory")
for shape, count in common_shapes:
if count > 1:
self.log(f" {shape}: {count} instances", category="memory")
# Module instances
if details['module_types']:
multi_instance = [(k, v) for k, v in details['module_types'].items() if v > 1]
if multi_instance:
self.log("MULTIPLE MODULE INSTANCES:", category="memory")
for mtype, count in sorted(multi_instance, key=lambda x: x[1], reverse=True)[:5]:
self.log(f" {mtype}: {count} instances", category="memory")
self.log("─" * 60, category="memory")
def _log_memory_diff(self, current_metrics: Dict[str, Any]) -> None:
"""Log memory changes from last checkpoint."""
last = self.memory_checkpoints[-1]
vram_diff = current_metrics['vram_allocated'] - last.get('vram_allocated', 0)
ram_diff = current_metrics['ram_process'] - last.get('ram_process', 0)
diffs = []
if abs(vram_diff) > 0.01:
sign = "+" if vram_diff > 0 else ""
diffs.append(f"VRAM {sign}{vram_diff:.2f}GB")
if abs(ram_diff) > 0.01:
sign = "+" if ram_diff > 0 else ""
diffs.append(f"RAM {sign}{ram_diff:.2f}GB")
if diffs:
self.log(f" Memory changes: {', '.join(diffs)}", category="memory")
def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None:
"""Store checkpoint with memory limit to prevent leaks."""
checkpoint = {
'label': label,
'timestamp': time.time(),
'vram_allocated': metrics['vram_allocated'],
'vram_reserved': metrics['vram_reserved'],
'vram_free': metrics['vram_free'],
'ram_process': metrics['ram_process'],
'ram_available': metrics['ram_available'],
'ram_others': metrics['ram_others']
}
self.memory_checkpoints.append(checkpoint)
# Prevent memory leak by limiting checkpoint history
if len(self.memory_checkpoints) > self.max_checkpoints:
# Keep first and last N/2 checkpoints for better history coverage
mid = self.max_checkpoints // 2
self.memory_checkpoints = (self.memory_checkpoints[:mid] +
self.memory_checkpoints[-mid:])
def log_swap_time(self, component_id: Union[int, str], duration: float,
component_type: str = "block") -> None: