refactor(WIP): Memory Management Overhaul
This is a work-in-progress commit that consolidates a series of changes to fix memory leaks, optimize VRAM/RAM usage, improve performance, and enhance code maintainability. * **BlockSwap Pinned Memory:** Disabled `use_non_blocking=True` for CPU-to-GPU transfers to resolve a memory leak where pinned memory was not being released. * **Logging-Induced Leaks:** Modified `log_memory_state()` to avoid holding references to tensors during analysis and added a history limit to the checkpoint system to prevent unbounded memory growth. * **Incomplete Model Cleanup:** Ensured models are completely deleted and their tensor storage is released when `cache_model=False`. * **Lingering Tensors:** Fixed an issue where a scalar tensor from sampling timesteps and text embeddings remained on the GPU between batches when `preserve_vram` is active. * **Centralized Cleanup Functions:** Introduced `clear_memory()` to replace `clear_vram_cache()` and all manual `torch.cuda.empty_cache()` calls, providing consistent VRAM/RAM cleanup logic. The function features a `full` parameter to distinguish between a fast, GPU-only cache clear (~1-5ms) for frequent operations and a full cleanup with garbage collection (~10-50ms) for critical stages. * **Direct-to-CPU Model Loading:** Modified DiT/VAE weight loading to load directly onto the CPU when `preserve_vram` or `BlockSwap` is active, avoiding unnecessary VRAM spikes during model preparation. * **VAE Device Management:** Created the `manage_vae_device()` helper function to centralize the logic for moving the VAE between the CPU and GPU, reducing code duplication. This also fixed a bug that incorrectly kept the VAE on the GPU when `preserve_vram` was active. * **CPU Offloading:** Implemented logic to move text embeddings and sampling timesteps to the CPU after each batch when `preserve_vram` is active, reducing idle VRAM usage. * **VAE Decode Performance:** Replaced proactive, frequent memory clearing during VAE decode with a reactive Out-of-Memory (OOM) handling system. This fixed a significant performance regression and eliminated the need for the `keep_vae_loaded_during_decode` flag. * **Reduced Overhead:** Removed redundant `gc.collect()` calls from multiple locations to decrease unnecessary processing overhead. * **Interval-Based VRAM Tracking:** Modified the logging system to reset peak VRAM statistics after each `log_memory_state()` call, enabling accurate tracking of peak memory usage for specific processing intervals (e.g., encode, inference, decode). * **Accurate RAM Monitoring:** Added the `get_ram_usage()` function for correct process-specific RAM tracking. * **Efficient Log Refactoring:** Refactored `log_memory_state()` into modular helper methods, optimizing tensor analysis into a single-pass `gc` iteration to improve both performance and maintainability. * **Log Clarity:** Refined memory state and debug logging to remove redundant snapshots and add new ones for critical operations like model loading, weight loading, VAE encoding, and decoding. Standardized log message conventions. * **Per-Batch Timers:** Implemented timer namespacing to ensure that performance timers for each batch are logged correctly without overwriting one another. * **Error Handling:** Added `try/except` blocks to key memory and device management functions to handle edge cases and improve robustness. * **Code Cleanup:** Removed deprecated code and outdated comments throughout the related modules. * **Documentation:** Updated comments and function docstrings to reflect the new memory management architecture.
This commit is contained in:
+2
-2
@@ -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
|
||||
|
||||
@@ -21,6 +21,7 @@ from typing import Callable
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch.nn import functional as F
|
||||
from src.optimization.memory_manager import clear_memory
|
||||
|
||||
#from ....models.dit_v2 import na
|
||||
|
||||
@@ -71,13 +72,9 @@ class EulerSampler(Sampler):
|
||||
|
||||
# Nettoyer les tenseurs temporaires
|
||||
del pred
|
||||
if torch.mps.is_available():
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
# Use debug if available from the sampler
|
||||
debug = getattr(self, 'debug', None)
|
||||
clear_memory(debug=debug, full=False, force=True)
|
||||
|
||||
i += 1
|
||||
progress.update()
|
||||
|
||||
+164
-112
@@ -26,7 +26,7 @@ from src.common.distributed import get_device
|
||||
|
||||
|
||||
# Import required modules
|
||||
from src.optimization.memory_manager import reset_vram_peak, clear_all_caches
|
||||
from src.optimization.memory_manager import manage_vae_device, clear_all_caches, clear_memory
|
||||
from src.optimization.performance import (
|
||||
optimized_video_rearrange, optimized_single_video_rearrange,
|
||||
optimized_sample_to_image_format, temporal_latent_blending
|
||||
@@ -114,6 +114,10 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
|
||||
cond_noise_scale = 0.0
|
||||
|
||||
def _add_noise(x, aug_noise):
|
||||
# Early return if no noise is being added
|
||||
if cond_noise_scale == 0.0:
|
||||
return x
|
||||
|
||||
# Use adaptive optimal dtype
|
||||
t = (
|
||||
torch.tensor([1000.0], device=device, dtype=dtype)
|
||||
@@ -122,6 +126,10 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
|
||||
shape = torch.tensor(x.shape[1:], device=device)[None]
|
||||
t = runner.timestep_transform(t, shape)
|
||||
x = runner.schedule.forward(x, aug_noise, t)
|
||||
|
||||
# Explicit cleanup of intermediate tensors
|
||||
del t, shape
|
||||
|
||||
return x
|
||||
|
||||
# Generate conditions with memory optimization
|
||||
@@ -137,17 +145,33 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
|
||||
|
||||
# Use adaptive autocast for optimal performance
|
||||
with torch.no_grad():
|
||||
# Restore timesteps to GPU if they were offloaded
|
||||
if preserve_vram and hasattr(runner, 'sampling_timesteps') and hasattr(runner.sampling_timesteps, 'timesteps'):
|
||||
if not runner.sampling_timesteps.timesteps.is_cuda:
|
||||
debug.log("Restoring timesteps tensor to GPU (preserve_vram)", category="memory")
|
||||
debug.start_timer("timesteps_to_gpu")
|
||||
runner.sampling_timesteps.timesteps = runner.sampling_timesteps.timesteps.to(device, non_blocking=True)
|
||||
debug.end_timer("timesteps_to_gpu", "Sampling timesteps restored to GPU")
|
||||
|
||||
with torch.autocast(str(get_device()), autocast_dtype, enabled=True):
|
||||
video_tensors = runner.inference(
|
||||
noises=noises,
|
||||
conditions=conditions,
|
||||
preserve_vram=preserve_vram # Memory offload optimization
|
||||
and not use_blockswap, # Disable dit_offload if BlockSwap active
|
||||
preserve_vram=preserve_vram, # Memory offload optimization
|
||||
temporal_overlap=temporal_overlap,
|
||||
use_blockswap=use_blockswap,
|
||||
**text_embeds_dict,
|
||||
)
|
||||
|
||||
# Clean up diffusion timesteps from GPU if preserve_vram is enabled
|
||||
if preserve_vram:
|
||||
if hasattr(runner, 'sampling_timesteps') and hasattr(runner.sampling_timesteps, 'timesteps'):
|
||||
if runner.sampling_timesteps.timesteps.is_cuda:
|
||||
debug.log("Moving timesteps tensor to CPU (preserve_vram)", category="memory")
|
||||
debug.start_timer("timesteps_to_cpu")
|
||||
runner.sampling_timesteps.timesteps = runner.sampling_timesteps.timesteps.cpu()
|
||||
debug.end_timer("timesteps_to_cpu", "Sampling timesteps offloaded to CPU")
|
||||
|
||||
# Process samples with advanced optimization
|
||||
samples = optimized_video_rearrange(video_tensors)
|
||||
#last_latents = samples[-temporal_overlap:] if temporal_overlap > 0 else samples[-1:]
|
||||
@@ -283,8 +307,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
|
||||
# Set random seed
|
||||
set_seed(seed)
|
||||
debug.log_memory_state("Model configuration - Memory")
|
||||
debug.end_timer("model_config", "Model configuration completed", show_breakdown=True)
|
||||
debug.log_memory_state("After model configuration", detailed_tensors=False)
|
||||
|
||||
# ───────────────────────────────────────────────────────────────
|
||||
# Step 2: Input Preparation & Transformation Setup
|
||||
@@ -315,8 +339,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
text_neg_embeds = torch.load(os.path.join(script_directory, 'neg_emb.pt')).to(device, dtype=compute_dtype)
|
||||
text_embeds = {"texts_pos": [text_pos_embeds], "texts_neg": [text_neg_embeds]}
|
||||
|
||||
debug.log_memory_state("Input preparation - Memory ")
|
||||
debug.end_timer("input_prep", "Input preparation completed", show_breakdown=True)
|
||||
debug.log_memory_state("After input preparation", detailed_tensors=False)
|
||||
|
||||
# ───────────────────────────────────────────────────────────────
|
||||
# Step 3: Batch Processing
|
||||
@@ -325,8 +349,6 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
debug.start_timer("batch_processing")
|
||||
|
||||
# Standard processing (non-TileVAE) continues below
|
||||
# Memory optimization
|
||||
reset_vram_peak(debug)
|
||||
|
||||
# Calculate processing parameters
|
||||
step = batch_size - temporal_overlap
|
||||
@@ -364,103 +386,124 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
break # Not enough new frames, stop
|
||||
|
||||
batch_number = (batch_idx // step + 1) if step > 0 else 1
|
||||
debug.start_timer(f"batch_{batch_number}")
|
||||
current_frames = end_idx - start_idx
|
||||
debug.log("", category="none")
|
||||
debug.log(f"━━━ Batch {batch_number}/{total_batches}: frames {start_idx}-{end_idx-1} ━━━", category="generation", force=True)
|
||||
|
||||
debug.log_memory_state(f"Before batch {batch_number} processing", detailed_tensors=False)
|
||||
|
||||
# Use timer context for this batch - all timers within will be namespaced
|
||||
with debug.timer_context(f"batch_{batch_number}"):
|
||||
debug.start_timer("batch") # This becomes "batch_1_batch" internally
|
||||
|
||||
# Process current batch
|
||||
video = images[start_idx:end_idx]
|
||||
debug.log(f"Video compute dtype: {compute_dtype}", category="generation")
|
||||
# Use adaptive computation dtype
|
||||
video = video.permute(0, 3, 1, 2).to(device, dtype=compute_dtype)
|
||||
|
||||
# Apply video transformations with memory optimization
|
||||
transformed_video = video_transform(video)
|
||||
del video
|
||||
#video = video.to("cpu")
|
||||
#del video
|
||||
ori_lengths = [transformed_video.size(1)]
|
||||
|
||||
# Handle correct format: frames % 4 == 1
|
||||
t = transformed_video.size(1)
|
||||
debug.log(f"Sequence of {t} frames", category="video", force=True)
|
||||
|
||||
|
||||
if len(images) >= 5 and t % 4 != 1:
|
||||
debug.log(f"Transformed video shape before cut: {transformed_video.shape}", category="video")
|
||||
transformed_video = cut_videos(transformed_video)
|
||||
debug.log(f"Transformed video shape: {transformed_video.shape}", category="video")
|
||||
|
||||
# Context-aware temporal strategy
|
||||
# First batch: standard complete diffusion
|
||||
debug.start_timer("vae_to_gpu")
|
||||
runner.vae = runner.vae.to(device)
|
||||
debug.end_timer("vae_to_gpu", "VAE to GPU")
|
||||
debug.start_timer("vae_encode")
|
||||
debug.log(f"VAE encoding precision: {autocast_dtype}", category="vae")
|
||||
with torch.autocast(str(device), autocast_dtype, enabled=True):
|
||||
cond_latents = runner.vae_encode([transformed_video])
|
||||
debug.end_timer("vae_encode", "VAE encoding")
|
||||
#tps = time.time()
|
||||
#transformed_video = transformed_video.to("cpu")
|
||||
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
|
||||
debug.log(f"Cond latents shape: {cond_latents[0].shape}", category="info")
|
||||
|
||||
# Normal generation
|
||||
samples = generation_step(runner, text_embeds, preserve_vram, cond_latents=cond_latents, temporal_overlap=temporal_overlap, debug=debug)
|
||||
#del cond_latents
|
||||
del cond_latents
|
||||
|
||||
|
||||
# Post-process samples
|
||||
sample = samples[0]
|
||||
del samples
|
||||
#del samples
|
||||
if ori_lengths[0] < sample.shape[0]:
|
||||
sample = sample[:ori_lengths[0]]
|
||||
#if temporal_overlap > 0 and not is_first_batch and sample.shape[0] > effective_batch_size - temporal_overlap:
|
||||
# sample = sample[temporal_overlap:] # Remove overlap frames from output
|
||||
|
||||
# Apply color correction if available
|
||||
debug.start_timer("video_to_device")
|
||||
transformed_video = transformed_video.to(device)
|
||||
debug.end_timer("video_to_device", "Transformed video to device")
|
||||
|
||||
input_video = [optimized_single_video_rearrange(transformed_video)]
|
||||
del transformed_video
|
||||
#transformed_video = transformed_video.to("cpu")
|
||||
#del transformed_video
|
||||
sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)], debug)
|
||||
del input_video
|
||||
# Move text embeddings back to GPU if they were offloaded
|
||||
if preserve_vram or (block_swap_config and block_swap_config.get("offload_io_components", False)):
|
||||
if text_pos_embeds.device.type == "cpu":
|
||||
reason = "preserve_vram" if preserve_vram else "BlockSwap I/O offload"
|
||||
debug.log(f"Restoring text embeddings to GPU ({reason} active)", category="memory")
|
||||
debug.start_timer("text_embeddings_to_gpu")
|
||||
text_pos_embeds = text_pos_embeds.to(device, dtype=compute_dtype)
|
||||
text_neg_embeds = text_neg_embeds.to(device, dtype=compute_dtype)
|
||||
text_embeds["texts_pos"][0] = text_pos_embeds
|
||||
text_embeds["texts_neg"][0] = text_neg_embeds
|
||||
debug.end_timer("text_embeddings_to_gpu", "Text embeddings restored to GPU")
|
||||
|
||||
# Process current batch
|
||||
video = images[start_idx:end_idx]
|
||||
debug.log(f"Video compute dtype: {compute_dtype}", category="precision")
|
||||
# Use adaptive computation dtype
|
||||
video = video.permute(0, 3, 1, 2).to(device, dtype=compute_dtype)
|
||||
|
||||
# Apply video transformations with memory optimization
|
||||
transformed_video = video_transform(video)
|
||||
del video
|
||||
#video = video.to("cpu")
|
||||
#del video
|
||||
ori_lengths = [transformed_video.size(1)]
|
||||
|
||||
# Handle correct format: frames % 4 == 1
|
||||
t = transformed_video.size(1)
|
||||
debug.log(f"Sequence of {t} frames", category="video", force=True)
|
||||
|
||||
|
||||
if len(images) >= 5 and t % 4 != 1:
|
||||
debug.log(f"Transformed video shape before cut: {transformed_video.shape}", category="video")
|
||||
transformed_video = cut_videos(transformed_video)
|
||||
debug.log(f"Transformed video shape: {transformed_video.shape}", category="video")
|
||||
|
||||
# Context-aware temporal strategy
|
||||
# First batch: standard complete diffusion
|
||||
|
||||
# Convert to final image format
|
||||
sample = optimized_sample_to_image_format(sample)
|
||||
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
|
||||
sample_cpu = sample.to(torch.float16).to("cpu")
|
||||
del sample
|
||||
batch_samples.append(sample_cpu)
|
||||
#del sample
|
||||
|
||||
# Aggressive cleanup after each batch
|
||||
# tps = time.time()
|
||||
# Progress callback - batch start
|
||||
if progress_callback:
|
||||
progress_callback(batch_count+1, total_batches, current_frames, "Processing batch...")
|
||||
#transformed_video = transformed_video.to("cpu")
|
||||
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
|
||||
# Clean VRAM after each batch when preserve_vram is active
|
||||
if preserve_vram:
|
||||
# Only offload the VAE when we are not keeping it resident
|
||||
offload = not getattr(runner, 'keep_vae_in_vram', False)
|
||||
clear_all_caches(runner, debug, offload_vae=offload)
|
||||
#del transformed_video
|
||||
#clear_vram_cache()
|
||||
# Log memory state at the end of each batch
|
||||
debug.log_memory_state(f"Batch {batch_number} - Memory")
|
||||
debug.end_timer(f"batch_{batch_number}", f"Batch {batch_number} processed", show_breakdown=True)
|
||||
# Move VAE to GPU if needed for encoding
|
||||
manage_vae_device(runner, str(device), preserve_vram=False, debug=debug)
|
||||
debug.log(f"VAE encoding precision: {autocast_dtype}", category="precision")
|
||||
debug.log("Encoding video to latents...", category="vae")
|
||||
debug.start_timer("vae_encoding")
|
||||
with torch.autocast(str(device), autocast_dtype, enabled=True):
|
||||
cond_latents = runner.vae_encode([transformed_video])
|
||||
debug.end_timer("vae_encoding", "VAE encoding")
|
||||
debug.log(f"Cond latents shape: {cond_latents[0].shape}", category="info")
|
||||
|
||||
# Move VAE back to CPU after encoding if preserve_vram is enabled
|
||||
if preserve_vram:
|
||||
manage_vae_device(runner, 'cpu', preserve_vram=preserve_vram, debug=debug)
|
||||
|
||||
debug.log_memory_state("After VAE encode", detailed_tensors=False)
|
||||
|
||||
# Normal generation
|
||||
samples = generation_step(runner, text_embeds, preserve_vram, cond_latents=cond_latents, temporal_overlap=temporal_overlap, debug=debug)
|
||||
#del cond_latents
|
||||
del cond_latents
|
||||
|
||||
# Post-process samples
|
||||
sample = samples[0]
|
||||
del samples
|
||||
#del samples
|
||||
if ori_lengths[0] < sample.shape[0]:
|
||||
sample = sample[:ori_lengths[0]]
|
||||
#if temporal_overlap > 0 and not is_first_batch and sample.shape[0] > effective_batch_size - temporal_overlap:
|
||||
# sample = sample[temporal_overlap:] # Remove overlap frames from output
|
||||
|
||||
# Apply color correction if available
|
||||
debug.start_timer("video_to_device")
|
||||
transformed_video = transformed_video.to(device)
|
||||
debug.end_timer("video_to_device", "Transformed video to device")
|
||||
|
||||
input_video = [optimized_single_video_rearrange(transformed_video)]
|
||||
del transformed_video
|
||||
#transformed_video = transformed_video.to("cpu")
|
||||
#del transformed_video
|
||||
sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)], debug)
|
||||
del input_video
|
||||
|
||||
# Convert to final image format
|
||||
sample = optimized_sample_to_image_format(sample)
|
||||
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
|
||||
sample_cpu = sample.to(torch.float16).to("cpu")
|
||||
del sample
|
||||
batch_samples.append(sample_cpu)
|
||||
|
||||
# Aggressive cleanup after each batch
|
||||
# tps = time.time()
|
||||
# Progress callback - batch start
|
||||
if progress_callback:
|
||||
progress_callback(batch_count+1, total_batches, current_frames, "Processing batch...")
|
||||
#transformed_video = transformed_video.to("cpu")
|
||||
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
|
||||
# Clean VRAM after each batch when preserve_vram is active
|
||||
if preserve_vram or (block_swap_config and block_swap_config.get("offload_io_components", False)):
|
||||
# Move text embeddings to CPU when using memory-saving features
|
||||
reason = "preserve_vram" if preserve_vram else "BlockSwap I/O offload"
|
||||
debug.log(f"Moving text embeddings to CPU ({reason})", category="memory")
|
||||
debug.start_timer("text_embeddings_to_cpu")
|
||||
text_pos_embeds = text_pos_embeds.to("cpu")
|
||||
text_neg_embeds = text_neg_embeds.to("cpu")
|
||||
text_embeds["texts_pos"][0] = text_pos_embeds
|
||||
text_embeds["texts_neg"][0] = text_neg_embeds
|
||||
debug.end_timer("text_embeddings_to_cpu", "Text embeddings moved to CPU")
|
||||
|
||||
# Log memory state at the end of each batch
|
||||
debug.end_timer("batch", f"Batch {batch_number} processed", show_breakdown=True)
|
||||
debug.log_memory_state(f"After batch {batch_number} processing", detailed_tensors=True)
|
||||
|
||||
finally:
|
||||
debug.log("", category="none")
|
||||
@@ -468,22 +511,24 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
debug.start_timer("generation_cleanup")
|
||||
|
||||
# Final cleanup of embeddings
|
||||
debug.log("Moving text embeddings to CPU (final cleanup)", category="memory")
|
||||
text_pos_embeds = text_pos_embeds.to("cpu")
|
||||
text_neg_embeds = text_neg_embeds.to("cpu")
|
||||
|
||||
# Move DiT to CPU
|
||||
debug.log("Moving DiT to CPU (final cleanup)", category="memory")
|
||||
debug.start_timer("dit_to_cpu_cleanup")
|
||||
runner.dit.to("cpu")
|
||||
if not getattr(runner, 'keep_vae_in_vram', False):
|
||||
runner.vae.to("cpu")
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
#del text_pos_embeds, text_neg_embeds
|
||||
#clear_vram_cache()
|
||||
debug.end_timer("dit_to_cpu_cleanup", "DiT moved to CPU (final cleanup)")
|
||||
|
||||
# Move VAE to CPU
|
||||
manage_vae_device(runner, 'cpu', preserve_vram=False, debug=debug, reason="final cleanup")
|
||||
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
|
||||
# Log final memory state
|
||||
debug.log_memory_state("Generation cleanup - Memory")
|
||||
debug.end_timer("generation_cleanup", "Batch generation cleanup")
|
||||
debug.log_memory_state("After batch generation cleanup", detailed_tensors=False)
|
||||
|
||||
debug.end_timer("batch_processing", "Batch processing completed", show_breakdown=True)
|
||||
|
||||
@@ -535,6 +580,15 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
|
||||
# Clean up merged batch memory
|
||||
del batch_group, merged_result
|
||||
|
||||
# Clean up batch_samples list completely
|
||||
for batch in batch_samples:
|
||||
if torch.is_tensor(batch):
|
||||
if batch.is_cuda:
|
||||
batch.cpu()
|
||||
del batch
|
||||
batch_samples.clear()
|
||||
del batch_samples
|
||||
|
||||
debug.log(f"Memory pre-allocation completed for output tensor: {final_video_images.shape}", category="success")
|
||||
debug.log("Pre-allocation ensures contiguous memory for final video output", category="info")
|
||||
@@ -565,11 +619,9 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
debug.log(f" Most swapped: Block {swap_summary['most_swapped_block']} "
|
||||
f"({swap_summary['most_swapped_count']} times)", category="blockswap")
|
||||
|
||||
debug.log_memory_state("Post-processing - Memory")
|
||||
debug.end_timer("post_processing", "Post-processing completed", show_breakdown=True)
|
||||
debug.log_memory_state("After post-processing", detailed_tensors=False)
|
||||
|
||||
# Cleanup batch_samples
|
||||
#del batch_samples
|
||||
return final_video_images
|
||||
|
||||
|
||||
@@ -661,4 +713,4 @@ def calculate_optimal_batch_params(total_frames, batch_size, temporal_overlap):
|
||||
'best_batch': best_batch,
|
||||
'padding_waste': padding_waste,
|
||||
'is_optimal': batch_size in optimal_batches
|
||||
}
|
||||
}
|
||||
+28
-82
@@ -18,7 +18,7 @@ import torch
|
||||
from einops import rearrange
|
||||
from omegaconf import DictConfig, ListConfig
|
||||
from torch import Tensor
|
||||
from src.optimization.memory_manager import clear_vram_cache
|
||||
from src.optimization.memory_manager import clear_memory, manage_vae_device
|
||||
|
||||
from src.common.diffusion import (
|
||||
classifier_free_guidance_dispatcher,
|
||||
@@ -68,7 +68,7 @@ def optimized_channels_to_second(tensor):
|
||||
return tensor.permute(*dims)
|
||||
|
||||
class VideoDiffusionInfer():
|
||||
def __init__(self, config: DictConfig, debug=None, vae_tiling_enabled: bool = False,
|
||||
def __init__(self, config: DictConfig, debug=None, vae_tiling_enabled: bool = False,
|
||||
vae_tile_size: Tuple[int, int] = (512, 512), vae_tile_overlap: Tuple[int, int] = (64, 64)):
|
||||
# Check if debug instance is available
|
||||
if debug is None:
|
||||
@@ -78,9 +78,6 @@ class VideoDiffusionInfer():
|
||||
self.vae_tiling_enabled = vae_tiling_enabled
|
||||
self.vae_tile_size = vae_tile_size
|
||||
self.vae_tile_overlap = vae_tile_overlap
|
||||
|
||||
# Keep the VAE on the GPU between decode calls
|
||||
self.keep_vae_in_vram: bool = False
|
||||
|
||||
def get_condition(self, latent: Tensor, latent_blur: Tensor, task: str) -> Tensor:
|
||||
t, h, w, c = latent.shape
|
||||
@@ -187,7 +184,6 @@ class VideoDiffusionInfer():
|
||||
"""🚀 VAE decode optimisé - décodage direct sans chunking, compatible avec autocast externe"""
|
||||
samples = []
|
||||
if len(latents) > 0:
|
||||
#t = time.time()
|
||||
device = get_device()
|
||||
dtype = getattr(torch, self.config.vae.dtype)
|
||||
scale = self.config.vae.scaling_factor
|
||||
@@ -209,9 +205,7 @@ class VideoDiffusionInfer():
|
||||
self.debug.log(f"Using VAE Tiled Decoding (Tile: {self.vae_tile_size}, Overlap: {self.vae_tile_overlap})", category="vae", force=True)
|
||||
|
||||
self.debug.log(f"Latents batch shape: {latents[0].shape}", category="info")
|
||||
self.debug.start_timer("vae_decode")
|
||||
# If the user wants to keep the VAE resident, do not let the VAE free its buffers
|
||||
internal_preserve_vram = preserve_vram and not getattr(self, 'keep_vae_in_vram', False)
|
||||
|
||||
for i, latent in enumerate(latents):
|
||||
effective_dtype = target_dtype if target_dtype is not None else dtype
|
||||
latent = latent.to(device, effective_dtype, non_blocking=True)
|
||||
@@ -220,7 +214,7 @@ class VideoDiffusionInfer():
|
||||
latent = latent.squeeze(2)
|
||||
|
||||
sample = self.vae.decode(
|
||||
latent, preserve_vram=internal_preserve_vram,
|
||||
latent, preserve_vram=preserve_vram,
|
||||
tiled=use_tiling, tile_size=self.vae_tile_size,
|
||||
tile_overlap=self.vae_tile_overlap).sample
|
||||
|
||||
@@ -229,8 +223,6 @@ class VideoDiffusionInfer():
|
||||
|
||||
samples.append(sample)
|
||||
|
||||
self.debug.end_timer("vae_decode", "VAE decode completed")
|
||||
|
||||
if self.config.vae.grouping:
|
||||
samples = na.unpack(samples, indices)
|
||||
else:
|
||||
@@ -271,30 +263,6 @@ class VideoDiffusionInfer():
|
||||
timesteps = timesteps * self.schedule.T
|
||||
return timesteps
|
||||
|
||||
def get_vram_usage(self):
|
||||
"""Obtenir l'utilisation VRAM actuelle (allouée et réservée)"""
|
||||
if torch.mps.is_available():
|
||||
allocated = torch.mps.current_allocated_memory() / (1024**3)
|
||||
reserved = torch.mps.driver_allocated_memory() / (1024**3)
|
||||
max_allocated = 0
|
||||
return allocated, reserved, max_allocated
|
||||
if torch.cuda.is_available():
|
||||
allocated = torch.cuda.memory_allocated() / (1024**3)
|
||||
reserved = torch.cuda.memory_reserved() / (1024**3)
|
||||
max_allocated = torch.cuda.max_memory_allocated() / (1024**3)
|
||||
return allocated, reserved, max_allocated
|
||||
return 0, 0, 0
|
||||
|
||||
def get_vram_peak(self):
|
||||
"""Obtenir le pic VRAM depuis le dernier reset"""
|
||||
if torch.cuda.is_available():
|
||||
return torch.cuda.max_memory_allocated() / (1024**3)
|
||||
return 0
|
||||
|
||||
def reset_vram_peak(self):
|
||||
"""Reset le compteur de pic VRAM"""
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.reset_peak_memory_stats()
|
||||
|
||||
@torch.no_grad()
|
||||
def inference(
|
||||
@@ -314,9 +282,6 @@ class VideoDiffusionInfer():
|
||||
# Return if empty.
|
||||
if batch_size == 0:
|
||||
return []
|
||||
|
||||
# Monitoring VRAM initial et reset des pics
|
||||
#self.reset_vram_peak()
|
||||
|
||||
# Set cfg scale
|
||||
if cfg_scale is None:
|
||||
@@ -369,12 +334,9 @@ class VideoDiffusionInfer():
|
||||
|
||||
|
||||
if preserve_vram:
|
||||
if conditions[0].shape[0] > 1:
|
||||
self.debug.start_timer("vae_to_cpu")
|
||||
self.vae = self.vae.to("cpu")
|
||||
self.debug.end_timer("vae_to_cpu", "VAE to CPU")
|
||||
# Before sampling, check if BlockSwap is active
|
||||
if not use_blockswap and not hasattr(self, "_blockswap_active"):
|
||||
self.debug.log("Moving DiT to GPU (inference requirement)", category="memory")
|
||||
self.debug.start_timer("dit_to_gpu")
|
||||
self.dit = self.dit.to(get_device())
|
||||
self.debug.end_timer("dit_to_gpu", "DiT to GPU")
|
||||
@@ -413,59 +375,43 @@ class VideoDiffusionInfer():
|
||||
)
|
||||
|
||||
self.debug.end_timer("dit_inference", "DiT inference completed")
|
||||
self.debug.log_memory_state("After inference upscale", detailed_tensors=False)
|
||||
|
||||
latents = na.unflatten(latents, latents_shapes)
|
||||
#self.debug.log(f"UNFLATTEN time: {time.time() - t} seconds", category="timing")
|
||||
|
||||
# 🎯 Pré-calcul des dtypes (une seule fois)
|
||||
# Pre-calculate dtypes (only once for efficiency)
|
||||
vae_dtype = getattr(torch, self.config.vae.dtype)
|
||||
decode_dtype = torch.float16 if (vae_dtype == torch.float16 or target_dtype == torch.float16) else vae_dtype
|
||||
self.debug.log(f"VAE decode precision: {decode_dtype}", category="precision")
|
||||
if preserve_vram:
|
||||
|
||||
if preserve_vram and not hasattr(self, "_blockswap_active"):
|
||||
self.debug.log("Moving DiT back to CPU (preserve_vram mode)", category="memory")
|
||||
self.debug.start_timer("dit_to_cpu")
|
||||
self.dit = self.dit.to("cpu")
|
||||
latents_cond = latents_cond.to("cpu")
|
||||
latents_shapes = latents_shapes.to("cpu")
|
||||
self.debug.end_timer("dit_to_cpu", "DiT moved to CPU (preserve_vram)")
|
||||
if latents[0].shape[0] > 1:
|
||||
clear_vram_cache(self.debug)
|
||||
self.debug.end_timer("dit_to_cpu", "DiT moved to CPU")
|
||||
clear_memory(debug=self.debug, full=True, force=True)
|
||||
|
||||
if latents[0].shape[0] > 1:
|
||||
self.debug.start_timer("vae_to_gpu")
|
||||
self.vae = self.vae.to(get_device())
|
||||
|
||||
self.debug.end_timer("vae_to_gpu", "VAE moved to GPU")
|
||||
|
||||
|
||||
|
||||
|
||||
#with torch.autocast("cuda", decode_dtype, enabled=True):
|
||||
# Move VAE to GPU if needed for decoding
|
||||
manage_vae_device(self, str(get_device()), preserve_vram=False, debug=self.debug)
|
||||
self.debug.log(f"VAE decode precision: {decode_dtype}", category="precision")
|
||||
self.debug.log("Decoding latents to video...", category="vae")
|
||||
self.debug.start_timer("vae_decode")
|
||||
samples = self.vae_decode(latents, target_dtype=decode_dtype, preserve_vram=preserve_vram)
|
||||
|
||||
self.debug.end_timer("vae_decode", "VAE decode completed")
|
||||
self.debug.log(f"Samples shape: {samples[0].shape}", category="vae")
|
||||
#self.debug.log(f"🔄 ULTRA-FAST VAE DECODE time: {time.time() - t} seconds", category="timing")
|
||||
#t = time.time()
|
||||
#self.dit.to(get_device())
|
||||
#self.vae.to("cpu")
|
||||
#self.debug.log(f"🔄 Dit to GPU time: {time.time() - t} seconds", category="timing")
|
||||
#t = time.time()
|
||||
# 🚀 CORRECTION CRITIQUE: Conversion batch Float16 pour ComfyUI (plus rapide)
|
||||
|
||||
# Move VAE back to CPU after decoding if preserve_vram is enabled
|
||||
if preserve_vram:
|
||||
manage_vae_device(self, 'cpu', preserve_vram=preserve_vram, debug=self.debug)
|
||||
|
||||
self.debug.log_memory_state("After VAE decode", detailed_tensors=False)
|
||||
|
||||
|
||||
# Converting batch Float16 for ComfyUI (faster)
|
||||
if samples and len(samples) > 0 and samples[0].dtype != torch.float16:
|
||||
self.debug.log(f"Converting {len(samples)} samples from {samples[0].dtype} to Float16", category="precision")
|
||||
samples = [sample.to(torch.float16, non_blocking=True) for sample in samples]
|
||||
|
||||
#self.debug.log(f"🚀 Conversion batch Float16 time: {time.time() - t} seconds", category="timing")
|
||||
|
||||
# 🚀 OPTIMISATION: Nettoyage final minimal
|
||||
#t = time.time()
|
||||
#if dit_offload:
|
||||
# self.vae.to("cpu")
|
||||
# torch.cuda.empty_cache()
|
||||
# self.dit.to(get_device())
|
||||
#else:
|
||||
# Garder VAE sur GPU pour les prochains appels
|
||||
#torch.cuda.empty_cache()
|
||||
#self.debug.log(f"🔄 FINAL CLEANUP time: {time.time() - t} seconds", category="timing")
|
||||
|
||||
|
||||
|
||||
return samples
|
||||
|
||||
+53
-41
@@ -29,7 +29,7 @@ except ImportError:
|
||||
print("⚠️ SafeTensors not available, recommended install: pip install safetensors")
|
||||
SAFETENSORS_AVAILABLE = False
|
||||
|
||||
from src.optimization.memory_manager import get_basic_vram_info, clear_vram_cache
|
||||
from src.optimization.memory_manager import get_basic_vram_info, clear_memory
|
||||
from src.optimization.compatibility import FP8CompatibleDiT
|
||||
from src.optimization.memory_manager import preinitialize_rope_cache, clear_rope_lru_caches
|
||||
from src.common.config import load_config, create_object
|
||||
@@ -183,12 +183,13 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=None,
|
||||
runner = configure_dit_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug, block_swap_config)
|
||||
|
||||
debug.end_timer("dit_model_infer", "DiT model configured")
|
||||
|
||||
debug.log_memory_state("After DiT model configuration", detailed_tensors=False)
|
||||
|
||||
debug.start_timer("vae_model_infer")
|
||||
checkpoint_path = os.path.join(base_cache_dir, f'./{config.vae.checkpoint}')
|
||||
runner = configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug)
|
||||
|
||||
runner = configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug, block_swap_config)
|
||||
debug.end_timer("vae_model_infer", "VAE model configured")
|
||||
debug.log_memory_state("After VAE model configuration", detailed_tensors=False)
|
||||
|
||||
debug.start_timer("vae_memory_limit")
|
||||
if hasattr(runner.vae, "set_memory_limit"):
|
||||
@@ -199,7 +200,10 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=None,
|
||||
blockswap_active = (
|
||||
block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0
|
||||
)
|
||||
|
||||
# Clear memory after model setup if using memory-saving features
|
||||
if preserve_vram:
|
||||
clear_memory(debug=debug, full=False, force=True)
|
||||
|
||||
# Pre-initialize RoPE cache for optimal performance if BlockSwap is NOT active
|
||||
if not blockswap_active:
|
||||
debug.start_timer("rope_cache_preinit")
|
||||
@@ -211,7 +215,6 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=None,
|
||||
# Apply BlockSwap if configured
|
||||
if blockswap_active:
|
||||
apply_block_swap_to_dit(runner, block_swap_config, debug)
|
||||
#clear_vram_cache()
|
||||
|
||||
# Store debug instance on runner for consistent access
|
||||
runner.debug = debug
|
||||
@@ -325,14 +328,15 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
|
||||
|
||||
runner.dit.set_gradient_checkpointing(config.dit.gradient_checkpoint)
|
||||
|
||||
# Detect and log model format
|
||||
debug.log(f"Loading model_weight: {model_weight}", category="model", force=True)
|
||||
|
||||
# Determine loading device and reason
|
||||
state_loading_device = "cpu" if (preserve_vram or blockswap_active) else device
|
||||
reason = f" ({('preserve_vram' if preserve_vram else 'BlockSwap active')})" if state_loading_device == "cpu" else ""
|
||||
debug.log(f"Loading DiT weights {model_weight} to {state_loading_device}{reason}", category="model", force=True)
|
||||
|
||||
debug.start_timer("dit_load_state_dict")
|
||||
state_loading_device = "cpu" if "7b" in model_weight and vram_info['total_gb'] < 25 else device
|
||||
state = load_quantized_state_dict(checkpoint, state_loading_device, keep_native_fp8=True)
|
||||
|
||||
debug.end_timer("dit_load_state_dict", "DiT state dict loaded")
|
||||
|
||||
debug.start_timer("dit_load")
|
||||
runner.dit.load_state_dict(state, strict=True, assign=True)
|
||||
|
||||
@@ -340,8 +344,7 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
|
||||
del state
|
||||
|
||||
debug.end_timer("dit_load", "DiT load")
|
||||
#state.to("cpu")
|
||||
#runner.dit = runner.dit.to(device)
|
||||
debug.log_memory_state("After DiT weights loaded", detailed_tensors=False)
|
||||
|
||||
# Apply universal compatibility wrapper to ALL models
|
||||
# This ensures RoPE compatibility and optimal performance across all architectures
|
||||
@@ -351,12 +354,12 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
|
||||
runner.dit = FP8CompatibleDiT(runner.dit, skip_conversion=False, debug=debug)
|
||||
debug.end_timer("FP8CompatibleDiT", "FP8/RoPE compatibility wrapper applied to DiT model")
|
||||
|
||||
# Move DiT to CPU to prevent VRAM leaks (especially for 3B model with complex RoPE)
|
||||
# Move DiT to CPU to prevent VRAM leaks when preserve_vram is enabled
|
||||
if preserve_vram and not blockswap_active:
|
||||
debug.log("Moving DiT model to CPU (preserve_vram enabled)", category="memory")
|
||||
debug.log("Moving DiT model to CPU (preserve_vram)", category="memory")
|
||||
runner.dit = runner.dit.to("cpu")
|
||||
if "7b" in model_weight:
|
||||
clear_vram_cache(debug)
|
||||
# Clear VRAM after moving models to CPU
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
else:
|
||||
if state_loading_device == "cpu" and not blockswap_active:
|
||||
runner.dit.to(device)
|
||||
@@ -370,25 +373,37 @@ def configure_dit_model_inference(runner, device, checkpoint, config,
|
||||
|
||||
def configure_vae_model_inference(runner, device, checkpoint_path, config,
|
||||
preserve_vram=False, model_weight=None,
|
||||
vram_info=None, debug=None):
|
||||
vram_info=None, debug=None, block_swap_config=None):
|
||||
"""
|
||||
Configure VAE model for inference without distributed decorators
|
||||
|
||||
Args:
|
||||
runner: VideoDiffusionInfer instance
|
||||
config: Model configuration
|
||||
device (str): Target device
|
||||
checkpoint_path (str): Path to VAE checkpoint
|
||||
config: Model configuration
|
||||
preserve_vram (bool): Whether to preserve VRAM by keeping model on CPU
|
||||
model_weight (str): Model weight identifier
|
||||
vram_info (dict): VRAM information dictionary
|
||||
debug: Debug instance for logging
|
||||
block_swap_config (dict): BlockSwap configuration dictionary
|
||||
|
||||
Features:
|
||||
- Dynamic path resolution for VAE checkpoints
|
||||
- SafeTensors and PyTorch format support
|
||||
- FP8 and FP16 VAE handling
|
||||
- Causal slicing configuration
|
||||
- BlockSwap-aware device placement
|
||||
"""
|
||||
# Check if debug instance is available
|
||||
if debug is None:
|
||||
raise ValueError("Debug instance must be provided to configure_vae_model_inference")
|
||||
|
||||
# Check if BlockSwap is active
|
||||
blockswap_active = (
|
||||
block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0
|
||||
)
|
||||
|
||||
# Create vae model
|
||||
if torch.mps.is_available():
|
||||
config.vae.dtype = "float16"
|
||||
@@ -397,19 +412,18 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config,
|
||||
|
||||
dtype = getattr(torch, config.vae.dtype)
|
||||
debug.start_timer("vae_model_create")
|
||||
loading_device = "cpu" if preserve_vram else device
|
||||
|
||||
# VAE should be on CPU when preserve_vram is True or BlockSwap is active
|
||||
loading_device = "cpu" if (preserve_vram or blockswap_active) else device
|
||||
|
||||
with torch.device(device):
|
||||
with torch.device(loading_device):
|
||||
runner.vae = create_object(config.vae.model)
|
||||
debug.end_timer("vae_model_create", f"VAE model created on {device} with dtype {dtype}")
|
||||
debug.end_timer("vae_model_create", f"VAE model created on {loading_device} with dtype {dtype}")
|
||||
|
||||
debug.start_timer("model_requires_grad")
|
||||
runner.vae.requires_grad_(False).eval()
|
||||
debug.end_timer("model_requires_grad", f"VAE model set to eval mode (gradients disabled)")
|
||||
|
||||
# t = time.time()
|
||||
#runner.vae.to(device=loading_device, dtype=dtype)
|
||||
#debug.log(f"🔄 CONFIG VAE : TO CPU TIME: {time.time() - t} seconds device: {device} dtype: {dtype}", category="timing")
|
||||
# Resolve VAE checkpoint path dynamically
|
||||
'''
|
||||
checkpoint_path = config.vae.checkpoint
|
||||
@@ -432,9 +446,12 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config,
|
||||
raise FileNotFoundError(f"VAE checkpoint not found. Tried paths: {possible_paths}")
|
||||
'''
|
||||
# Load VAE with format detection
|
||||
# Determine loading device and reason
|
||||
state_loading_device = "cpu" if (preserve_vram or blockswap_active) else device
|
||||
reason = f" ({('preserve_vram' if preserve_vram else 'BlockSwap active')})" if state_loading_device == "cpu" else ""
|
||||
debug.log(f"Loading VAE SafeTensors to {state_loading_device}{reason}: {checkpoint_path}", category="vae", force=True)
|
||||
|
||||
debug.start_timer("vae_load")
|
||||
state_loading_device = "cpu" if "7b" in model_weight and vram_info['total_gb'] < 25 else device
|
||||
debug.log(f"Loading VAE SafeTensors: {checkpoint_path}", category="vae", force=True)
|
||||
# Use optimized loading for all SafeTensors formats
|
||||
if "fp8_e4m3fn" in checkpoint_path:
|
||||
state = load_quantized_state_dict(checkpoint_path, state_loading_device, keep_native_fp8=True)
|
||||
@@ -445,29 +462,24 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config,
|
||||
debug.end_timer("vae_load", "VAE loaded")
|
||||
debug.start_timer("vae_load_state_dict")
|
||||
runner.vae.load_state_dict(state)
|
||||
|
||||
if torch.mps.is_available():
|
||||
runner.vae = runner.vae.to(dtype=getattr(torch, config.vae.dtype))
|
||||
|
||||
if state_loading_device == "cpu":
|
||||
runner.vae.to(device)
|
||||
# Apply correct dtype after loading weights
|
||||
vae_dtype = getattr(torch, config.vae.dtype)
|
||||
runner.vae = runner.vae.to(dtype=vae_dtype)
|
||||
|
||||
if 'state' in locals():
|
||||
del state
|
||||
debug.end_timer("vae_load_state_dict", "VAE state dict loaded")
|
||||
debug.log_memory_state("After VAE weights loaded", detailed_tensors=False)
|
||||
|
||||
# Set causal slicing if available
|
||||
debug.start_timer("vae_set_causal_slicing")
|
||||
if hasattr(runner.vae, "set_causal_slicing") and hasattr(config.vae, "slicing"):
|
||||
debug.start_timer("vae_set_causal_slicing")
|
||||
debug.log("Configuring VAE causal slicing for temporal processing", category="vae")
|
||||
runner.vae.set_causal_slicing(**config.vae.slicing)
|
||||
|
||||
debug.end_timer("vae_set_causal_slicing", "VAE causal slicing configured")
|
||||
debug.end_timer("vae_set_causal_slicing", "VAE causal slicing configured")
|
||||
|
||||
# Attach debug to VAE
|
||||
runner.vae.debug = debug
|
||||
|
||||
# Propagate debug to all modules efficiently
|
||||
for module in runner.vae.modules():
|
||||
module.debug = debug
|
||||
|
||||
return runner
|
||||
#runner.vae.to("cpu")
|
||||
return runner
|
||||
@@ -14,14 +14,14 @@ from src.utils.constants import get_script_directory
|
||||
from src.utils.debug import Debug
|
||||
from src.core.model_manager import configure_runner
|
||||
from src.core.generation import generation_loop
|
||||
from src.optimization.memory_manager import fast_model_cleanup, fast_ram_cleanup, get_vram_usage
|
||||
from src.optimization.blockswap import cleanup_blockswap
|
||||
from src.optimization.memory_manager import (
|
||||
clear_rope_lru_caches,
|
||||
fast_model_cleanup,
|
||||
fast_ram_cleanup,
|
||||
complete_model_deletion,
|
||||
clear_memory,
|
||||
clear_all_caches,
|
||||
get_device_list
|
||||
get_device_list,
|
||||
reset_vram_peak
|
||||
)
|
||||
|
||||
# Import ComfyUI progress reporting
|
||||
@@ -139,11 +139,15 @@ class SeedVR2:
|
||||
self.debug = Debug(enabled=enable_debug)
|
||||
else:
|
||||
self.debug.enabled = enable_debug
|
||||
|
||||
|
||||
self.debug.start_timer("total_execution")
|
||||
self.debug.log("\n─── Model Preparation ───", category="none")
|
||||
|
||||
# Reset PyTorch's global peak memory stats for clean generation metrics
|
||||
reset_vram_peak(self.debug)
|
||||
|
||||
self.debug.log_memory_state("Before model preparation", detailed_tensors=False)
|
||||
self.debug.start_timer("model_preparation")
|
||||
self.debug.log_memory_state("Execution start")
|
||||
self.debug.log(f"Preparing model: {model}", category="model", force=True)
|
||||
|
||||
# Check if download succeeded
|
||||
@@ -160,16 +164,15 @@ class SeedVR2:
|
||||
preserve_vram, keep_vae_loaded, temporal_overlap,
|
||||
cache_model, device, block_swap_config)
|
||||
except Exception as e:
|
||||
self.cleanup(force_ram_cleanup=True, cache_model=cache_model, debug=self.debug)
|
||||
self.cleanup(cache_model=cache_model, debug=self.debug)
|
||||
raise e
|
||||
|
||||
|
||||
def cleanup(self, force_ram_cleanup: bool = True, cache_model: bool = False, debug=None):
|
||||
def cleanup(self, cache_model: bool = False, debug=None):
|
||||
"""
|
||||
Comprehensive cleanup with memory tracking
|
||||
|
||||
Args:
|
||||
force_ram_cleanup (bool): Whether to perform aggressive RAM cleanup
|
||||
cache_model (bool): Whether to keep the model in RAM
|
||||
debug: Optional Debug instance for logging
|
||||
"""
|
||||
@@ -186,7 +189,6 @@ class SeedVR2:
|
||||
|
||||
cleanup_type = "partial" if should_keep_model else "full"
|
||||
debug.log(f"Starting {cleanup_type} cleanup", category="cleanup")
|
||||
debug.log_memory_state(f"Before {cleanup_type} cleanup")
|
||||
|
||||
# Perform partial or full cleanup based on model caching
|
||||
if should_keep_model:
|
||||
@@ -194,13 +196,12 @@ class SeedVR2:
|
||||
if hasattr(self.runner, "_blockswap_active") and self.runner._blockswap_active:
|
||||
cleanup_blockswap(self.runner, keep_state_for_cache=True)
|
||||
if self.runner:
|
||||
offload = not getattr(self.runner, 'keep_vae_in_vram', False)
|
||||
clear_all_caches(self.runner, debug, offload_vae=offload)
|
||||
clear_all_caches(self.runner, debug, offload_vae=True)
|
||||
|
||||
debug.log("Models kept in RAM for next run", category="store")
|
||||
|
||||
else:
|
||||
# Full cleanup - existing implementation
|
||||
# Full cleanup
|
||||
debug.log("Performing full cleanup", category="cleanup")
|
||||
|
||||
if self.runner:
|
||||
@@ -218,49 +219,32 @@ class SeedVR2:
|
||||
del value
|
||||
self.runner.cache.cache.clear()
|
||||
|
||||
# Clear DiT model
|
||||
# Clear DiT model completely
|
||||
if hasattr(self.runner, 'dit') and self.runner.dit is not None:
|
||||
# Handle FP8CompatibleDiT wrapper
|
||||
if hasattr(self.runner.dit, 'dit_model'):
|
||||
# Clean inner model first
|
||||
clear_rope_lru_caches(self.runner.dit.dit_model)
|
||||
# Ensure RoPE modules are on CPU
|
||||
for name, module in self.runner.dit.dit_model.named_modules():
|
||||
if hasattr(module, 'rope') and hasattr(module.rope, 'to'):
|
||||
module.rope = module.rope.to('cpu')
|
||||
if hasattr(module.rope, 'freqs'):
|
||||
module.rope.freqs = module.rope.freqs.to('cpu')
|
||||
fast_model_cleanup(self.runner.dit.dit_model)
|
||||
# Aggressively clear the wrapper too
|
||||
self.runner.dit.dit_model = None
|
||||
# Delete the wrapper's __dict__ to break any circular refs
|
||||
self.runner.dit.__dict__.clear()
|
||||
# Break all references
|
||||
self.runner.dit.dit_model = None
|
||||
if hasattr(self.runner.dit, 'debug'):
|
||||
self.runner.dit.debug = None
|
||||
else:
|
||||
# Direct model cleanup
|
||||
clear_rope_lru_caches(self.runner.dit)
|
||||
# Ensure RoPE modules are on CPU
|
||||
for name, module in self.runner.dit.named_modules():
|
||||
if hasattr(module, 'rope') and hasattr(module.rope, 'to'):
|
||||
module.rope = module.rope.to('cpu')
|
||||
if hasattr(module.rope, 'freqs'):
|
||||
module.rope.freqs = module.rope.freqs.to('cpu')
|
||||
fast_model_cleanup(self.runner.dit)
|
||||
|
||||
del self.runner.dit
|
||||
self.runner.dit = None
|
||||
try:
|
||||
# Handle FP8CompatibleDiT wrapper
|
||||
if hasattr(self.runner.dit, 'dit_model'):
|
||||
clear_rope_lru_caches(self.runner.dit.dit_model)
|
||||
complete_model_deletion(self.runner.dit.dit_model)
|
||||
self.runner.dit.dit_model = None
|
||||
else:
|
||||
clear_rope_lru_caches(self.runner.dit)
|
||||
|
||||
# Delete the entire dit (wrapper or direct)
|
||||
complete_model_deletion(self.runner.dit)
|
||||
except Exception as e:
|
||||
debug.log(f"Warning during DiT cleanup: {e}", category="warning")
|
||||
finally:
|
||||
self.runner.dit = None
|
||||
|
||||
# Clear VAE model
|
||||
# Clear VAE model completely
|
||||
if hasattr(self.runner, 'vae') and self.runner.vae is not None:
|
||||
#from src.optimization.memory_manager import fast_model_cleanup
|
||||
fast_model_cleanup(self.runner.vae)
|
||||
# Clear VAE's internal dict
|
||||
self.runner.vae.__dict__.clear()
|
||||
del self.runner.vae
|
||||
self.runner.vae = None
|
||||
try:
|
||||
complete_model_deletion(self.runner.vae)
|
||||
except Exception as e:
|
||||
debug.log(f"Warning during VAE cleanup: {e}", category="warning")
|
||||
finally:
|
||||
self.runner.vae = None
|
||||
|
||||
# Clear other components
|
||||
for component in ['sampler', 'sampling_timesteps', 'schedule', 'config']:
|
||||
@@ -285,9 +269,8 @@ class SeedVR2:
|
||||
|
||||
self.current_model_name = ""
|
||||
|
||||
# Fast RAM cleanup
|
||||
if force_ram_cleanup:
|
||||
fast_ram_cleanup()
|
||||
# Final memory cleanup
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
|
||||
|
||||
def _internal_execute(self, images, model, seed, new_resolution, cfg_scale, batch_size,
|
||||
@@ -306,7 +289,6 @@ class SeedVR2:
|
||||
if model_changed and self.runner is not None:
|
||||
debug.log(f"Model changed from {current_model} to {model}, clearing cache...", category="cache")
|
||||
self.cleanup(
|
||||
force_ram_cleanup=True,
|
||||
cache_model=False, # Don't keep old model
|
||||
debug=debug,
|
||||
)
|
||||
@@ -324,12 +306,8 @@ class SeedVR2:
|
||||
vae_tile_overlap=(vae_tile_overlap, vae_tile_overlap),
|
||||
cached_runner=self.runner if cache_model else None
|
||||
)
|
||||
|
||||
# Set whether to keep the VAE in VRAM for this run
|
||||
self.runner.keep_vae_in_vram = bool(keep_vae_loaded)
|
||||
|
||||
self.current_model_name = model
|
||||
debug.log_memory_state("Model preparation completed")
|
||||
|
||||
debug.end_timer("model_preparation", "Model preparation", force=True, show_breakdown=True)
|
||||
|
||||
@@ -355,19 +333,25 @@ class SeedVR2:
|
||||
if swap_summary and swap_summary.get('total_swaps', 0) > 0:
|
||||
total_time = swap_summary.get('block_total_ms', 0) + swap_summary.get('io_total_ms', 0)
|
||||
debug.log(f"BlockSwap overhead: {total_time:.1f}ms across {swap_summary['total_swaps']} swaps", category="blockswap")
|
||||
|
||||
# Log memory usage summary
|
||||
allocated, reserved, peak = get_vram_usage()
|
||||
debug.log(f"Final VRAM usage - Allocated: {allocated:.2f}GB, Peak: {peak:.2f}GB", category="memory")
|
||||
debug.log_memory_state("Video generation - Memory")
|
||||
|
||||
debug.end_timer("generation_loop", "Video generation completed", show_breakdown=True)
|
||||
debug.log_memory_state("After video generation", detailed_tensors=False)
|
||||
|
||||
debug.log("\n─── Final Cleanup ───", category="none")
|
||||
debug.start_timer("final_cleanup")
|
||||
self.cleanup(force_ram_cleanup=True, cache_model=cache_model, debug=debug)
|
||||
debug.log_memory_state("Final cleanup - Memory", detailed_tensors=False)
|
||||
|
||||
# Perform cleanup (this already calls clear_memory internally)
|
||||
self.cleanup(cache_model=cache_model, debug=debug)
|
||||
|
||||
# Ensure sample is on CPU (ComfyUI expects CPU tensors)
|
||||
if torch.is_tensor(sample) and sample.is_cuda:
|
||||
sample = sample.cpu()
|
||||
|
||||
# Log final memory state after ALL cleanup is done
|
||||
debug.end_timer("final_cleanup", "Final cleanup completed", show_breakdown=True)
|
||||
# Cleanup
|
||||
debug.log_memory_state("After final cleanup", detailed_tensors=True)
|
||||
|
||||
# Final timing summary
|
||||
debug.log("\n─────────", category="none")
|
||||
child_times = {
|
||||
"Model preparation": debug.timer_durations.get("model_preparation", 0),
|
||||
@@ -376,9 +360,10 @@ class SeedVR2:
|
||||
}
|
||||
debug.end_timer("total_execution", "Total execution", show_breakdown=True, custom_children=child_times)
|
||||
debug.log("─────────", category="none")
|
||||
# Clear history for next run
|
||||
|
||||
# Clear history for next run (do this last, after all logging)
|
||||
debug.clear_history()
|
||||
|
||||
|
||||
return (sample,)
|
||||
|
||||
def _progress_callback(self, batch_idx, total_batches, current_batch_frames, message=""):
|
||||
@@ -408,12 +393,21 @@ class SeedVR2:
|
||||
def __del__(self):
|
||||
"""Destructor"""
|
||||
try:
|
||||
debug = self.debug
|
||||
self.cleanup(force_ram_cleanup=True, cache_model=False, debug=debug)
|
||||
# Store debug reference
|
||||
debug = self.debug if hasattr(self, 'debug') else None
|
||||
|
||||
# Full cleanup
|
||||
if hasattr(self, 'cleanup'):
|
||||
self.cleanup(cache_model=False, debug=debug)
|
||||
|
||||
# Clear all remaining references
|
||||
for attr in ['runner', 'text_pos_embeds', 'text_neg_embeds',
|
||||
'current_model_name', 'debug', 'last_batch_time']:
|
||||
if hasattr(self, attr):
|
||||
delattr(self, attr)
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
class SeedVR2BlockSwap:
|
||||
"""Configure block swapping to reduce VRAM usage"""
|
||||
|
||||
|
||||
@@ -52,6 +52,7 @@ from .types import (
|
||||
_memory_device_t,
|
||||
_receptive_field_t,
|
||||
)
|
||||
from src.optimization.memory_manager import clear_memory
|
||||
|
||||
logger = get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -131,22 +132,32 @@ class Upsample3D(Upsample2D):
|
||||
)
|
||||
else:
|
||||
hidden_states = [hidden_states]
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
for i in range(len(hidden_states)):
|
||||
hidden_states[i] = self.upscale_conv(hidden_states[i])
|
||||
hidden_states[i] = rearrange(
|
||||
hidden_states[i],
|
||||
"b (x y z c) f h w -> b c (f z) (h x) (w y)",
|
||||
x=self.spatial_ratio,
|
||||
y=self.spatial_ratio,
|
||||
z=self.temporal_ratio,
|
||||
)
|
||||
# OOM recovery attempt
|
||||
try:
|
||||
hidden_states[i] = self.upscale_conv(hidden_states[i])
|
||||
hidden_states[i] = rearrange(
|
||||
hidden_states[i],
|
||||
"b (x y z c) f h w -> b c (f z) (h x) (w y)",
|
||||
x=self.spatial_ratio,
|
||||
y=self.spatial_ratio,
|
||||
z=self.temporal_ratio,
|
||||
)
|
||||
except Exception as e:
|
||||
debug = getattr(self, 'debug', None)
|
||||
if debug:
|
||||
debug.log("OOM recovery: Upsample3D upscale_conv", category="warning", force=True)
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
time.sleep(1)
|
||||
hidden_states[i] = self.upscale_conv(hidden_states[i])
|
||||
hidden_states[i] = rearrange(
|
||||
hidden_states[i],
|
||||
"b (x y z c) f h w -> b c (f z) (h x) (w y)",
|
||||
x=self.spatial_ratio,
|
||||
y=self.spatial_ratio,
|
||||
z=self.temporal_ratio,
|
||||
)
|
||||
|
||||
# [Overridden] For causal temporal conv
|
||||
if self.temporal_up and memory_state != MemoryState.ACTIVE:
|
||||
@@ -154,18 +165,24 @@ class Upsample3D(Upsample2D):
|
||||
|
||||
if not self.slicing:
|
||||
hidden_states = hidden_states[0]
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
if self.use_conv:
|
||||
if self.name == "conv":
|
||||
hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
|
||||
else:
|
||||
hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
|
||||
# OOM recovery attempt
|
||||
try:
|
||||
if self.name == "conv":
|
||||
hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
|
||||
else:
|
||||
hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
|
||||
except Exception as e:
|
||||
debug = getattr(self, 'debug', None)
|
||||
if debug:
|
||||
debug.log("OOM recovery: Upsample3D conv", category="warning", force=True)
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
time.sleep(1)
|
||||
if self.name == "conv":
|
||||
hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
|
||||
else:
|
||||
hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
|
||||
|
||||
if not self.slicing:
|
||||
return hidden_states
|
||||
@@ -313,19 +330,16 @@ class ResnetBlock3D(ResnetBlock2D):
|
||||
hidden_states = input_tensor
|
||||
|
||||
hidden_states = causal_norm_wrapper(self.norm1, hidden_states, preserve_vram=preserve_vram)
|
||||
# ADD BY NUMZ
|
||||
# OOM recovery attempt
|
||||
try:
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
except Exception as e:
|
||||
if hasattr(self, 'debug') and self.debug:
|
||||
self.debug.log("OOM second chance: ResnetBlock3D", category="warning", force=True)
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(1)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
debug = getattr(self, 'debug', None)
|
||||
if debug:
|
||||
debug.log("OOM recovery: ResnetBlock3D", category="warning", force=True)
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
time.sleep(1)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
if self.upsample is not None:
|
||||
# upsample_nearest_nhwc fails with large batch sizes.
|
||||
|
||||
@@ -27,6 +27,7 @@ from .context_parallel_lib import cache_send_recv, get_cache_size
|
||||
from .global_config import get_norm_limit
|
||||
from .types import MemoryState, _inflation_mode_t, _memory_device_t
|
||||
from ....common.half_precision_fixes import safe_pad_operation
|
||||
from src.optimization.memory_manager import clear_memory
|
||||
|
||||
# Single GPU inference - no distributed processing needed
|
||||
#print("Warning: Using single GPU inference mode - distributed features disabled in causal_inflation_lib")
|
||||
@@ -118,12 +119,11 @@ class InflatedCausalConv3d(Conv3d):
|
||||
x = list(x.split(split_sizes, dim=split_dim))
|
||||
if prev_cache is not None:
|
||||
prev_cache = list(prev_cache.split(split_sizes, dim=split_dim))
|
||||
if preserve_vram:
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
# Memory cleanup when preserve_vram is true
|
||||
'''if preserve_vram:
|
||||
# Use debug if passed through the module
|
||||
debug = getattr(self, 'debug', None)
|
||||
clear_memory(debug=debug, full=False, force=True)'''
|
||||
# Loop Fwd.
|
||||
cache = None
|
||||
for idx in range(len(x)):
|
||||
@@ -167,26 +167,15 @@ class InflatedCausalConv3d(Conv3d):
|
||||
# Update cache.
|
||||
cache = next_cache
|
||||
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
#print("empty cache 1")
|
||||
#time.sleep(2)
|
||||
# OOM recovery attempt
|
||||
try:
|
||||
output = torch.cat(x, split_dim)
|
||||
except Exception as e:
|
||||
if hasattr(self, 'debug') and self.debug:
|
||||
self.debug.log("OOM Second Chance", category="warning", force=True)
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(2)
|
||||
debug = getattr(self, 'debug', None)
|
||||
if debug:
|
||||
debug.log("OOM recovery: Concatenating conv splits", category="warning", force=True)
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
time.sleep(1)
|
||||
output = torch.cat(x, split_dim)
|
||||
return output
|
||||
|
||||
@@ -363,38 +352,26 @@ def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor, preserve_vram: b
|
||||
weights = norm_layer.weight.chunk(num_chunks, dim=0)
|
||||
biases = norm_layer.bias.chunk(num_chunks, dim=0)
|
||||
for i, (w, b) in enumerate(zip(weights, biases)):
|
||||
# OOM recovery attempt
|
||||
try:
|
||||
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
|
||||
except Exception as e:
|
||||
if hasattr(norm_layer, 'debug') and norm_layer.debug:
|
||||
norm_layer.debug.log("OOM Second Chance: Group Norm", category="warning", force=True)
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(2)
|
||||
debug = getattr(norm_layer, 'debug', None)
|
||||
if debug:
|
||||
debug.log("OOM recovery: Group Norm chunk", category="warning", force=True)
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
time.sleep(1)
|
||||
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
|
||||
x[i] = x[i].to(input_dtype)
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
# ADD BY NUMZ
|
||||
# OOM recovery attempt
|
||||
try:
|
||||
x = torch.cat(x, dim=1)
|
||||
except Exception as e:
|
||||
if hasattr(norm_layer, 'debug') and norm_layer.debug:
|
||||
norm_layer.debug.log("OOM Second Chance: Cat", category="warning", force=True)
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(2)
|
||||
debug = getattr(norm_layer, 'debug', None)
|
||||
if debug:
|
||||
debug.log("OOM recovery: Concatenating norm chunks", category="warning", force=True)
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
time.sleep(1)
|
||||
x = torch.cat(x, dim=1)
|
||||
else:
|
||||
x = norm_layer(x)
|
||||
|
||||
@@ -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",
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ import gc
|
||||
import psutil
|
||||
|
||||
from typing import Dict, Any, List, Tuple, Optional, Union
|
||||
from src.optimization.memory_manager import get_vram_usage
|
||||
from src.optimization.memory_manager import clear_memory
|
||||
from src.optimization.compatibility import call_rope_with_stability
|
||||
from src.common.distributed import get_device
|
||||
|
||||
@@ -72,7 +72,6 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) ->
|
||||
raise ValueError("Debug instance must be provided to apply_block_swap_to_dit")
|
||||
|
||||
debug.start_timer("apply_blockswap")
|
||||
debug.log_memory_state("Before BlockSwap")
|
||||
|
||||
# Get the actual model (handle FP8CompatibleDiT wrapper)
|
||||
model = runner.dit
|
||||
@@ -152,8 +151,8 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) ->
|
||||
_protect_model_from_move(model, runner, debug)
|
||||
|
||||
debug.log("BlockSwap configuration complete", category="success")
|
||||
debug.log_memory_state("After BlockSwap")
|
||||
debug.end_timer("apply_blockswap", "BlockSwap configuration applied")
|
||||
debug.log_memory_state("After BlockSwap", detailed_tensors=False)
|
||||
|
||||
|
||||
def _configure_io_components(model, device: str, offload_device: str,
|
||||
@@ -166,7 +165,8 @@ def _configure_io_components(model, device: str, offload_device: str,
|
||||
for name, param in model.named_parameters():
|
||||
if "block" not in name:
|
||||
target_device = offload_device if offload_io_components else device
|
||||
param.data = param.data.to(target_device, non_blocking=use_non_blocking)
|
||||
# Never use non_blocking for initial setup to avoid pinned memory
|
||||
param.data = param.data.to(target_device, non_blocking=False)
|
||||
status = "(offloaded)" if offload_io_components else ""
|
||||
debug.log(f" {name} → {target_device} {status}", category="blockswap")
|
||||
|
||||
@@ -192,6 +192,7 @@ def _configure_blocks(model, device: str, offload_device: str,
|
||||
total_main_memory = 0.0
|
||||
|
||||
# Move blocks based on swap configuration
|
||||
# NEVER use non_blocking for initial CPU offload to avoid pinned memory
|
||||
for b, block in enumerate(model.blocks):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
@@ -199,7 +200,7 @@ def _configure_blocks(model, device: str, offload_device: str,
|
||||
block.to(device)
|
||||
total_main_memory += block_memory
|
||||
else:
|
||||
block.to(offload_device, non_blocking=use_non_blocking)
|
||||
block.to(offload_device, non_blocking=False)
|
||||
total_offload_memory += block_memory
|
||||
|
||||
# Ensure all buffers match their containing module's device
|
||||
@@ -207,13 +208,13 @@ def _configure_blocks(model, device: str, offload_device: str,
|
||||
target_device = device if b > model.blocks_to_swap else offload_device
|
||||
for name, buffer in block.named_buffers():
|
||||
if buffer.device != torch.device(target_device):
|
||||
buffer.data = buffer.data.to(target_device)
|
||||
buffer.data = buffer.data.to(target_device, non_blocking=False)
|
||||
|
||||
# Clean up memory
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
gc.collect()
|
||||
# Only force clear memory if we actually moved blocks
|
||||
if model.blocks_to_swap > 0:
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
else:
|
||||
clear_memory(debug=debug, full=True, force=False)
|
||||
|
||||
return {
|
||||
"offload_memory": total_offload_memory,
|
||||
@@ -278,17 +279,18 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
|
||||
target_device = torch.device(model.main_device)
|
||||
|
||||
if current_device != target_device:
|
||||
self.to(model.main_device, non_blocking=model.use_non_blocking)
|
||||
|
||||
# Synchronize if needed
|
||||
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking and torch.cuda.is_available():
|
||||
torch.cuda.synchronize(get_device())
|
||||
# CPU->GPU: never use non_blocking (prevents pinned memory allocation)
|
||||
self.to(model.main_device, non_blocking=False)
|
||||
|
||||
# Execute forward pass with OOM protection
|
||||
output = original_forward(*args, **kwargs)
|
||||
|
||||
# Move back to offload device
|
||||
self.to(model.offload_device, non_blocking=model.use_non_blocking)
|
||||
# Only use non_blocking for GPU->GPU transfers (following WanVideo pattern)
|
||||
if model.use_non_blocking and model.offload_device != "cpu":
|
||||
self.to(model.offload_device, non_blocking=True)
|
||||
else:
|
||||
self.to(model.offload_device, non_blocking=False)
|
||||
|
||||
# Log timing if debug is available
|
||||
if debug and t_start is not None:
|
||||
@@ -299,13 +301,7 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
|
||||
)
|
||||
|
||||
# Only clear cache under memory pressure
|
||||
if torch.mps.is_available():
|
||||
mem = psutil.virtual_memory()
|
||||
if torch.mps.current_allocated_memory() > mem.total * 0.9:
|
||||
torch.mps.empty_cache()
|
||||
if torch.cuda.is_available() and torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
clear_memory(debug=debug, full=True, force=False)
|
||||
else:
|
||||
output = original_forward(*args, **kwargs)
|
||||
|
||||
@@ -353,20 +349,18 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
|
||||
|
||||
# Move to GPU for computation if needed
|
||||
if current_device != target_device:
|
||||
self.to(model.main_device)
|
||||
|
||||
# Synchronize if not using non-blocking transfers
|
||||
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking:
|
||||
if torch.mps.is_available():
|
||||
torch.mps.synchronize()
|
||||
else:
|
||||
torch.cuda.synchronize(get_device())
|
||||
# CPU->GPU: never use non_blocking
|
||||
self.to(model.main_device, non_blocking=False)
|
||||
|
||||
# Execute forward pass
|
||||
output = self._original_forward(*args, **kwargs)
|
||||
|
||||
# Move back to offload device
|
||||
self.to(model.offload_device, non_blocking=model.use_non_blocking)
|
||||
# Only use non_blocking for GPU->GPU transfers
|
||||
if model.use_non_blocking and model.offload_device != "cpu":
|
||||
self.to(model.offload_device, non_blocking=True)
|
||||
else:
|
||||
self.to(model.offload_device, non_blocking=False)
|
||||
|
||||
# Log timing if debug is available
|
||||
if debug and t_start is not None:
|
||||
@@ -377,13 +371,7 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
|
||||
)
|
||||
|
||||
# Only clear cache under memory pressure
|
||||
if torch.mps.is_available():
|
||||
mem = psutil.virtual_memory()
|
||||
if torch.mps.current_allocated_memory() > mem.total * 0.9:
|
||||
torch.mps.empty_cache()
|
||||
if torch.cuda.is_available() and torch.cuda.memory_allocated(get_device()) > torch.cuda.get_device_properties(get_device()).total_memory * 0.9:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
clear_memory(debug=debug, full=True, force=False)
|
||||
|
||||
return output
|
||||
|
||||
@@ -666,6 +654,16 @@ def cleanup_blockswap(runner, keep_state_for_cache: bool = False) -> None:
|
||||
# Move model to CPU to free VRAM
|
||||
if not keep_state_for_cache:
|
||||
model.to("cpu")
|
||||
# Ensure all sub-modules are also on CPU
|
||||
for module in model.modules():
|
||||
if hasattr(module, '_parameters'):
|
||||
for param in module._parameters.values():
|
||||
if param is not None and param.is_cuda:
|
||||
param.data = param.data.cpu()
|
||||
if hasattr(module, '_buffers'):
|
||||
for buffer in module._buffers.values():
|
||||
if buffer is not None and buffer.is_cuda:
|
||||
buffer.data = buffer.data.cpu()
|
||||
debug.log("Moved model to CPU", category="store")
|
||||
|
||||
# Clean up runner attributes
|
||||
@@ -684,15 +682,3 @@ def cleanup_blockswap(runner, keep_state_for_cache: bool = False) -> None:
|
||||
|
||||
# Clear local debug reference
|
||||
debug = None
|
||||
|
||||
# Force garbage collection (multiple passes for thorough cleanup)
|
||||
gc.collect(2) # Full collection including oldest generation
|
||||
gc.collect()
|
||||
gc.collect()
|
||||
|
||||
# Final memory cleanup
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
+338
-105
@@ -4,13 +4,13 @@ Handles VRAM usage, cache management, and memory optimization
|
||||
|
||||
Extracted from: seedvr2.py (lines 373-405, 607-626, 1016-1044)
|
||||
"""
|
||||
|
||||
import os
|
||||
\
|
||||
import torch
|
||||
import gc
|
||||
import sys
|
||||
import time
|
||||
import psutil
|
||||
from typing import Tuple, Optional
|
||||
from typing import Tuple, Optional, Dict, Any
|
||||
from src.common.cache import Cache
|
||||
from src.models.dit_v2.rope import RotaryEmbeddingBase
|
||||
from src.common.distributed import get_device
|
||||
@@ -38,74 +38,283 @@ def get_device_list():
|
||||
return devs[1:]
|
||||
return devs
|
||||
|
||||
def get_basic_vram_info():
|
||||
if torch.mps.is_available():
|
||||
mem = psutil.virtual_memory()
|
||||
free_memory = mem.total - mem.used
|
||||
total_memory = mem.total
|
||||
else:
|
||||
"""🔍 Méthode basique avec PyTorch natif"""
|
||||
if not torch.cuda.is_available():
|
||||
return {"error": "CUDA not available"}
|
||||
# Mémoire libre et totale (en bytes)
|
||||
free_memory, total_memory = torch.cuda.mem_get_info(get_device())
|
||||
def get_basic_vram_info() -> Dict[str, Any]:
|
||||
"""
|
||||
Get basic VRAM availability info (free and total memory).
|
||||
Used for capacity planning and initial checks.
|
||||
|
||||
# Conversion en GB
|
||||
free_gb = free_memory / (1024**3)
|
||||
total_gb = total_memory / (1024**3)
|
||||
|
||||
return {
|
||||
"free_gb": free_gb,
|
||||
"total_gb": total_gb
|
||||
}
|
||||
Returns:
|
||||
dict: {"free_gb": float, "total_gb": float} or {"error": str}
|
||||
"""
|
||||
try:
|
||||
if torch.cuda.is_available():
|
||||
device = get_device()
|
||||
free_memory, total_memory = torch.cuda.mem_get_info(device)
|
||||
elif torch.mps.is_available():
|
||||
mem = psutil.virtual_memory()
|
||||
free_memory = mem.total - mem.used
|
||||
total_memory = mem.total
|
||||
else:
|
||||
return {"error": "No GPU backend available (CUDA/MPS)"}
|
||||
|
||||
return {
|
||||
"free_gb": free_memory / (1024**3),
|
||||
"total_gb": total_memory / (1024**3)
|
||||
}
|
||||
except Exception as e:
|
||||
return {"error": f"Failed to get memory info: {str(e)}"}
|
||||
|
||||
# Initial VRAM check at module load
|
||||
vram_info = get_basic_vram_info()
|
||||
if "error" not in vram_info:
|
||||
print(f"📊 Initial VRAM status: {vram_info['free_gb']:.2f}GB free / {vram_info['total_gb']:.2f}GB total")
|
||||
backend = "MPS" if torch.mps.is_available() else "CUDA"
|
||||
print(f"📊 Initial {backend} memory: {vram_info['free_gb']:.2f}GB free / {vram_info['total_gb']:.2f}GB total")
|
||||
else:
|
||||
print(f"⚠️ VRAM check: {vram_info['error']} - No available backend!")
|
||||
print(f"⚠️ Memory check failed: {vram_info['error']} - No available backend!")
|
||||
|
||||
def get_vram_usage() -> Tuple[float, float, float]:
|
||||
"""
|
||||
Get current VRAM usage (allocated, reserved, peak)
|
||||
Get current VRAM usage metrics for monitoring.
|
||||
Used for tracking memory consumption during processing.
|
||||
|
||||
Returns:
|
||||
tuple: (allocated_gb, reserved_gb, max_allocated_gb)
|
||||
Returns (0, 0, 0) if CUDA not available
|
||||
Returns (0, 0, 0) if no GPU available
|
||||
"""
|
||||
if torch.mps.is_available():
|
||||
allocated = torch.mps.current_allocated_memory() / (1024**3)
|
||||
reserved = torch.mps.driver_allocated_memory() / (1024**3)
|
||||
max_allocated = 0
|
||||
return allocated, reserved, max_allocated
|
||||
if torch.cuda.is_available():
|
||||
allocated = torch.cuda.memory_allocated(get_device()) / (1024**3)
|
||||
reserved = torch.cuda.memory_reserved(get_device()) / (1024**3)
|
||||
max_allocated = torch.cuda.max_memory_allocated(get_device()) / (1024**3)
|
||||
return allocated, reserved, max_allocated
|
||||
return 0, 0, 0
|
||||
try:
|
||||
if torch.cuda.is_available():
|
||||
device = get_device()
|
||||
allocated = torch.cuda.memory_allocated(device) / (1024**3)
|
||||
reserved = torch.cuda.memory_reserved(device) / (1024**3)
|
||||
max_allocated = torch.cuda.max_memory_allocated(device) / (1024**3)
|
||||
return allocated, reserved, max_allocated
|
||||
elif torch.mps.is_available():
|
||||
allocated = torch.mps.current_allocated_memory() / (1024**3)
|
||||
reserved = torch.mps.driver_allocated_memory() / (1024**3)
|
||||
max_allocated = allocated # MPS doesn't track peak separately
|
||||
return allocated, reserved, max_allocated
|
||||
except Exception:
|
||||
pass
|
||||
return 0.0, 0.0, 0.0
|
||||
|
||||
|
||||
def clear_vram_cache(debug) -> None:
|
||||
"""Clear VRAM cache and run garbage collection"""
|
||||
def get_ram_usage() -> Tuple[float, float, float, float]:
|
||||
"""
|
||||
Get current RAM usage metrics for the current process.
|
||||
Provides accurate tracking of process-specific memory consumption.
|
||||
|
||||
Returns:
|
||||
tuple: (process_gb, available_gb, total_gb, used_by_others_gb)
|
||||
Returns (0, 0, 0, 0) if psutil not available
|
||||
"""
|
||||
try:
|
||||
if not psutil:
|
||||
return 0.0, 0.0, 0.0, 0.0
|
||||
|
||||
# Get current process memory
|
||||
process = psutil.Process()
|
||||
process_memory = process.memory_info()
|
||||
process_gb = process_memory.rss / (1024**3)
|
||||
|
||||
debug.log("Clearing VRAM cache...", category="cleanup")
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
# Get system memory
|
||||
sys_memory = psutil.virtual_memory()
|
||||
total_gb = sys_memory.total / (1024**3)
|
||||
available_gb = sys_memory.available / (1024**3)
|
||||
|
||||
# Calculate memory used by other processes
|
||||
# This is the CORRECT calculation:
|
||||
total_used_gb = total_gb - available_gb # Total memory used by ALL processes
|
||||
used_by_others_gb = max(0, total_used_gb - process_gb) # Subtract current process
|
||||
|
||||
return process_gb, available_gb, total_gb, used_by_others_gb
|
||||
|
||||
except Exception:
|
||||
return 0.0, 0.0, 0.0, 0.0
|
||||
|
||||
|
||||
# Global cache for OS libraries (initialized once)
|
||||
_os_memory_lib = None
|
||||
|
||||
|
||||
def clear_memory(debug=None, full=False, force=True) -> None:
|
||||
"""
|
||||
Clear memory caches with two-tier approach for optimal performance.
|
||||
|
||||
Args:
|
||||
debug: Debug instance for logging (optional)
|
||||
force: If True, always clear. If False, only clear when <15% free
|
||||
full: If True, perform full cleanup including GC and OS operations.
|
||||
If False (default), only perform minimal GPU cache clearing.
|
||||
|
||||
Two-tier approach:
|
||||
- Minimal mode (full=False): GPU cache operations (~1-5ms)
|
||||
Used for frequent calls during batch processing
|
||||
- Full mode (full=True): Complete cleanup with GC and OS operations (~10-50ms)
|
||||
Used at key points like model switches or final cleanup
|
||||
"""
|
||||
global _os_memory_lib
|
||||
|
||||
# Check if we should clear based on memory pressure
|
||||
if not force:
|
||||
should_clear = False
|
||||
|
||||
# Use existing function for memory info
|
||||
mem_info = get_basic_vram_info()
|
||||
|
||||
if "error" not in mem_info:
|
||||
# Check VRAM/MPS memory pressure (15% free threshold)
|
||||
free_ratio = mem_info["free_gb"] / mem_info["total_gb"]
|
||||
if free_ratio < 0.15:
|
||||
should_clear = True
|
||||
if debug:
|
||||
backend = "MPS" if torch.mps.is_available() else "VRAM"
|
||||
debug.log(f"{backend} pressure: {mem_info['free_gb']:.1f}GB free of {mem_info['total_gb']:.1f}GB", category="memory")
|
||||
|
||||
# For non-MPS systems, also check system RAM separately
|
||||
if not should_clear and not torch.mps.is_available():
|
||||
mem = psutil.virtual_memory()
|
||||
if mem.available < mem.total * 0.15:
|
||||
should_clear = True
|
||||
if debug:
|
||||
debug.log(f"RAM pressure: {mem.available/(1024**3):.1f}GB free of {mem.total/(1024**3):.1f}GB", category="memory")
|
||||
|
||||
if not should_clear:
|
||||
return
|
||||
|
||||
# Determine cleanup level
|
||||
cleanup_mode = "full" if full else "minimal"
|
||||
if debug:
|
||||
debug.log(f"Clearing memory caches ({cleanup_mode})...", category="cleanup")
|
||||
|
||||
# ===== MINIMAL OPERATIONS (Always performed) =====
|
||||
# Step 1: Clear GPU caches - Fast operations (~1-5ms)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
gc.collect()
|
||||
elif torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
|
||||
# ===== FULL OPERATIONS (Only when full=True) =====
|
||||
if full:
|
||||
# Step 2: Clear PyTorch internal caches
|
||||
if hasattr(torch, '_C'):
|
||||
try:
|
||||
torch._C._clear_cache()
|
||||
except:
|
||||
pass
|
||||
|
||||
# Step 3: Full garbage collection (expensive ~5-20ms)
|
||||
gc.collect(2)
|
||||
|
||||
# Step 4: Return memory to OS (platform-specific, ~5-30ms)
|
||||
try:
|
||||
if sys.platform == 'linux':
|
||||
# Linux: malloc_trim
|
||||
import ctypes # Import only when needed
|
||||
if _os_memory_lib is None:
|
||||
_os_memory_lib = ctypes.CDLL("libc.so.6")
|
||||
_os_memory_lib.malloc_trim(0)
|
||||
|
||||
elif sys.platform == 'win32':
|
||||
# Windows: Trim working set
|
||||
import ctypes # Import only when needed
|
||||
if _os_memory_lib is None:
|
||||
_os_memory_lib = ctypes.windll.kernel32
|
||||
handle = _os_memory_lib.GetCurrentProcess()
|
||||
_os_memory_lib.SetProcessWorkingSetSize(handle, -1, -1)
|
||||
|
||||
elif torch.mps.is_available():
|
||||
# macOS with MPS
|
||||
import ctypes # Import only when needed
|
||||
import ctypes.util
|
||||
if _os_memory_lib is None:
|
||||
libc_path = ctypes.util.find_library('c')
|
||||
if libc_path:
|
||||
_os_memory_lib = ctypes.CDLL(libc_path)
|
||||
|
||||
if _os_memory_lib:
|
||||
_os_memory_lib.sync()
|
||||
except:
|
||||
# OS-specific memory operations are optional
|
||||
pass
|
||||
|
||||
|
||||
def manage_vae_device(runner, target_device: str, preserve_vram: bool = False,
|
||||
debug=None, reason: str = None) -> bool:
|
||||
"""
|
||||
Manage VAE device placement with intelligent movement and logging.
|
||||
|
||||
Args:
|
||||
runner: Runner instance containing the VAE
|
||||
target_device: Target device ('cuda:0', 'cpu', etc.)
|
||||
preserve_vram: Whether preserve_vram mode is active
|
||||
debug: Debug instance for logging
|
||||
reason: Optional custom reason for the movement
|
||||
|
||||
Returns:
|
||||
bool: True if VAE was moved, False if already on target device
|
||||
"""
|
||||
if not hasattr(runner, 'vae') or runner.vae is None:
|
||||
return False
|
||||
|
||||
# Get current VAE device
|
||||
current_device = next(runner.vae.parameters()).device if hasattr(runner.vae, 'parameters') else None
|
||||
if current_device is None:
|
||||
return False
|
||||
|
||||
# Normalize device strings for comparison
|
||||
target_type = target_device.split(':')[0] if ':' in target_device else target_device
|
||||
current_type = str(current_device.type)
|
||||
|
||||
# Skip if already on target device
|
||||
if current_type == target_type:
|
||||
return False
|
||||
|
||||
# Determine reason for movement
|
||||
if reason:
|
||||
reason = reason
|
||||
elif preserve_vram:
|
||||
reason = "preserve_vram"
|
||||
else:
|
||||
reason = "inference requirement"
|
||||
|
||||
# Start timer based on direction
|
||||
timer_name = "vae_to_gpu" if target_type != 'cpu' else "vae_to_cpu"
|
||||
if debug:
|
||||
debug.start_timer(timer_name)
|
||||
|
||||
# Log the movement
|
||||
if debug:
|
||||
if target_type == 'cpu':
|
||||
debug.log(f"Moving VAE to CPU ({reason})", category="memory")
|
||||
else:
|
||||
debug.log(f"Moving VAE from {current_type} to {target_device} ({reason})", category="memory")
|
||||
|
||||
# Move VAE
|
||||
runner.vae = runner.vae.to(target_device)
|
||||
|
||||
# End timer
|
||||
if debug:
|
||||
if target_type == 'cpu':
|
||||
debug.end_timer(timer_name, "VAE moved to CPU")
|
||||
else:
|
||||
debug.end_timer(timer_name, "VAE moved to GPU")
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def reset_vram_peak(debug) -> None:
|
||||
"""
|
||||
Reset VRAM peak counter for new tracking
|
||||
Reset VRAM peak memory statistics for fresh tracking.
|
||||
"""
|
||||
debug.log("Resetting VRAM peak memory statistics", category="memory")
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.reset_peak_memory_stats(get_device())
|
||||
try:
|
||||
if torch.cuda.is_available():
|
||||
device = get_device()
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
# MPS doesn't support peak memory reset
|
||||
except Exception as e:
|
||||
debug.log(f"Failed to reset peak memory stats: {e}", category="warning")
|
||||
|
||||
def preinitialize_rope_cache(runner, debug) -> None:
|
||||
"""
|
||||
@@ -171,8 +380,8 @@ def preinitialize_rope_cache(runner, debug) -> None:
|
||||
except Exception as e:
|
||||
debug.log(f"Failed for {cache_key}: {e}", level="WARNING", category="cache")
|
||||
# Return empty tensors as fallback
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
time.sleep(1)
|
||||
clear_vram_cache(debug)
|
||||
|
||||
return torch.zeros(1, 64)
|
||||
|
||||
@@ -200,57 +409,92 @@ def clear_rope_lru_caches(model) -> int:
|
||||
"""Clear ALL LRU caches from RoPE modules"""
|
||||
cleared_count = 0
|
||||
|
||||
for name, module in model.named_modules():
|
||||
if hasattr(module, 'get_axial_freqs') and hasattr(module.get_axial_freqs, 'cache_clear'):
|
||||
module.get_axial_freqs.cache_clear()
|
||||
cleared_count += 1
|
||||
if model is None:
|
||||
return 0
|
||||
|
||||
try:
|
||||
for name, module in model.named_modules():
|
||||
if hasattr(module, 'get_axial_freqs') and hasattr(module.get_axial_freqs, 'cache_clear'):
|
||||
module.get_axial_freqs.cache_clear()
|
||||
cleared_count += 1
|
||||
except AttributeError:
|
||||
# Model structure already damaged, skip
|
||||
pass
|
||||
|
||||
return cleared_count
|
||||
|
||||
|
||||
def fast_model_cleanup(model):
|
||||
"""Fast model cleanup without logs"""
|
||||
def complete_model_deletion(model):
|
||||
"""Completely delete a model and free all its memory"""
|
||||
if model is None:
|
||||
return
|
||||
|
||||
# Move to CPU
|
||||
model.to("cpu")
|
||||
|
||||
# Clear parameters and buffers recursively
|
||||
def clear_recursive(m):
|
||||
for child in m.children():
|
||||
clear_recursive(child)
|
||||
for param in m.parameters():
|
||||
if param is not None:
|
||||
param.data = param.data.cpu()
|
||||
param.grad = None
|
||||
for buffer in m.buffers():
|
||||
if buffer is not None:
|
||||
buffer.data = buffer.data.cpu()
|
||||
|
||||
clear_recursive(model)
|
||||
|
||||
|
||||
def fast_ram_cleanup():
|
||||
"""Fast RAM cleanup without excessive logging"""
|
||||
# Garbage collection
|
||||
gc.collect()
|
||||
|
||||
# Clear MPS cache
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
# Clear CUDA cache
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
torch.cuda.reset_peak_memory_stats(get_device())
|
||||
|
||||
# Clear PyTorch internal caches
|
||||
try:
|
||||
torch._C._clear_cache()
|
||||
except:
|
||||
# Move to CPU first
|
||||
model.to("cpu")
|
||||
|
||||
# Clear parameters and buffers recursively and release storage
|
||||
def clear_recursive(m):
|
||||
# Process children first
|
||||
for child in m.children():
|
||||
clear_recursive(child)
|
||||
|
||||
# Clear parameters and release storage
|
||||
if hasattr(m, '_parameters'):
|
||||
for param_name, param in list(m._parameters.items()):
|
||||
if param is not None:
|
||||
param.data = param.data.cpu()
|
||||
param.grad = None
|
||||
# Release underlying storage
|
||||
if param.data.numel() > 0:
|
||||
param.data.set_()
|
||||
|
||||
# Clear buffers and release storage
|
||||
if hasattr(m, '_buffers'):
|
||||
for buffer_name, buffer in list(m._buffers.items()):
|
||||
if buffer is not None:
|
||||
buffer.data = buffer.data.cpu()
|
||||
# Release underlying storage
|
||||
if buffer.data.numel() > 0:
|
||||
buffer.data.set_()
|
||||
|
||||
clear_recursive(model)
|
||||
|
||||
# Clear all module dicts but keep the structure
|
||||
if hasattr(model, 'modules'):
|
||||
for module in model.modules():
|
||||
# Clear custom attributes but keep PyTorch internals
|
||||
if hasattr(module, '__dict__'):
|
||||
keys_to_delete = []
|
||||
for key in module.__dict__.keys():
|
||||
# Keep PyTorch internal attributes
|
||||
if not key.startswith('_') or key.startswith('_original_'):
|
||||
keys_to_delete.append(key)
|
||||
for key in keys_to_delete:
|
||||
try:
|
||||
delattr(module, key)
|
||||
except:
|
||||
pass
|
||||
|
||||
# Now clear the model's dict
|
||||
if hasattr(model, '__dict__'):
|
||||
# Clear everything except PyTorch internals
|
||||
keys_to_delete = []
|
||||
for key in model.__dict__.keys():
|
||||
if not key in ['_modules', '_parameters', '_buffers', 'training']:
|
||||
keys_to_delete.append(key)
|
||||
for key in keys_to_delete:
|
||||
try:
|
||||
delattr(model, key)
|
||||
except:
|
||||
pass
|
||||
except AttributeError:
|
||||
# Model already partially cleaned, that's OK
|
||||
pass
|
||||
|
||||
# Final cleanup - now we can clear everything
|
||||
if hasattr(model, '__dict__'):
|
||||
model.__dict__.clear()
|
||||
|
||||
def clear_all_caches(runner, debug, offload_vae=False) -> int:
|
||||
"""
|
||||
@@ -374,9 +618,7 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
|
||||
|
||||
# Handle VAE offloading if requested
|
||||
if offload_vae and hasattr(runner, 'vae') and runner.vae is not None:
|
||||
debug.log("Moving VAE to CPU and clearing intermediate tensors", category="cleanup")
|
||||
|
||||
# Clear any intermediate tensors/buffers in VAE
|
||||
# Clear intermediate tensors BEFORE moving to CPU (more efficient)
|
||||
vae_caches_cleared = 0
|
||||
for module in runner.vae.modules():
|
||||
# Clear module-specific caches
|
||||
@@ -388,7 +630,7 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
|
||||
# Clear any CUDA tensors in module attributes
|
||||
for attr_name in list(vars(module).keys()):
|
||||
attr = getattr(module, attr_name, None)
|
||||
if torch.is_tensor(attr) and attr.is_cuda:
|
||||
if torch.is_tensor(attr) and (attr.is_cuda or attr.is_mps):
|
||||
# Move tensor to CPU if it's not a parameter/buffer
|
||||
if attr_name not in module._parameters and attr_name not in module._buffers:
|
||||
setattr(module, attr_name, attr.cpu())
|
||||
@@ -397,21 +639,12 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
|
||||
if vae_caches_cleared > 0:
|
||||
debug.log(f"Cleared {vae_caches_cleared} VAE caches", category="success")
|
||||
|
||||
# Move entire VAE to CPU (preserves model for reuse)
|
||||
runner.vae = runner.vae.to('cpu')
|
||||
debug.log("VAE moved to CPU, intermediate tensors cleared", category="success")
|
||||
# Now move VAE to CPU using helper
|
||||
manage_vae_device(runner, 'cpu', preserve_vram=True, debug=debug)
|
||||
cleaned_items += vae_caches_cleared
|
||||
|
||||
# Force garbage collection
|
||||
gc.collect(2) # Collect all generations
|
||||
|
||||
# Clear MPS cache
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
# Clear CUDA cache
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
# Final memory cleanup
|
||||
clear_memory(debug=debug, full=True, force=True)
|
||||
|
||||
return cleaned_items
|
||||
|
||||
+260
-193
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user