Merge pull request #103 from lihaoyun6/main
Added MPS backend support (for running on macOS)
This commit is contained in:
+2
-1
@@ -22,4 +22,5 @@ src/core/isolated_generation.py
|
||||
src/core/subprocess_runner.py
|
||||
models/video_vae_v3_mine_bad/
|
||||
src/processing/
|
||||
TILE_VAE*
|
||||
TILE_VAE*
|
||||
.DS_Store
|
||||
+37
-23
@@ -7,6 +7,7 @@ import sys
|
||||
import os
|
||||
import argparse
|
||||
import time
|
||||
import platform
|
||||
import multiprocessing as mp
|
||||
|
||||
# Set up path before any other imports to fix module resolution
|
||||
@@ -22,17 +23,18 @@ if mp.get_start_method(allow_none=True) != 'spawn':
|
||||
mp.set_start_method('spawn', force=True)
|
||||
# -------------------------------------------------------------
|
||||
# 1) Gestion VRAM (cudaMallocAsync) déjà en place
|
||||
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
|
||||
if platform.system() != "Darwin":
|
||||
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
|
||||
|
||||
# 2) Pré-parse de la ligne de commande pour récupérer --cuda_device
|
||||
_pre_parser = argparse.ArgumentParser(add_help=False)
|
||||
_pre_parser.add_argument("--cuda_device", type=str, default=None)
|
||||
_pre_args, _ = _pre_parser.parse_known_args()
|
||||
if _pre_args.cuda_device is not None:
|
||||
device_list_env = [x.strip() for x in _pre_args.cuda_device.split(',') if x.strip()!='']
|
||||
if len(device_list_env) == 1:
|
||||
# Single GPU: restrict visibility now
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = device_list_env[0]
|
||||
# 2) Pré-parse de la ligne de commande pour récupérer --cuda_device
|
||||
_pre_parser = argparse.ArgumentParser(add_help=False)
|
||||
_pre_parser.add_argument("--cuda_device", type=str, default=None)
|
||||
_pre_args, _ = _pre_parser.parse_known_args()
|
||||
if _pre_args.cuda_device is not None:
|
||||
device_list_env = [x.strip() for x in _pre_args.cuda_device.split(',') if x.strip()!='']
|
||||
if len(device_list_env) == 1:
|
||||
# Single GPU: restrict visibility now
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = device_list_env[0]
|
||||
|
||||
# -------------------------------------------------------------
|
||||
# 3) Imports lourds (torch, etc.) après la configuration env
|
||||
@@ -273,10 +275,11 @@ def apply_temporal_overlap_blending(frames_tensor, batch_size, overlap):
|
||||
|
||||
def _worker_process(proc_idx, device_id, frames_np, shared_args, return_queue):
|
||||
"""Worker process that performs upscaling on a slice of frames using a dedicated GPU."""
|
||||
# 1. Limit CUDA visibility to the chosen GPU BEFORE importing torch-heavy deps
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(device_id)
|
||||
# Keep same cudaMallocAsync setting
|
||||
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
|
||||
if platform.system() != "Darwin":
|
||||
# 1. Limit CUDA visibility to the chosen GPU BEFORE importing torch-heavy deps
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(device_id)
|
||||
# Keep same cudaMallocAsync setting
|
||||
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
|
||||
|
||||
import torch # local import inside subprocess
|
||||
from src.core.model_manager import configure_runner
|
||||
@@ -471,7 +474,8 @@ def parse_arguments():
|
||||
help="Enable VRAM preservation mode")
|
||||
parser.add_argument("--debug", action="store_true",
|
||||
help="Enable debug logging")
|
||||
parser.add_argument("--cuda_device", type=str, default=None,
|
||||
if platform.system() != "Darwin":
|
||||
parser.add_argument("--cuda_device", type=str, default=None,
|
||||
help="CUDA device id(s). Single id (e.g., '0') or comma-separated list '0,1' for multi-GPU")
|
||||
parser.add_argument("--blocks_to_swap", type=int, default=0,
|
||||
help="Number of blocks to swap for VRAM optimization (default: 0, disabled), up to 32 for 3B model, 36 for 7B")
|
||||
@@ -508,11 +512,15 @@ def main():
|
||||
print(f"Error: VAE tile overlap {args.vae_tile_overlap} must be smaller than tile size {args.vae_tile_size}")
|
||||
sys.exit(1)
|
||||
|
||||
# Show actual CUDA device visibility
|
||||
debug.log(f"CUDA_VISIBLE_DEVICES: {os.environ.get('CUDA_VISIBLE_DEVICES', 'Not set (all)')}", category="device")
|
||||
if torch.cuda.is_available():
|
||||
debug.log(f"torch.cuda.device_count(): {torch.cuda.device_count()}", category="device")
|
||||
debug.log(f"Using device index 0 inside script (mapped to selected GPU)", category="device")
|
||||
if args.debug:
|
||||
if platform.system() == "Darwin":
|
||||
print("You are running on macOS and will use the MPS backend!")
|
||||
else:
|
||||
# Show actual CUDA device visibility
|
||||
debug.log(f"CUDA_VISIBLE_DEVICES: {os.environ.get('CUDA_VISIBLE_DEVICES', 'Not set (all)')}", category="device")
|
||||
if torch.cuda.is_available():
|
||||
debug.log(f"torch.cuda.device_count(): {torch.cuda.device_count()}", category="device")
|
||||
debug.log(f"Using device index 0 inside script (mapped to selected GPU)", category="device")
|
||||
|
||||
try:
|
||||
# Ensure --output is a directory when using PNG format
|
||||
@@ -537,15 +545,21 @@ def main():
|
||||
# debug.log(f"Initial VRAM: {torch.cuda.memory_allocated() / 1024**3:.2f}GB", category="memory")
|
||||
|
||||
# Parse GPU list
|
||||
device_list = [d.strip() for d in str(args.cuda_device).split(',') if d.strip()] if args.cuda_device else ["0"]
|
||||
debug.log(f"Using devices: {device_list}", category="device")
|
||||
if platform.system() == "Darwin":
|
||||
device_list = ["0"]
|
||||
else:
|
||||
device_list = [d.strip() for d in str(args.cuda_device).split(',') if d.strip()] if args.cuda_device else ["0"]
|
||||
|
||||
if args.debug:
|
||||
debug.log(f"Using devices: {device_list}", category="device")
|
||||
processing_start = time.time()
|
||||
download_weight(args.model, args.model_dir)
|
||||
result = _gpu_processing(frames_tensor, device_list, args)
|
||||
generation_time = time.time() - processing_start
|
||||
|
||||
debug.log(f"Generation time: {generation_time:.2f}s", category="general")
|
||||
debug.log(f"Peak VRAM usage: {torch.cuda.max_memory_allocated() / 1024**3:.2f}GB", category="memory")
|
||||
if platform.system() != "Darwin":
|
||||
debug.log(f"Peak VRAM usage: {torch.cuda.max_memory_allocated() / 1024**3:.2f}GB", category="memory")
|
||||
|
||||
if args.temporal_overlap > 0:
|
||||
debug.log(f"Applying temporal overlap with blending", category="generation")
|
||||
|
||||
@@ -21,6 +21,7 @@ from typing import Callable
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch.nn import functional as F
|
||||
import platform
|
||||
|
||||
#from ....models.dit_v2 import na
|
||||
|
||||
@@ -71,9 +72,13 @@ class EulerSampler(Sampler):
|
||||
|
||||
# Nettoyer les tenseurs temporaires
|
||||
del pred
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
if torch.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
i += 1
|
||||
progress.update()
|
||||
@@ -126,3 +131,4 @@ class EulerSampler(Sampler):
|
||||
pred_x_s = pred_x_s.where(s >= 0, pred_x_0)
|
||||
pred_x_s = pred_x_s.where(s <= T, pred_x_T)
|
||||
return pred_x_s
|
||||
|
||||
@@ -21,7 +21,7 @@ from datetime import timedelta
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
|
||||
import platform
|
||||
|
||||
def get_global_rank() -> int:
|
||||
"""
|
||||
@@ -48,7 +48,10 @@ def get_device() -> torch.device:
|
||||
"""
|
||||
Get current rank device.
|
||||
"""
|
||||
return torch.device("cuda", get_local_rank())
|
||||
device = "cuda"
|
||||
if platform.system() == "Darwin":
|
||||
device = "mps"
|
||||
return torch.device(device, get_local_rank())
|
||||
|
||||
|
||||
def barrier_if_distributed(*args, **kwargs):
|
||||
@@ -82,3 +85,4 @@ def convert_to_ddp(module: torch.nn.Module, **kwargs) -> DistributedDataParallel
|
||||
output_device=get_local_rank(),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
+25
-7
@@ -20,6 +20,7 @@ import os
|
||||
import gc
|
||||
import torch
|
||||
import time
|
||||
import platform
|
||||
from src.utils.constants import get_script_directory
|
||||
from torchvision.transforms import Compose, Lambda, Normalize
|
||||
|
||||
@@ -75,6 +76,8 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
|
||||
raise ValueError("Debug instance must be provided to generation_step")
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if platform.system() == "Darwin":
|
||||
device = "mps"
|
||||
|
||||
# Adaptive dtype detection for optimal performance
|
||||
model_dtype = next(runner.dit.parameters()).dtype
|
||||
@@ -95,12 +98,17 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
|
||||
def _move_to_cuda(x):
|
||||
"""Move tensors to CUDA with adaptive optimal dtype"""
|
||||
return [i.to(device, dtype=dtype) for i in x]
|
||||
|
||||
|
||||
# Memory optimization: Generate noise once and reuse to save VRAM
|
||||
with torch.cuda.device(device):
|
||||
if platform.system() == "Darwin":
|
||||
base_noise = torch.randn_like(cond_latents[0], dtype=dtype)
|
||||
noises = [base_noise]
|
||||
aug_noises = [base_noise * 0.1 + torch.randn_like(base_noise) * 0.05]
|
||||
else:
|
||||
with torch.cuda.device(device):
|
||||
base_noise = torch.randn_like(cond_latents[0], dtype=dtype)
|
||||
noises = [base_noise]
|
||||
aug_noises = [base_noise * 0.1 + torch.randn_like(base_noise) * 0.05]
|
||||
|
||||
# Move tensors with adaptive dtype (optimized for FP8/FP16/BFloat16)
|
||||
noises, aug_noises, cond_latents = _move_to_cuda(noises), _move_to_cuda(aug_noises), _move_to_cuda(cond_latents)
|
||||
@@ -131,7 +139,11 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
|
||||
|
||||
# Use adaptive autocast for optimal performance
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", autocast_dtype, enabled=True):
|
||||
d = "cuda"
|
||||
if platform.system() == "Darwin":
|
||||
d = "mps"
|
||||
|
||||
with torch.autocast(d, autocast_dtype, enabled=True):
|
||||
video_tensors = runner.inference(
|
||||
noises=noises,
|
||||
conditions=conditions,
|
||||
@@ -219,6 +231,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
raise ValueError("Debug instance must be provided to generation_loop")
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if platform.system() == "Darwin":
|
||||
device = "mps"
|
||||
|
||||
# ───────────────────────────────────────────────────────────────
|
||||
# Step 1: Model Configuration & Precision Detection
|
||||
@@ -256,7 +270,7 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
vae_dtype = torch.bfloat16
|
||||
|
||||
# Optimization tips for users
|
||||
if torch.cuda.is_available():
|
||||
if torch.cuda.is_available() or torch.mps.is_available():
|
||||
total_frames = len(images)
|
||||
optimal_batches = [x for x in [i for i in range(1, 200) if i % 4 == 1] if x <= total_frames]
|
||||
if optimal_batches:
|
||||
@@ -395,7 +409,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
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("cuda", autocast_dtype, enabled=True):
|
||||
_device = "mps" if platform.system() == "Darwin" else "cuda"
|
||||
with torch.autocast(_device, autocast_dtype, enabled=True):
|
||||
cond_latents = runner.vae_encode([transformed_video])
|
||||
debug.end_timer("vae_encode", "VAE encoding")
|
||||
#tps = time.time()
|
||||
@@ -464,8 +479,11 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
|
||||
text_neg_embeds = text_neg_embeds.to("cpu")
|
||||
runner.dit.to("cpu")
|
||||
runner.vae.to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
#del text_pos_embeds, text_neg_embeds
|
||||
#clear_vram_cache()
|
||||
|
||||
|
||||
+12
-2
@@ -19,6 +19,7 @@ from einops import rearrange
|
||||
from omegaconf import DictConfig, ListConfig
|
||||
from torch import Tensor
|
||||
from src.optimization.memory_manager import clear_vram_cache
|
||||
import platform
|
||||
|
||||
from src.common.diffusion import (
|
||||
classifier_free_guidance_dispatcher,
|
||||
@@ -268,6 +269,11 @@ class VideoDiffusionInfer():
|
||||
|
||||
def get_vram_usage(self):
|
||||
"""Obtenir l'utilisation VRAM actuelle (allouée et réservée)"""
|
||||
if platform.system() == "Darwin":
|
||||
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)
|
||||
@@ -374,7 +380,11 @@ class VideoDiffusionInfer():
|
||||
|
||||
self.debug.start_timer("dit_inference")
|
||||
|
||||
with torch.autocast("cuda", target_dtype, enabled=True):
|
||||
d = "cuda"
|
||||
if platform.system() == "Darwin":
|
||||
d = "mps"
|
||||
|
||||
with torch.autocast(d, target_dtype, enabled=True):
|
||||
latents = self.sampler.sample(
|
||||
x=latents,
|
||||
f=lambda args: classifier_free_guidance_dispatcher(
|
||||
@@ -458,4 +468,4 @@ class VideoDiffusionInfer():
|
||||
#self.debug.log(f"🔄 FINAL CLEANUP time: {time.time() - t} seconds", category="timing")
|
||||
|
||||
|
||||
return samples
|
||||
return samples
|
||||
|
||||
@@ -18,6 +18,7 @@ Key Features:
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
import platform
|
||||
from src.utils.constants import get_script_directory
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
|
||||
@@ -178,6 +179,8 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=None,
|
||||
|
||||
# Set device
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if platform.system() == "Darwin":
|
||||
device = "mps"
|
||||
|
||||
# Configure models
|
||||
checkpoint_path = os.path.join(base_cache_dir, f'./{model}')
|
||||
@@ -258,7 +261,7 @@ def load_quantized_state_dict(checkpoint_path, device="cpu", keep_native_fp8=Tru
|
||||
if hasattr(tensor, 'dtype') and tensor.dtype in fp8_types:
|
||||
fp8_detected = True
|
||||
break
|
||||
|
||||
|
||||
if fp8_detected:
|
||||
if keep_native_fp8:
|
||||
# Keep native FP8 format for optimal performance
|
||||
@@ -393,6 +396,10 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config,
|
||||
raise ValueError("Debug instance must be provided to configure_vae_model_inference")
|
||||
|
||||
# Create vae model
|
||||
if platform.system() == "Darwin":
|
||||
config.vae.dtype = "float16"
|
||||
if "fp8_e4m3fn" in runner._model_name:
|
||||
config.vae.dtype = "bfloat16"
|
||||
|
||||
dtype = getattr(torch, config.vae.dtype)
|
||||
debug.start_timer("vae_model_create")
|
||||
@@ -444,6 +451,10 @@ 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 platform.system() == "Darwin":
|
||||
runner.vae = runner.vae.to(dtype=getattr(torch, config.vae.dtype))
|
||||
|
||||
if state_loading_device == "cpu":
|
||||
runner.vae.to(device)
|
||||
if 'state' in locals():
|
||||
|
||||
@@ -31,6 +31,8 @@ class AreaResize:
|
||||
self.max_area = max_area
|
||||
self.downsample_only = downsample_only
|
||||
self.interpolation = interpolation
|
||||
if platform.system() == "Darwin":
|
||||
self.interpolation = InterpolationMode.BILINEAR
|
||||
|
||||
def __call__(self, image: Union[torch.Tensor, Image.Image]):
|
||||
|
||||
@@ -133,3 +135,4 @@ class ScaleResize:
|
||||
antialias=antialias,
|
||||
)
|
||||
return image
|
||||
|
||||
@@ -17,7 +17,7 @@ from torchvision.transforms import CenterCrop, Compose, InterpolationMode, Resiz
|
||||
|
||||
from .area_resize import AreaResize
|
||||
from .side_resize import SideResize
|
||||
|
||||
import platform
|
||||
|
||||
def NaResize(
|
||||
resolution: int,
|
||||
@@ -25,24 +25,25 @@ def NaResize(
|
||||
downsample_only: bool,
|
||||
interpolation: InterpolationMode = InterpolationMode.BICUBIC,
|
||||
):
|
||||
Interpolation = InterpolationMode.BILINEAR if platform.system() == "Darwin" else interpolation
|
||||
if mode == "area":
|
||||
return AreaResize(
|
||||
max_area=resolution**2,
|
||||
downsample_only=downsample_only,
|
||||
interpolation=interpolation,
|
||||
interpolation=Interpolation,
|
||||
)
|
||||
if mode == "side":
|
||||
return SideResize(
|
||||
size=resolution,
|
||||
downsample_only=downsample_only,
|
||||
interpolation=interpolation,
|
||||
interpolation=Interpolation,
|
||||
)
|
||||
if mode == "square":
|
||||
return Compose(
|
||||
[
|
||||
Resize(
|
||||
size=resolution,
|
||||
interpolation=interpolation,
|
||||
interpolation=Interpolation,
|
||||
),
|
||||
CenterCrop(resolution),
|
||||
]
|
||||
|
||||
@@ -17,7 +17,7 @@ import torch
|
||||
from PIL import Image
|
||||
from torchvision.transforms import InterpolationMode
|
||||
from torchvision.transforms import functional as TVF
|
||||
|
||||
import platform
|
||||
|
||||
class SideResize:
|
||||
def __init__(
|
||||
@@ -29,6 +29,8 @@ class SideResize:
|
||||
self.size = size
|
||||
self.downsample_only = downsample_only
|
||||
self.interpolation = interpolation
|
||||
if platform.system() == "Darwin":
|
||||
self.interpolation = InterpolationMode.BILINEAR
|
||||
|
||||
def __call__(self, image: Union[torch.Tensor, Image.Image]):
|
||||
"""
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
import platform
|
||||
from typing import Tuple, Dict, Any
|
||||
|
||||
from src.utils.constants import get_base_cache_dir
|
||||
@@ -196,8 +197,6 @@ class SeedVR2:
|
||||
# Clean BlockSwap with state preservation
|
||||
if hasattr(self.runner, "_blockswap_active") and self.runner._blockswap_active:
|
||||
cleanup_blockswap(self.runner, keep_state_for_cache=True)
|
||||
|
||||
# Clear caches but keep models
|
||||
if self.runner:
|
||||
clear_all_caches(self.runner, debug, offload_vae=True)
|
||||
|
||||
@@ -421,7 +420,7 @@ class SeedVR2BlockSwap:
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Use non-blocking GPU transfers for better performance.",
|
||||
"tooltip": "Use non-blocking GPU transfers for better performance.\n(This will always False on macOS to prevent Nan tensors)",
|
||||
},
|
||||
),
|
||||
"offload_io_components": (
|
||||
@@ -464,10 +463,10 @@ The actual memory savings depend on your specific model architecture and will be
|
||||
"""Create BlockSwap configuration"""
|
||||
if blocks_to_swap == 0:
|
||||
return (None,)
|
||||
|
||||
_use_non_blocking = False if platform.system() == "Darwin" else use_non_blocking
|
||||
config = {
|
||||
"blocks_to_swap": blocks_to_swap,
|
||||
"use_non_blocking": use_non_blocking,
|
||||
"use_non_blocking": _use_non_blocking,
|
||||
"offload_io_components": offload_io_components,
|
||||
}
|
||||
|
||||
|
||||
@@ -29,6 +29,7 @@ from diffusers.utils import is_torch_version
|
||||
from diffusers.utils.accelerate_utils import apply_forward_hook
|
||||
from einops import rearrange
|
||||
from ....common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation
|
||||
import platform
|
||||
|
||||
from ....common.distributed.advanced import get_sequence_parallel_world_size
|
||||
from ....common.logger import get_logger
|
||||
@@ -133,8 +134,11 @@ class Upsample3D(Upsample2D):
|
||||
hidden_states = [hidden_states]
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
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(
|
||||
@@ -153,8 +157,11 @@ class Upsample3D(Upsample2D):
|
||||
hidden_states = hidden_states[0]
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
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)
|
||||
@@ -313,8 +320,11 @@ class ResnetBlock3D(ResnetBlock2D):
|
||||
except Exception as e:
|
||||
if hasattr(self, 'debug') and self.debug:
|
||||
self.debug.log("OOM second chance: ResnetBlock3D", category="warning", force=True)
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(1)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
@@ -1612,3 +1622,4 @@ class VideoAutoencoderKLWrapper(VideoAutoencoderKL):
|
||||
for m in self.modules():
|
||||
if isinstance(m, InflatedCausalConv3d):
|
||||
m.set_memory_limit(conv_max_mem if conv_max_mem is not None else float("inf"))
|
||||
|
||||
@@ -22,6 +22,7 @@ from diffusers.models.normalization import RMSNorm
|
||||
from einops import rearrange
|
||||
from torch import Tensor, nn
|
||||
from torch.nn import Conv3d
|
||||
import platform
|
||||
|
||||
from .context_parallel_lib import cache_send_recv, get_cache_size
|
||||
from .global_config import get_norm_limit
|
||||
@@ -119,8 +120,11 @@ class InflatedCausalConv3d(Conv3d):
|
||||
if prev_cache is not None:
|
||||
prev_cache = list(prev_cache.split(split_sizes, dim=split_dim))
|
||||
if preserve_vram:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
# Loop Fwd.
|
||||
cache = None
|
||||
for idx in range(len(x)):
|
||||
@@ -166,8 +170,11 @@ class InflatedCausalConv3d(Conv3d):
|
||||
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
#print("empty cache 1")
|
||||
#time.sleep(2)
|
||||
try:
|
||||
@@ -175,8 +182,11 @@ class InflatedCausalConv3d(Conv3d):
|
||||
except Exception as e:
|
||||
if hasattr(self, 'debug') and self.debug:
|
||||
self.debug.log("OOM Second Chance", category="warning", force=True)
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(2)
|
||||
output = torch.cat(x, split_dim)
|
||||
return output
|
||||
@@ -359,23 +369,32 @@ def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor, preserve_vram: b
|
||||
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)
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(2)
|
||||
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:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
# ADD BY NUMZ
|
||||
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)
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
time.sleep(2)
|
||||
x = torch.cat(x, dim=1)
|
||||
else:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,480 @@
|
||||
# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
|
||||
# //
|
||||
# // Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# // you may not use this file except in compliance with the License.
|
||||
# // You may obtain a copy of the License at
|
||||
# //
|
||||
# // http://www.apache.org/licenses/LICENSE-2.0
|
||||
# //
|
||||
# // Unless required by applicable law or agreed to in writing, software
|
||||
# // distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# // See the License for the specific language governing permissions and
|
||||
# // limitations under the License.
|
||||
|
||||
import math
|
||||
from contextlib import contextmanager
|
||||
import time
|
||||
from typing import List, Optional, Union
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models.normalization import RMSNorm
|
||||
from einops import rearrange
|
||||
from torch import Tensor, nn
|
||||
from torch.nn import Conv3d
|
||||
import platform
|
||||
|
||||
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
|
||||
|
||||
# Single GPU inference - no distributed processing needed
|
||||
#print("Warning: Using single GPU inference mode - distributed features disabled in causal_inflation_lib")
|
||||
|
||||
# Mock distributed functions for single GPU inference
|
||||
def get_sequence_parallel_group():
|
||||
return None
|
||||
|
||||
def get_sequence_parallel_rank():
|
||||
return 0
|
||||
|
||||
def get_sequence_parallel_world_size():
|
||||
return 1
|
||||
|
||||
def get_next_sequence_parallel_rank():
|
||||
return 0
|
||||
|
||||
def get_prev_sequence_parallel_rank():
|
||||
return 0
|
||||
|
||||
|
||||
@contextmanager
|
||||
def ignore_padding(model):
|
||||
orig_padding = model.padding
|
||||
model.padding = (0, 0, 0)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
model.padding = orig_padding
|
||||
|
||||
|
||||
class InflatedCausalConv3d(Conv3d):
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
inflation_mode: _inflation_mode_t,
|
||||
memory_device: _memory_device_t = "same",
|
||||
**kwargs,
|
||||
):
|
||||
self.inflation_mode = inflation_mode
|
||||
self.memory = None
|
||||
super().__init__(*args, **kwargs)
|
||||
self.temporal_padding = self.padding[0]
|
||||
self.memory_device = memory_device
|
||||
self.padding = (0, *self.padding[1:]) # Remove temporal pad to keep causal.
|
||||
self.memory_limit = float("inf")
|
||||
|
||||
def set_memory_limit(self, value: float):
|
||||
self.memory_limit = value
|
||||
|
||||
def set_memory_device(self, memory_device: _memory_device_t):
|
||||
self.memory_device = memory_device
|
||||
|
||||
def memory_limit_conv(
|
||||
self,
|
||||
x,
|
||||
*,
|
||||
split_dim=3,
|
||||
padding=(0, 0, 0, 0, 0, 0),
|
||||
prev_cache=None,
|
||||
preserve_vram = False,
|
||||
):
|
||||
# Compatible with no limit.
|
||||
if math.isinf(self.memory_limit):
|
||||
if prev_cache is not None:
|
||||
x = torch.cat([prev_cache, x], dim=split_dim - 1)
|
||||
return super().forward(x)
|
||||
|
||||
# Compute tensor shape after concat & padding.
|
||||
shape = torch.tensor(x.size())
|
||||
if prev_cache is not None:
|
||||
shape[split_dim - 1] += prev_cache.size(split_dim - 1)
|
||||
shape[-3:] += torch.tensor(padding).view(3, 2).sum(-1).flip(0)
|
||||
memory_occupy = shape.prod() * x.element_size() / 1024**3 # GiB
|
||||
if memory_occupy < self.memory_limit or split_dim == x.ndim:
|
||||
if prev_cache is not None:
|
||||
x = torch.cat([prev_cache, x], dim=split_dim - 1)
|
||||
x = safe_pad_operation(x, padding, mode='constant', value=0.0)
|
||||
with ignore_padding(self):
|
||||
return super().forward(x)
|
||||
|
||||
# Exceed memory limit, splitting tensor
|
||||
|
||||
# Split input (& prev_cache).
|
||||
num_splits = math.ceil(memory_occupy / self.memory_limit)
|
||||
size_per_split = x.size(split_dim) // num_splits
|
||||
split_sizes = [size_per_split] * (num_splits - 1)
|
||||
split_sizes += [x.size(split_dim) - sum(split_sizes)]
|
||||
|
||||
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 platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
#print("empty cache 0")
|
||||
# Loop Fwd.
|
||||
cache = None
|
||||
for idx in range(len(x)):
|
||||
# Concat prev cache from last dim
|
||||
if prev_cache is not None:
|
||||
x[idx] = torch.cat([prev_cache[idx], x[idx]], dim=split_dim - 1)
|
||||
|
||||
# Get padding pattern.
|
||||
lpad_dim = (x[idx].ndim - split_dim - 1) * 2
|
||||
rpad_dim = lpad_dim + 1
|
||||
padding = list(padding)
|
||||
padding[lpad_dim] = self.padding[split_dim - 2] if idx == 0 else 0
|
||||
padding[rpad_dim] = self.padding[split_dim - 2] if idx == len(x) - 1 else 0
|
||||
pad_len = padding[lpad_dim] + padding[rpad_dim]
|
||||
padding = tuple(padding)
|
||||
|
||||
# Prepare cache for next slice (this dim).
|
||||
next_cache = None
|
||||
cache_len = cache.size(split_dim) if cache is not None else 0
|
||||
next_catch_size = get_cache_size(
|
||||
conv_module=self,
|
||||
input_len=x[idx].size(split_dim) + cache_len,
|
||||
pad_len=pad_len,
|
||||
dim=split_dim - 2,
|
||||
)
|
||||
if next_catch_size != 0:
|
||||
assert next_catch_size <= x[idx].size(split_dim)
|
||||
next_cache = (
|
||||
x[idx].transpose(0, split_dim)[-next_catch_size:].transpose(0, split_dim)
|
||||
)
|
||||
|
||||
# Recursive.
|
||||
x[idx] = self.memory_limit_conv(
|
||||
x[idx],
|
||||
split_dim=split_dim + 1,
|
||||
padding=padding,
|
||||
prev_cache=cache,
|
||||
preserve_vram=preserve_vram
|
||||
)
|
||||
|
||||
# Update cache.
|
||||
cache = next_cache
|
||||
# ADD BY NUMZ
|
||||
if preserve_vram:
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
#print("empty cache 1")
|
||||
#time.sleep(2)
|
||||
try:
|
||||
output = torch.cat(x, split_dim)
|
||||
except Exception as e:
|
||||
print("OOM second chance")
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
time.sleep(2)
|
||||
output = torch.cat(x, split_dim)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input: Union[Tensor, List[Tensor]],
|
||||
memory_state: MemoryState = MemoryState.UNSET,
|
||||
preserve_vram: bool = False,
|
||||
) -> Tensor:
|
||||
assert memory_state != MemoryState.UNSET
|
||||
if memory_state != MemoryState.ACTIVE:
|
||||
self.memory = None
|
||||
if (
|
||||
math.isinf(self.memory_limit)
|
||||
and torch.is_tensor(input)
|
||||
and get_sequence_parallel_group() is None
|
||||
):
|
||||
return self.basic_forward(input, memory_state)
|
||||
return self.slicing_forward(input, memory_state, preserve_vram)
|
||||
|
||||
def basic_forward(self, input: Tensor, memory_state: MemoryState = MemoryState.UNSET):
|
||||
mem_size = self.stride[0] - self.kernel_size[0]
|
||||
if (self.memory is not None) and (memory_state == MemoryState.ACTIVE):
|
||||
input = extend_head(input, memory=self.memory, times=-1)
|
||||
else:
|
||||
input = extend_head(input, times=self.temporal_padding * 2)
|
||||
memory = (
|
||||
input[:, :, mem_size:].detach()
|
||||
if (mem_size != 0 and memory_state != MemoryState.DISABLED)
|
||||
else None
|
||||
)
|
||||
if (
|
||||
memory_state != MemoryState.DISABLED
|
||||
and not self.training
|
||||
and (self.memory_device is not None)
|
||||
):
|
||||
self.memory = memory
|
||||
if self.memory_device == "cpu" and self.memory is not None:
|
||||
self.memory = self.memory.to("cpu")
|
||||
return super().forward(input)
|
||||
|
||||
def slicing_forward(
|
||||
self,
|
||||
input: Union[Tensor, List[Tensor]],
|
||||
memory_state: MemoryState = MemoryState.UNSET,
|
||||
preserve_vram: bool = False,
|
||||
) -> Tensor:
|
||||
squeeze_out = False
|
||||
if torch.is_tensor(input):
|
||||
input = [input]
|
||||
squeeze_out = True
|
||||
|
||||
cache_size = self.kernel_size[0] - self.stride[0]
|
||||
cache = cache_send_recv(
|
||||
input, cache_size=cache_size, memory=self.memory, times=self.temporal_padding * 2
|
||||
)
|
||||
|
||||
# Single GPU inference - simplified memory management
|
||||
if (
|
||||
memory_state in [MemoryState.INITIALIZING, MemoryState.ACTIVE] # use_slicing
|
||||
and not self.training
|
||||
and (self.memory_device is not None)
|
||||
and cache_size != 0
|
||||
):
|
||||
if cache_size > input[-1].size(2) and cache is not None and len(input) == 1:
|
||||
input[0] = torch.cat([cache, input[0]], dim=2)
|
||||
cache = None
|
||||
if cache_size <= input[-1].size(2):
|
||||
self.memory = input[-1][:, :, -cache_size:].detach().contiguous()
|
||||
if self.memory_device == "cpu" and self.memory is not None:
|
||||
self.memory = self.memory.to("cpu")
|
||||
|
||||
padding = tuple(x for x in reversed(self.padding) for _ in range(2))
|
||||
for i in range(len(input)):
|
||||
# Prepare cache for next input slice.
|
||||
next_cache = None
|
||||
cache_size = 0
|
||||
if i < len(input) - 1:
|
||||
cache_len = cache.size(2) if cache is not None else 0
|
||||
cache_size = get_cache_size(self, input[i].size(2) + cache_len, pad_len=0)
|
||||
if cache_size != 0:
|
||||
if cache_size > input[i].size(2) and cache is not None:
|
||||
input[i] = torch.cat([cache, input[i]], dim=2)
|
||||
cache = None
|
||||
assert cache_size <= input[i].size(2), f"{cache_size} > {input[i].size(2)}"
|
||||
next_cache = input[i][:, :, -cache_size:]
|
||||
|
||||
# Conv forward for this input slice.
|
||||
input[i] = self.memory_limit_conv(
|
||||
input[i],
|
||||
padding=padding,
|
||||
prev_cache=cache,
|
||||
preserve_vram=preserve_vram
|
||||
)
|
||||
|
||||
# Update cache.
|
||||
cache = next_cache
|
||||
|
||||
return input[0] if squeeze_out else input
|
||||
|
||||
def tflops(self, args, kwargs, output) -> float:
|
||||
if torch.is_tensor(output):
|
||||
output_numel = output.numel()
|
||||
elif isinstance(output, list):
|
||||
output_numel = sum(o.numel() for o in output)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
return (2 * math.prod(self.kernel_size) * self.in_channels * (output_numel / 1e6)) / 1e6
|
||||
|
||||
def _load_from_state_dict(
|
||||
self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs
|
||||
):
|
||||
if self.inflation_mode != "none":
|
||||
state_dict = modify_state_dict(
|
||||
self,
|
||||
state_dict,
|
||||
prefix,
|
||||
inflate_weight_fn=inflate_weight,
|
||||
inflate_bias_fn=inflate_bias,
|
||||
)
|
||||
super()._load_from_state_dict(
|
||||
state_dict,
|
||||
prefix,
|
||||
local_metadata,
|
||||
(strict and self.inflation_mode == "none"),
|
||||
missing_keys,
|
||||
unexpected_keys,
|
||||
error_msgs,
|
||||
)
|
||||
|
||||
|
||||
def init_causal_conv3d(
|
||||
*args,
|
||||
inflation_mode: _inflation_mode_t,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize a Causal-3D convolution layer.
|
||||
Parameters:
|
||||
inflation_mode: Listed as below. It's compatible with all the 3D-VAE checkpoints we have.
|
||||
- none: No inflation will be conducted.
|
||||
The loading logic of state dict will fall back to default.
|
||||
- tail / replicate: Refer to the definition of `InflatedCausalConv3d`.
|
||||
"""
|
||||
return InflatedCausalConv3d(*args, inflation_mode=inflation_mode, **kwargs)
|
||||
|
||||
|
||||
def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor, preserve_vram: bool = False) -> torch.Tensor:
|
||||
input_dtype = x.dtype
|
||||
if isinstance(norm_layer, (nn.LayerNorm, RMSNorm)):
|
||||
if x.ndim == 4:
|
||||
x = rearrange(x, "b c h w -> b h w c")
|
||||
x = norm_layer(x)
|
||||
x = rearrange(x, "b h w c -> b c h w")
|
||||
return x.to(input_dtype)
|
||||
if x.ndim == 5:
|
||||
x = rearrange(x, "b c t h w -> b t h w c")
|
||||
x = norm_layer(x)
|
||||
x = rearrange(x, "b t h w c -> b c t h w")
|
||||
return x.to(input_dtype)
|
||||
if isinstance(norm_layer, (nn.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)):
|
||||
if x.ndim <= 4:
|
||||
return norm_layer(x).to(input_dtype)
|
||||
if x.ndim == 5:
|
||||
t = x.size(2)
|
||||
x = rearrange(x, "b c t h w -> (b t) c h w")
|
||||
memory_occupy = x.numel() * x.element_size() / 1024**3
|
||||
if isinstance(norm_layer, nn.GroupNorm) and memory_occupy > get_norm_limit():
|
||||
num_chunks = min(4 if x.element_size() == 2 else 2, norm_layer.num_groups)
|
||||
assert norm_layer.num_groups % num_chunks == 0
|
||||
num_groups_per_chunk = norm_layer.num_groups // num_chunks
|
||||
|
||||
x = list(x.chunk(num_chunks, dim=1))
|
||||
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)):
|
||||
try:
|
||||
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
|
||||
except Exception as e:
|
||||
print("OOM Second Chance : Group Norm")
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
time.sleep(2)
|
||||
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 platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
else:
|
||||
torch.cuda.empty_cache()
|
||||
x = torch.cat(x, dim=1)
|
||||
else:
|
||||
x = norm_layer(x)
|
||||
x = rearrange(x, "(b t) c h w -> b c t h w", t=t)
|
||||
return x.to(input_dtype)
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def remove_head(tensor: Tensor, times: int = 1) -> Tensor:
|
||||
"""
|
||||
Remove duplicated first frame features in the up-sampling process.
|
||||
"""
|
||||
# Single GPU inference - always process
|
||||
if times == 0:
|
||||
return tensor
|
||||
return torch.cat(tensors=(tensor[:, :, :1], tensor[:, :, times + 1 :]), dim=2)
|
||||
|
||||
|
||||
def extend_head(tensor: Tensor, times: int = 2, memory: Optional[Tensor] = None) -> Tensor:
|
||||
"""
|
||||
When memory is None:
|
||||
- Duplicate first frame features in the down-sampling process.
|
||||
When memory is not None:
|
||||
- Concatenate memory features with the input features to keep temporal consistency.
|
||||
"""
|
||||
if memory is not None:
|
||||
return torch.cat((memory.to(tensor), tensor), dim=2)
|
||||
assert times >= 0, "Invalid input for function 'extend_head'!"
|
||||
if times == 0:
|
||||
return tensor
|
||||
else:
|
||||
tile_repeat = [1] * tensor.ndim
|
||||
tile_repeat[2] = times
|
||||
return torch.cat(tensors=(torch.tile(tensor[:, :, :1], tile_repeat), tensor), dim=2)
|
||||
|
||||
|
||||
def inflate_weight(weight_2d: torch.Tensor, weight_3d: torch.Tensor, inflation_mode: str):
|
||||
"""
|
||||
Inflate a 2D convolution weight matrix to a 3D one.
|
||||
Parameters:
|
||||
weight_2d: The weight matrix of 2D conv to be inflated.
|
||||
weight_3d: The weight matrix of 3D conv to be initialized.
|
||||
inflation_mode: the mode of inflation
|
||||
"""
|
||||
assert inflation_mode in ["tail", "replicate"]
|
||||
assert weight_3d.shape[:2] == weight_2d.shape[:2]
|
||||
with torch.no_grad():
|
||||
if inflation_mode == "replicate":
|
||||
depth = weight_3d.size(2)
|
||||
weight_3d.copy_(weight_2d.unsqueeze(2).repeat(1, 1, depth, 1, 1) / depth)
|
||||
else:
|
||||
weight_3d.fill_(0.0)
|
||||
weight_3d[:, :, -1].copy_(weight_2d)
|
||||
return weight_3d
|
||||
|
||||
|
||||
def inflate_bias(bias_2d: torch.Tensor, bias_3d: torch.Tensor, inflation_mode: str):
|
||||
"""
|
||||
Inflate a 2D convolution bias tensor to a 3D one
|
||||
Parameters:
|
||||
bias_2d: The bias tensor of 2D conv to be inflated.
|
||||
bias_3d: The bias tensor of 3D conv to be initialized.
|
||||
inflation_mode: Placeholder to align `inflate_weight`.
|
||||
"""
|
||||
assert bias_3d.shape == bias_2d.shape
|
||||
with torch.no_grad():
|
||||
bias_3d.copy_(bias_2d)
|
||||
return bias_3d
|
||||
|
||||
|
||||
def modify_state_dict(layer, state_dict, prefix, inflate_weight_fn, inflate_bias_fn):
|
||||
"""
|
||||
the main function to inflated 2D parameters to 3D.
|
||||
"""
|
||||
weight_name = prefix + "weight"
|
||||
bias_name = prefix + "bias"
|
||||
if weight_name in state_dict:
|
||||
weight_2d = state_dict[weight_name]
|
||||
if weight_2d.dim() == 4:
|
||||
# Assuming the 2D weights are 4D tensors (out_channels, in_channels, h, w)
|
||||
weight_3d = inflate_weight_fn(
|
||||
weight_2d=weight_2d,
|
||||
weight_3d=layer.weight,
|
||||
inflation_mode=layer.inflation_mode,
|
||||
)
|
||||
state_dict[weight_name] = weight_3d
|
||||
else:
|
||||
return state_dict
|
||||
# It's a 3d state dict, should not do inflation on both bias and weight.
|
||||
if bias_name in state_dict:
|
||||
bias_2d = state_dict[bias_name]
|
||||
if bias_2d.dim() == 1:
|
||||
# Assuming the 2D biases are 1D tensors (out_channels,)
|
||||
bias_3d = inflate_bias_fn(
|
||||
bias_2d=bias_2d,
|
||||
bias_3d=layer.bias,
|
||||
inflation_mode=layer.inflation_mode,
|
||||
)
|
||||
state_dict[bias_name] = bias_3d
|
||||
return state_dict
|
||||
@@ -17,6 +17,9 @@ import types
|
||||
import torch
|
||||
import weakref
|
||||
import gc
|
||||
import platform
|
||||
import psutil
|
||||
|
||||
from typing import Dict, Any, List, Tuple, Optional, Union
|
||||
from src.optimization.memory_manager import get_vram_usage
|
||||
from src.optimization.compatibility import call_rope_with_stability
|
||||
@@ -78,6 +81,8 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) ->
|
||||
|
||||
# Determine devices
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if platform.system() == "Darwin":
|
||||
device = "mps"
|
||||
offload_device = "cpu"
|
||||
use_non_blocking = block_swap_config.get("use_non_blocking", True)
|
||||
|
||||
@@ -278,7 +283,7 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
|
||||
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:
|
||||
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking and platform.system() != "Darwin":
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Execute forward pass with OOM protection
|
||||
@@ -296,6 +301,10 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.
|
||||
)
|
||||
|
||||
# Only clear cache under memory pressure
|
||||
if platform.system() == "Darwin":
|
||||
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()
|
||||
@@ -350,7 +359,10 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
|
||||
|
||||
# Synchronize if not using non-blocking transfers
|
||||
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking:
|
||||
torch.cuda.synchronize()
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.synchronize()
|
||||
else:
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Execute forward pass
|
||||
output = self._original_forward(*args, **kwargs)
|
||||
@@ -367,6 +379,10 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.
|
||||
)
|
||||
|
||||
# Only clear cache under memory pressure
|
||||
if platform.system() == "Darwin":
|
||||
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()
|
||||
@@ -415,7 +431,8 @@ def _patch_rope_for_blockswap(model, debug) -> None:
|
||||
debug.log(f"RoPE device issue for {module_name}: {e}", level="WARNING", category="blockswap")
|
||||
|
||||
# Get current device from parameters
|
||||
current_device = next(self.parameters()).device if list(self.parameters()) else torch.device("cuda")
|
||||
_device = "mps" if platform.system() == "Darwin" else "cuda"
|
||||
current_device = next(self.parameters()).device if list(self.parameters()) else torch.device(_device)
|
||||
|
||||
# Try clearing cache first (non-invasive fix)
|
||||
if hasattr(current_fn, 'cache_clear'):
|
||||
@@ -677,6 +694,8 @@ def cleanup_blockswap(runner, keep_state_for_cache: bool = False) -> None:
|
||||
gc.collect()
|
||||
|
||||
# Final memory cleanup
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
@@ -6,6 +6,7 @@ Extracted from: seedvr2.py (lines 1045-1630)
|
||||
"""
|
||||
|
||||
import torch
|
||||
import platform
|
||||
import types
|
||||
from typing import List, Tuple, Union, Any, Optional
|
||||
|
||||
@@ -304,17 +305,24 @@ class FP8CompatibleDiT(torch.nn.Module):
|
||||
k = k.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
|
||||
v = v.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
|
||||
|
||||
# Use optimized SDPA
|
||||
with torch.backends.cuda.sdp_kernel(
|
||||
enable_flash=True,
|
||||
enable_math=True,
|
||||
enable_mem_efficient=True
|
||||
):
|
||||
if platform.system() == "Darwin":
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
dropout_p=0.0,
|
||||
is_causal=False
|
||||
)
|
||||
else:
|
||||
# Use optimized SDPA
|
||||
with torch.backends.cuda.sdp_kernel(
|
||||
enable_flash=True,
|
||||
enable_math=True,
|
||||
enable_mem_efficient=True
|
||||
):
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
dropout_p=0.0,
|
||||
is_causal=False
|
||||
)
|
||||
|
||||
# Reshape back
|
||||
attn_output = attn_output.transpose(1, 2).contiguous().view(
|
||||
@@ -389,4 +397,4 @@ class FP8CompatibleDiT(torch.nn.Module):
|
||||
if hasattr(self, 'dit_model'):
|
||||
setattr(self.dit_model, name, value)
|
||||
else:
|
||||
super().__setattr__(name, value)
|
||||
super().__setattr__(name, value)
|
||||
|
||||
@@ -9,6 +9,8 @@ import os
|
||||
import torch
|
||||
import gc
|
||||
import time
|
||||
import platform
|
||||
import psutil
|
||||
from typing import Tuple, Optional
|
||||
from src.common.cache import Cache
|
||||
from src.models.dit_v2.rope import RotaryEmbeddingBase
|
||||
@@ -21,12 +23,16 @@ except:
|
||||
pass
|
||||
|
||||
def get_basic_vram_info():
|
||||
"""🔍 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()
|
||||
if platform.system() == "Darwin":
|
||||
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()
|
||||
|
||||
# Conversion en GB
|
||||
free_gb = free_memory / (1024**3)
|
||||
@@ -52,6 +58,11 @@ def get_vram_usage() -> Tuple[float, float, float]:
|
||||
tuple: (allocated_gb, reserved_gb, max_allocated_gb)
|
||||
Returns (0, 0, 0) if CUDA not available
|
||||
"""
|
||||
if platform.system() == "Darwin":
|
||||
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)
|
||||
@@ -64,10 +75,12 @@ def clear_vram_cache(debug) -> None:
|
||||
"""Clear VRAM cache and run garbage collection"""
|
||||
|
||||
debug.log("Clearing VRAM cache...", category="cleanup")
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
gc.collect()
|
||||
gc.collect()
|
||||
|
||||
|
||||
def reset_vram_peak(debug) -> None:
|
||||
@@ -207,6 +220,9 @@ def fast_ram_cleanup():
|
||||
# Garbage collection
|
||||
gc.collect()
|
||||
|
||||
# Clear MPS cache
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
# Clear CUDA cache
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
@@ -248,7 +264,7 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
|
||||
for key, value in list(runner.cache.cache.items()):
|
||||
if torch.is_tensor(value):
|
||||
# Force deallocation of tensor storage
|
||||
if value.is_cuda:
|
||||
if value.is_cuda or value.is_mps:
|
||||
value.data = value.data.cpu()
|
||||
value.grad = None
|
||||
if value.numel() > 0:
|
||||
@@ -256,7 +272,7 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
|
||||
elif isinstance(value, (list, tuple)):
|
||||
for item in value:
|
||||
if torch.is_tensor(item):
|
||||
if item.is_cuda:
|
||||
if item.is_cuda or item.is_mps:
|
||||
item.data = item.data.cpu()
|
||||
item.grad = None
|
||||
if item.numel() > 0:
|
||||
@@ -373,9 +389,13 @@ def clear_all_caches(runner, debug, offload_vae=False) -> int:
|
||||
# Force garbage collection
|
||||
gc.collect(2) # Collect all generations
|
||||
|
||||
# Clear MPS cache
|
||||
if platform.system() == "Darwin":
|
||||
torch.mps.empty_cache()
|
||||
# Clear CUDA cache
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
return cleaned_items
|
||||
|
||||
Reference in New Issue
Block a user