From 9b7c113681218bda285ce7d5060cbc282733c8e5 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Fri, 22 Aug 2025 09:20:41 -0400 Subject: [PATCH] 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. --- src/__init__.py | 4 +- src/common/diffusion/samplers/euler.py | 11 +- src/core/generation.py | 276 ++++++----- src/core/infer.py | 110 ++--- src/core/model_manager.py | 94 ++-- src/interfaces/comfyui_node.py | 142 +++--- .../video_vae_v3/modules/attn_video_vae.py | 86 ++-- .../modules/causal_inflation_lib.py | 71 +-- src/optimization/__init__.py | 5 +- src/optimization/blockswap.py | 90 ++-- src/optimization/memory_manager.py | 443 +++++++++++++---- src/utils/debug.py | 453 ++++++++++-------- 12 files changed, 1032 insertions(+), 753 deletions(-) diff --git a/src/__init__.py b/src/__init__.py index 13475a0..10512c8 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -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 diff --git a/src/common/diffusion/samplers/euler.py b/src/common/diffusion/samplers/euler.py index 0b0c9e9..1ce8768 100644 --- a/src/common/diffusion/samplers/euler.py +++ b/src/common/diffusion/samplers/euler.py @@ -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() diff --git a/src/core/generation.py b/src/core/generation.py index e3e3ddf..6b36c9f 100644 --- a/src/core/generation.py +++ b/src/core/generation.py @@ -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 - } + } \ No newline at end of file diff --git a/src/core/infer.py b/src/core/infer.py index bf4e6d1..d4b0a71 100644 --- a/src/core/infer.py +++ b/src/core/infer.py @@ -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 diff --git a/src/core/model_manager.py b/src/core/model_manager.py index 187b516..e8295bd 100644 --- a/src/core/model_manager.py +++ b/src/core/model_manager.py @@ -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") \ No newline at end of file + return runner \ No newline at end of file diff --git a/src/interfaces/comfyui_node.py b/src/interfaces/comfyui_node.py index 44bdb4d..7629375 100644 --- a/src/interfaces/comfyui_node.py +++ b/src/interfaces/comfyui_node.py @@ -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""" diff --git a/src/models/video_vae_v3/modules/attn_video_vae.py b/src/models/video_vae_v3/modules/attn_video_vae.py index ed3003c..f13ac47 100644 --- a/src/models/video_vae_v3/modules/attn_video_vae.py +++ b/src/models/video_vae_v3/modules/attn_video_vae.py @@ -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. diff --git a/src/models/video_vae_v3/modules/causal_inflation_lib.py b/src/models/video_vae_v3/modules/causal_inflation_lib.py index 3da3912..9dae215 100644 --- a/src/models/video_vae_v3/modules/causal_inflation_lib.py +++ b/src/models/video_vae_v3/modules/causal_inflation_lib.py @@ -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) diff --git a/src/optimization/__init__.py b/src/optimization/__init__.py index b28eefa..5681871 100644 --- a/src/optimization/__init__.py +++ b/src/optimization/__init__.py @@ -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", diff --git a/src/optimization/blockswap.py b/src/optimization/blockswap.py index a5cc0fb..97b89d8 100644 --- a/src/optimization/blockswap.py +++ b/src/optimization/blockswap.py @@ -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() diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index e0fe372..c974f47 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -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 \ No newline at end of file diff --git a/src/utils/debug.py b/src/utils/debug.py index 81192bd..dee79bd 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -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: