From a8d7153bc33df75997a2b2642551f92cf86f78ce Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Sun, 19 Oct 2025 08:41:45 -0400 Subject: [PATCH] Refactor: unify tensor management and enforce compute_dtype throughout pipeline - Rename manage_tensor_device -> manage_tensor with unified device/dtype handling - Convert VAE outputs (float16) to compute_dtype (bfloat16) immediately after encode/decode - Align alpha channel to compute_dtype at RGBA concatenation point - Maintain float32 precision for alpha processing numerical stability - Optimize conversions during offload operations to minimize overhead - speed/VRAM improvements through reduced dtype conversions and better consistency --- src/core/alpha_upscaling.py | 44 ++++++++-------- src/core/generation.py | 80 +++++++++++++++++++----------- src/optimization/memory_manager.py | 29 +++++++---- 3 files changed, 93 insertions(+), 60 deletions(-) diff --git a/src/core/alpha_upscaling.py b/src/core/alpha_upscaling.py index 37a8abe..5764cf6 100644 --- a/src/core/alpha_upscaling.py +++ b/src/core/alpha_upscaling.py @@ -12,7 +12,7 @@ import cv2 import numpy as np from typing import Optional, Any, List from ..common.half_precision_fixes import ensure_float32_precision -from ..optimization.memory_manager import manage_tensor_device +from ..optimization.memory_manager import manage_tensor def process_alpha_for_batch( @@ -20,6 +20,7 @@ def process_alpha_for_batch( alpha_original: torch.Tensor, rgb_original: torch.Tensor, device: torch.device, + compute_dtype: torch.dtype, debug: Optional[Any] = None ) -> List[torch.Tensor]: """ @@ -27,24 +28,25 @@ def process_alpha_for_batch( Called during postprocess phase when VRAM is available. Args: - rgb_samples: List of decoded RGB samples + rgb_samples: List of decoded RGB samples from VAE alpha_original: Original Alpha channel (C, T, H, W) rgb_original: Original RGB for guidance (C, T, H, W) - device: Target device for processing - debug: Debug instance + device: Target device for processing (typically CUDA) + compute_dtype: Pipeline compute dtype (e.g., bfloat16) for final output. + debug: Debug instance for logging Returns: - List of RGBA samples with Alpha merged + List of RGBA samples with Alpha merged, in compute_dtype """ # Move alpha and RGB guidance tensors to processing device (GPU) - alpha_original = manage_tensor_device( + alpha_original = manage_tensor( tensor=alpha_original, target_device=device, tensor_name="alpha_original", debug=debug, reason="Alpha processing" ) - rgb_original = manage_tensor_device( + rgb_original = manage_tensor( tensor=rgb_original, target_device=device, tensor_name="rgb_original", @@ -57,7 +59,7 @@ def process_alpha_for_batch( for rgb_sample in rgb_samples: # Move RGB sample to processing device before use - rgb_sample = manage_tensor_device( + rgb_sample = manage_tensor( tensor=rgb_sample, target_device=device, tensor_name="rgb_sample", @@ -87,6 +89,17 @@ def process_alpha_for_batch( debug=debug ) + # Convert Alpha from float32 to compute dtype + if alpha_upscaled.dtype != compute_dtype: + alpha_upscaled = manage_tensor( + tensor=alpha_upscaled, + target_device=alpha_upscaled.device, + tensor_name="alpha_upscaled", + dtype=compute_dtype, + debug=debug, + reason="dtype alignment for RGBA concatenation" + ) + # Concatenate RGB and upscaled alpha to create RGBA output (T, 4, H, W) rgba_sample = torch.cat([rgb_sample_4d[:, :3, :, :], alpha_upscaled], dim=1) @@ -158,7 +171,7 @@ def detect_edges_batch( edges = edges.unsqueeze(1) # Move back to images device after CPU numpy processing - edges = manage_tensor_device( + edges = manage_tensor( tensor=edges, target_device=images.device, tensor_name="edge_map", @@ -203,7 +216,7 @@ def guided_filter_pytorch(guide: torch.Tensor, src: torch.Tensor, # Restore original dtype after float32 processing if output.dtype != guide_dtype: - output = manage_tensor_device( + output = manage_tensor( tensor=output, target_device=output.device, tensor_name="filtered_output", @@ -400,17 +413,6 @@ def edge_guided_alpha_upscale( # Clamp output to valid alpha range [0, 1] alpha_final = alpha_final.clamp(0, 1) - - # Restore original dtype after float32 processing - if alpha_final.dtype != alpha_dtype: - alpha_final = manage_tensor_device( - tensor=alpha_final, - target_device=alpha_final.device, - tensor_name="alpha_final", - dtype=alpha_dtype, - debug=debug, - reason="dtype restoration" - ) if debug: final_values = alpha_final.flatten() diff --git a/src/core/generation.py b/src/core/generation.py index 7a3e8d5..c41b0b1 100644 --- a/src/core/generation.py +++ b/src/core/generation.py @@ -32,6 +32,7 @@ from .alpha_upscaling import process_alpha_for_batch from .infer import VideoDiffusionInfer from .model_manager import configure_runner, materialize_model, apply_model_specific_config from ..common.seed import set_seed +from ..common.half_precision_fixes import ensure_float32_precision from ..data.image.transforms.divisible_crop import DivisibleCrop from ..data.image.transforms.na_resize import NaResize from ..optimization.memory_manager import ( @@ -39,7 +40,7 @@ from ..optimization.memory_manager import ( cleanup_vae, cleanup_text_embeddings, clear_memory, - manage_tensor_device, + manage_tensor, manage_model_device, release_tensor_memory, release_tensor_collection @@ -135,7 +136,7 @@ def load_text_embeddings(script_directory: str, device: torch.device, text_pos_embeds = torch.load(os.path.join(script_directory, 'pos_emb.pt')) text_neg_embeds = torch.load(os.path.join(script_directory, 'neg_emb.pt')) - text_pos_embeds = manage_tensor_device( + text_pos_embeds = manage_tensor( tensor=text_pos_embeds, target_device=device, tensor_name="text_pos_embeds", @@ -143,7 +144,7 @@ def load_text_embeddings(script_directory: str, device: torch.device, debug=debug, reason="DiT inference" ) - text_neg_embeds = manage_tensor_device( + text_neg_embeds = manage_tensor( tensor=text_neg_embeds, target_device=device, tensor_name="text_neg_embeds", @@ -690,7 +691,7 @@ def encode_all_batches( # Process current batch video = images[start_idx:end_idx] video = video.permute(0, 3, 1, 2) - video = manage_tensor_device( + video = manage_tensor( tensor=video, target_device=ctx['vae_device'], tensor_name=f"video_batch_{encode_idx+1}", @@ -751,7 +752,7 @@ def encode_all_batches( # Store transformed video on offload device if needed for color correction if color_correction != "none": if ctx['tensor_offload_device'] is not None and (transformed_video.is_cuda or transformed_video.is_mps): - ctx['all_transformed_videos'][encode_idx] = manage_tensor_device( + ctx['all_transformed_videos'][encode_idx] = manage_tensor( tensor=transformed_video, target_device=ctx['tensor_offload_device'], tensor_name=f"transformed_video_{encode_idx+1}", @@ -774,14 +775,14 @@ def encode_all_batches( # Store on tensor_offload_device to save VRAM (or keep on device if none) if ctx['tensor_offload_device'] is not None: - ctx['all_alpha_channels'][encode_idx] = manage_tensor_device( + ctx['all_alpha_channels'][encode_idx] = manage_tensor( tensor=alpha_channel, target_device=ctx['tensor_offload_device'], tensor_name=f"alpha_channel_{encode_idx+1}", debug=debug, reason="storing Alpha channel for upscaling" ) - ctx['all_input_rgb'][encode_idx] = manage_tensor_device( + ctx['all_input_rgb'][encode_idx] = manage_tensor( tensor=rgb_video_original, target_device=ctx['tensor_offload_device'], tensor_name=f"rgb_original_{encode_idx+1}", @@ -797,7 +798,7 @@ def encode_all_batches( del video # Move to VAE device with correct dtype for encoding (no-op if already there) - transformed_video = manage_tensor_device( + transformed_video = manage_tensor( tensor=transformed_video, target_device=ctx['vae_device'], tensor_name=f"transformed_video_{encode_idx+1}", @@ -808,17 +809,26 @@ def encode_all_batches( # Encode to latents cond_latents = runner.vae_encode([transformed_video]) - # Offload latents to avoid VRAM accumulation + # Convert from VAE dtype to compute dtype and offload to avoid VRAM accumulation if ctx['tensor_offload_device'] is not None and (cond_latents[0].is_cuda or cond_latents[0].is_mps): - ctx['all_latents'][encode_idx] = manage_tensor_device( + ctx['all_latents'][encode_idx] = manage_tensor( tensor=cond_latents[0], target_device=ctx['tensor_offload_device'], tensor_name=f"latent_{encode_idx+1}", + dtype=ctx['compute_dtype'], debug=debug, - reason="storing encoded latents for upscaling" + reason="storing encoded latents for upscaling (VAE dtype → compute dtype)" ) else: - ctx['all_latents'][encode_idx] = cond_latents[0] + # Stay on current device but convert to compute dtype + ctx['all_latents'][encode_idx] = manage_tensor( + tensor=cond_latents[0], + target_device=cond_latents[0].device, + tensor_name=f"latent_{encode_idx+1}", + dtype=ctx['compute_dtype'], + debug=debug, + reason="VAE dtype → compute dtype" + ) del cond_latents, transformed_video @@ -970,7 +980,7 @@ def upscale_all_batches( debug.start_timer(f"upscale_batch_{upscale_idx+1}") # Move to DiT device with correct dtype for upscaling (no-op if already there) - latent = manage_tensor_device( + latent = manage_tensor( tensor=latent, target_device=ctx['dit_device'], tensor_name=f"latent_{upscale_idx+1}", @@ -1020,7 +1030,7 @@ def upscale_all_batches( # Offload upscaled latents to avoid VRAM accumulation if ctx['tensor_offload_device'] is not None and (upscaled_latents[0].is_cuda or upscaled_latents[0].is_mps): - ctx['all_upscaled_latents'][upscale_idx] = manage_tensor_device( + ctx['all_upscaled_latents'][upscale_idx] = manage_tensor( tensor=upscaled_latents[0], target_device=ctx['tensor_offload_device'], tensor_name=f"upscaled_latent_{upscale_idx+1}", @@ -1169,7 +1179,7 @@ def decode_all_batches( debug.start_timer(f"decode_batch_{decode_idx+1}") # Move to VAE device with correct dtype for decoding (no-op if already there) - upscaled_latent = manage_tensor_device( + upscaled_latent = manage_tensor( tensor=upscaled_latent, target_device=ctx['vae_device'], tensor_name=f"upscaled_latent_{decode_idx+1}", @@ -1183,22 +1193,33 @@ def decode_all_batches( samples = runner.vae_decode([upscaled_latent]) debug.end_timer("vae_decode", "VAE decode") - # Process samples + # Process samples debug.start_timer("optimized_video_rearrange") samples = optimized_video_rearrange(samples) debug.end_timer("optimized_video_rearrange", "Video rearrange") - # Offload decoded samples to avoid VRAM accumulation + # Convert from VAE dtype to compute dtype and offload to avoid VRAM accumulation if ctx['tensor_offload_device'] is not None: # samples is always a single-element list from vae_decode([upscaled_latent])] if samples[0].is_cuda or samples[0].is_mps: - samples[0] = manage_tensor_device( + samples[0] = manage_tensor( tensor=samples[0], target_device=ctx['tensor_offload_device'], tensor_name=f"sample_{decode_idx+1}", + dtype=ctx['compute_dtype'], debug=debug, - reason="storing decoded latents for post-processing" + reason="storing decoded samples for post-processing (VAE dtype → compute dtype)" ) + else: + # No offload device, but still convert to compute dtype + samples[0] = manage_tensor( + tensor=samples[0], + target_device=samples[0].device, + tensor_name=f"sample_{decode_idx+1}", + dtype=ctx['compute_dtype'], + debug=debug, + reason="VAE dtype → compute dtype" + ) ctx['batch_samples'][decode_idx] = samples # Free the upscaled latent and GPU samples @@ -1322,6 +1343,7 @@ def postprocess_all_batches( alpha_original=ctx['all_alpha_channels'][batch_idx], rgb_original=ctx['all_input_rgb'][batch_idx], device=ctx['vae_device'], + compute_dtype=ctx['compute_dtype'], debug=debug ) @@ -1350,7 +1372,7 @@ def postprocess_all_batches( # Post-process each sample in the batch for i, sample in enumerate(samples): # Move to VAE device with correct dtype for processing (no-op if already there) - sample = manage_tensor_device( + sample = manage_tensor( tensor=sample, target_device=ctx['vae_device'], tensor_name=f"sample_{batch_idx+1}_{i}", @@ -1389,7 +1411,7 @@ def postprocess_all_batches( # Ensure both tensors are on same device (GPU) for color correction if input_video.device != sample.device: - input_video = manage_tensor_device( + input_video = manage_tensor( tensor=input_video, target_device=sample.device, tensor_name=f"input_video_{batch_idx+1}", @@ -1456,15 +1478,13 @@ def postprocess_all_batches( # Move to tensor_offload_device if specified if ctx['tensor_offload_device'] is not None: - if sample.is_cuda or sample.is_mps: - sample = manage_tensor_device( - tensor=sample, - target_device=ctx['tensor_offload_device'], - tensor_name=f"final_sample_{batch_idx+1}", - debug=debug, - reason="storing final assembly" - ) - # If no tensor_offload_device, sample stays on VAE device (no movement needed) + sample = manage_tensor( + tensor=sample, + target_device=ctx['tensor_offload_device'], + tensor_name=f"sample_{batch_idx+1}_{i}_final", + debug=debug, + reason="storing final processed samples" + ) # Get batch dimensions batch_frames = sample.shape[0] diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index 5dc9f69..b1e85f5 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -549,7 +549,7 @@ def release_model_memory(model: Optional[torch.nn.Module], debug: Optional[Any] debug.log(f"Failed to release model memory: {e}", level="WARNING", category="memory", force=True) -def manage_tensor_device( +def manage_tensor( tensor: torch.Tensor, target_device: torch.device, tensor_name: str = "tensor", @@ -559,24 +559,27 @@ def manage_tensor_device( reason: Optional[str] = None ) -> torch.Tensor: """ - Move tensor to target device with consistent logging. + Unified tensor management for device movement and dtype conversion. + + Handles both device transfers (CPU ↔ GPU) and dtype conversions (e.g., float16 ↔ bfloat16) + with intelligent early-exit optimization and comprehensive logging. Args: - tensor: Tensor to move + tensor: Tensor to manage target_device: Target device (torch.device object) tensor_name: Descriptive name for logging (e.g., "latent", "sample", "alpha_channel") dtype: Optional target dtype to cast to (if None, keeps original dtype) non_blocking: Whether to use non-blocking transfer debug: Debug instance for logging - reason: Optional reason for the movement (e.g., "inference", "offload", "color correction") + reason: Optional reason for the operation (e.g., "inference", "offload", "dtype alignment") Returns: Tensor on target device with optional dtype conversion Note: - - Skips movement if tensor already on target device and dtype - - Logs movements consistently with model movements for tracking - - Optimized to avoid unnecessary data transfers + - Skips operation if tensor already has target device and dtype (zero-copy) + - Uses PyTorch's optimized .to() for efficient device/dtype handling + - Logs all operations consistently for tracking and debugging """ if tensor is None: return tensor @@ -617,8 +620,16 @@ def manage_tensor_device( category="general" ) - # Perform the movement - return tensor.to(target_device, dtype=target_dtype, non_blocking=non_blocking) + # Perform the operation based on what needs to change + if needs_device_move and needs_dtype_change: + # Both device and dtype need to change + return tensor.to(target_device, dtype=target_dtype, non_blocking=non_blocking) + elif needs_device_move: + # Only device needs to change + return tensor.to(target_device, non_blocking=non_blocking) + else: + # Only dtype needs to change + return tensor.to(dtype=target_dtype) def manage_model_device(model: torch.nn.Module, target_device: torch.device, model_name: str,