perf(mps): eliminate sync overhead from CPU tensor offload on unified memory
- Skip CPU tensor offload on MPS (no memory benefit, causes sync stall) - Keep input_images and final_video on MPS device - Add explicit MPS sync at phase boundaries for accurate timing - Preload text embeddings before Phase 1 to avoid Phase 2 stall - Skip model→CPU movement before deletion on MPS cleanup
This commit is contained in:
+7
-1
@@ -118,7 +118,9 @@ from src.core.generation_utils import (
|
||||
prepare_runner,
|
||||
compute_generation_info,
|
||||
log_generation_start,
|
||||
blend_overlapping_frames
|
||||
blend_overlapping_frames,
|
||||
load_text_embeddings,
|
||||
script_directory
|
||||
)
|
||||
from src.core.generation_phases import (
|
||||
encode_all_batches,
|
||||
@@ -858,6 +860,10 @@ def _process_frames_core(
|
||||
if runner_cache is not None:
|
||||
runner_cache['runner'] = runner
|
||||
|
||||
# Preload text embeddings before Phase 1 to avoid sync stall in Phase 2
|
||||
ctx['text_embeds'] = load_text_embeddings(script_directory, ctx['dit_device'], ctx['compute_dtype'], debug)
|
||||
debug.log("Loaded text embeddings for DiT", category="dit")
|
||||
|
||||
# Compute generation info and log start (handles prepending internally)
|
||||
frames_tensor, gen_info = compute_generation_info(
|
||||
ctx=ctx,
|
||||
|
||||
@@ -231,7 +231,11 @@ def encode_all_batches(
|
||||
if images is None:
|
||||
raise ValueError("Images to encode must be provided")
|
||||
else:
|
||||
ctx['input_images'] = images
|
||||
# MPS: keep on device to avoid sync overhead in Phase 4 color correction
|
||||
if ctx['vae_device'].type == 'mps' and images.device.type != 'mps':
|
||||
ctx['input_images'] = images.to(ctx['vae_device'])
|
||||
else:
|
||||
ctx['input_images'] = images
|
||||
|
||||
# Get total frame count from context (set in video_upscaler before encoding)
|
||||
total_frames = ctx.get('total_frames', len(images))
|
||||
@@ -529,6 +533,10 @@ def encode_all_batches(
|
||||
manage_model_device(model=runner.vae, target_device=ctx['vae_offload_device'],
|
||||
model_name="VAE", debug=debug, reason="VAE offload", runner=runner)
|
||||
|
||||
# MPS: sync to get accurate timing and free memory before Phase 2
|
||||
if ctx['vae_device'].type == 'mps':
|
||||
torch.mps.synchronize()
|
||||
|
||||
debug.end_timer("phase1_encoding", "Phase 1: VAE encoding complete", show_breakdown=True)
|
||||
debug.log_memory_state("After phase 1 (VAE encoding)", show_tensors=False)
|
||||
|
||||
@@ -860,7 +868,13 @@ def decode_all_batches(
|
||||
|
||||
# Pre-allocate final_video at the START of decode phase (before any batch processing)
|
||||
# This ensures we only need memory for final_video + 1 batch, not final_video + all batch_samples
|
||||
target_device = ctx['tensor_offload_device'] if ctx['tensor_offload_device'] is not None else 'cpu'
|
||||
# MPS: keep on device (unified memory, no benefit to CPU offload)
|
||||
if ctx['tensor_offload_device'] is not None:
|
||||
target_device = ctx['tensor_offload_device']
|
||||
elif ctx['vae_device'].type == 'mps':
|
||||
target_device = ctx['vae_device']
|
||||
else:
|
||||
target_device = 'cpu'
|
||||
channels_str = "RGBA" if C == 4 else "RGB"
|
||||
required_gb = (total_frames * true_h * true_w * C * 2) / (1024**3)
|
||||
debug.log(f"Pre-allocating output tensor: {total_frames} frames, {true_w}x{true_h}px, {channels_str} ({required_gb:.2f}GB)",
|
||||
@@ -1040,6 +1054,10 @@ def decode_all_batches(
|
||||
if 'all_upscaled_latents' in ctx:
|
||||
release_tensor_collection(ctx['all_upscaled_latents'])
|
||||
del ctx['all_upscaled_latents']
|
||||
|
||||
# MPS: sync to get accurate timing and free memory before Phase 4
|
||||
if ctx['vae_device'].type == 'mps':
|
||||
torch.mps.synchronize()
|
||||
|
||||
debug.end_timer("phase3_decoding", "Phase 3: VAE decoding complete", show_breakdown=True)
|
||||
debug.log_memory_state("After phase 3 (VAE decoding)", show_tensors=False)
|
||||
|
||||
@@ -350,7 +350,12 @@ def setup_generation_context(
|
||||
vae_device = _normalize_device(vae_device)
|
||||
dit_offload_device = _normalize_device(dit_offload_device) if dit_offload_device is not None else None
|
||||
vae_offload_device = _normalize_device(vae_offload_device) if vae_offload_device is not None else None
|
||||
tensor_offload_device = _normalize_device(tensor_offload_device) if tensor_offload_device is not None else None
|
||||
# MPS unified memory: CPU offload causes sync overhead with no memory benefit
|
||||
is_mps = dit_device.type == 'mps' or vae_device.type == 'mps'
|
||||
if is_mps and tensor_offload_device is not None and str(tensor_offload_device) == 'cpu':
|
||||
tensor_offload_device = None
|
||||
else:
|
||||
tensor_offload_device = _normalize_device(tensor_offload_device) if tensor_offload_device is not None else None
|
||||
|
||||
# Set LOCAL_RANK to 0 for single-GPU inference mode
|
||||
# CLI multi-GPU uses CUDA_VISIBLE_DEVICES to restrict visibility per worker
|
||||
|
||||
@@ -19,7 +19,9 @@ from ..core.generation_utils import (
|
||||
setup_generation_context,
|
||||
prepare_runner,
|
||||
compute_generation_info,
|
||||
log_generation_start
|
||||
log_generation_start,
|
||||
load_text_embeddings,
|
||||
script_directory
|
||||
)
|
||||
from ..optimization.memory_manager import (
|
||||
cleanup_text_embeddings,
|
||||
@@ -437,6 +439,10 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
|
||||
# Store cache context in ctx for use in generation phases
|
||||
ctx['cache_context'] = cache_context
|
||||
|
||||
# Preload text embeddings before Phase 1 to avoid sync stall in Phase 2
|
||||
ctx['text_embeds'] = load_text_embeddings(script_directory, ctx['dit_device'], ctx['compute_dtype'], debug)
|
||||
debug.log("Loaded text embeddings for DiT", category="dit")
|
||||
|
||||
debug.log_memory_state("After model preparation", show_tensors=False, detailed_tensors=False)
|
||||
debug.end_timer("model_preparation", "Model preparation", force=True, show_breakdown=True)
|
||||
|
||||
|
||||
@@ -1050,15 +1050,17 @@ def cleanup_dit(runner: Any, debug: Optional['Debug'] = None, cache_model: bool
|
||||
|
||||
# Move model off GPU if needed
|
||||
if param_device.type not in ['meta', 'cpu']:
|
||||
# Get offload target - default to 'cpu' if not configured or set to 'none'
|
||||
offload_target = getattr(runner, '_dit_offload_device', None)
|
||||
if offload_target is None or offload_target == 'none':
|
||||
offload_target = torch.device('cpu')
|
||||
|
||||
# Move model off GPU (either for caching or before deletion)
|
||||
reason = "model caching" if cache_model else "releasing GPU memory"
|
||||
manage_model_device(model=runner.dit, target_device=offload_target, model_name="DiT",
|
||||
debug=debug, reason=reason, runner=runner)
|
||||
# MPS: skip CPU movement before deletion (unified memory, just causes sync)
|
||||
if param_device.type == 'mps' and not cache_model:
|
||||
if debug:
|
||||
debug.log("DiT on MPS - skipping CPU movement before deletion", category="cleanup")
|
||||
else:
|
||||
offload_target = getattr(runner, '_dit_offload_device', None)
|
||||
if offload_target is None or offload_target == 'none':
|
||||
offload_target = torch.device('cpu')
|
||||
reason = "model caching" if cache_model else "releasing GPU memory"
|
||||
manage_model_device(model=runner.dit, target_device=offload_target, model_name="DiT",
|
||||
debug=debug, reason=reason, runner=runner)
|
||||
elif param_device.type == 'meta' and debug:
|
||||
debug.log("DiT on meta device - keeping structure for cache", category="cleanup")
|
||||
except StopIteration:
|
||||
@@ -1126,15 +1128,17 @@ def cleanup_vae(runner: Any, debug: Optional['Debug'] = None, cache_model: bool
|
||||
|
||||
# Move model off GPU if needed
|
||||
if param_device.type not in ['meta', 'cpu']:
|
||||
# Get offload target - default to 'cpu' if not configured or set to 'none'
|
||||
offload_target = getattr(runner, '_vae_offload_device', None)
|
||||
if offload_target is None or offload_target == 'none':
|
||||
offload_target = torch.device('cpu')
|
||||
|
||||
# Move model off GPU (either for caching or before deletion)
|
||||
reason = "model caching" if cache_model else "releasing GPU memory"
|
||||
manage_model_device(model=runner.vae, target_device=offload_target, model_name="VAE",
|
||||
debug=debug, reason=reason, runner=runner)
|
||||
# MPS: skip CPU movement before deletion (unified memory, just causes sync)
|
||||
if param_device.type == 'mps' and not cache_model:
|
||||
if debug:
|
||||
debug.log("VAE on MPS - skipping CPU movement before deletion", category="cleanup")
|
||||
else:
|
||||
offload_target = getattr(runner, '_vae_offload_device', None)
|
||||
if offload_target is None or offload_target == 'none':
|
||||
offload_target = torch.device('cpu')
|
||||
reason = "model caching" if cache_model else "releasing GPU memory"
|
||||
manage_model_device(model=runner.vae, target_device=offload_target, model_name="VAE",
|
||||
debug=debug, reason=reason, runner=runner)
|
||||
elif param_device.type == 'meta' and debug:
|
||||
debug.log("VAE on meta device - keeping structure for cache", category="cleanup")
|
||||
except StopIteration:
|
||||
|
||||
Reference in New Issue
Block a user