Merge pull request #103 from lihaoyun6/main

Added MPS backend support (for running on macOS)
This commit is contained in:
Adrien Toupet
2025-08-12 16:44:24 +02:00
committed by GitHub
18 changed files with 2081 additions and 87 deletions
+2 -1
View File
@@ -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
View File
@@ -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")
+9 -3
View File
@@ -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
+6 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
+12 -1
View File
@@ -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():
+3
View File
@@ -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
+5 -4
View File
@@ -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),
]
+3 -1
View File
@@ -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]):
"""
+4 -5
View File
@@ -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
+23 -4
View File
@@ -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()
+15 -7
View File
@@ -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)
+29 -9
View File
@@ -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