Improve memory management

This commit is contained in:
kijai
2024-04-25 01:53:28 +03:00
parent 98f08a4229
commit 3d73bbe626
3 changed files with 35 additions and 18 deletions
+29 -15
View File
@@ -5,8 +5,16 @@ import torch.nn as nn
from torch import Tensor from torch import Tensor
from functools import partial from functools import partial
from torchdiffeq import odeint from torchdiffeq import odeint
from contextlib import nullcontext
try:
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
is_accelerate_available = True
except:
is_accelerate_available = False
pass
#from unet import UNetModel
from .unet.openaimodel import UNetModel from .unet.openaimodel import UNetModel
def exists(val): def exists(val):
@@ -14,19 +22,24 @@ def exists(val):
class DepthFM(nn.Module): class DepthFM(nn.Module):
def __init__(self, vae, ckpt_path: str): def __init__(self, vae, ckpt_path: str, device, offload_device,dtype):
super().__init__() super().__init__()
#vae_id = "runwayml/stable-diffusion-v1-5"
#self.vae = AutoencoderKL.from_pretrained(vae_id, subfolder="vae")
self.vae = vae self.vae = vae
self.scale_factor = 0.18215 self.scale_factor = 0.18215
self.device = device
self.offload_device = offload_device
# set with checkpoint # set with checkpoint
ckpt = torch.load(ckpt_path, map_location="cpu") state_dict = torch.load(ckpt_path)
self.noising_step = ckpt['noising_step'] self.noising_step = state_dict['noising_step']
self.empty_text_embed = ckpt['empty_text_embedding'] self.empty_text_embed = state_dict['empty_text_embedding']
self.model = UNetModel(**ckpt['ldm_hparams']) with (init_empty_weights() if is_accelerate_available else nullcontext()):
self.model.load_state_dict(ckpt['state_dict']) self.model = UNetModel(**state_dict['ldm_hparams'])
if is_accelerate_available:
for key in state_dict['state_dict']:
set_module_tensor_to_device(self.model, key, device=device, dtype=dtype, value=state_dict['state_dict'][key])
else:
self.model.load_state_dict(state_dict['state_dict'])
def ode_fn(self, t: Tensor, x: Tensor, **kwargs): def ode_fn(self, t: Tensor, x: Tensor, **kwargs):
if t.numel() == 1: if t.numel() == 1:
@@ -67,7 +80,9 @@ class DepthFM(nn.Module):
bs, dev = ims.shape[0], ims.device bs, dev = ims.shape[0], ims.device
self.vae.first_stage_model = self.vae.first_stage_model.to(self.device)
ims_z = self.encode(ims, sample_posterior=False) ims_z = self.encode(ims, sample_posterior=False)
self.vae.first_stage_model = self.vae.first_stage_model.to(self.offload_device)
conditioning = torch.tensor(self.empty_text_embed).to(dev).repeat(bs, 1, 1) conditioning = torch.tensor(self.empty_text_embed).to(dev).repeat(bs, 1, 1)
context = ims_z context = ims_z
@@ -78,9 +93,14 @@ class DepthFM(nn.Module):
x_source = q_sample(x_source, self.noising_step) x_source = q_sample(x_source, self.noising_step)
# solve ODE # solve ODE
self.model.to(self.device)
depth_z = self.generate(x_source, num_steps=num_steps, context=context, context_ca=conditioning) depth_z = self.generate(x_source, num_steps=num_steps, context=context, context_ca=conditioning)
self.model.to(self.offload_device)
self.vae.first_stage_model = self.vae.first_stage_model.to(self.device)
depth = self.decode(depth_z) depth = self.decode(depth_z)
self.vae.first_stage_model = self.vae.first_stage_model.to(self.offload_device)
depth = depth.mean(dim=1, keepdim=True) depth = depth.mean(dim=1, keepdim=True)
if ensemble_size > 1: if ensemble_size > 1:
@@ -98,13 +118,7 @@ class DepthFM(nn.Module):
@torch.no_grad() @torch.no_grad()
def encode(self, x: Tensor, sample_posterior: bool = True): def encode(self, x: Tensor, sample_posterior: bool = True):
self.vae.first_stage_model = self.vae.first_stage_model.to(x.device)
z = self.vae.first_stage_model.encode(x) z = self.vae.first_stage_model.encode(x)
# if sample_posterior:
# z = posterior.latent_dist.sample()
# else:
# z = posterior.latent_dist.mode()
# normalize latent code
z = z * self.scale_factor z = z * self.scale_factor
return z return z
+3
View File
@@ -9,6 +9,9 @@ from typing import Optional, Any
from .util import checkpoint from .util import checkpoint
import model_management import model_management
if model_management.XFORMERS_IS_AVAILABLE:
import xformers
class Conv2d(torch.nn.Conv2d): class Conv2d(torch.nn.Conv2d):
def reset_parameters(self): def reset_parameters(self):
return None return None
+2 -2
View File
@@ -1,4 +1,3 @@
import os
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
from .depthfm.dfm import DepthFM from .depthfm.dfm import DepthFM
@@ -46,6 +45,7 @@ class Depth_fm:
def process(self, depthfm_model, vae, images, ensemble_size, steps, dtype, invert, per_batch): def process(self, depthfm_model, vae, images, ensemble_size, steps, dtype, invert, per_batch):
device = model_management.get_torch_device() device = model_management.get_torch_device()
offload_device = model_management.unet_offload_device()
dtype = convert_dtype(dtype) dtype = convert_dtype(dtype)
custom_config = { custom_config = {
@@ -55,7 +55,7 @@ class Depth_fm:
if not hasattr(self, "model") or custom_config != self.current_config: if not hasattr(self, "model") or custom_config != self.current_config:
self.current_config = custom_config self.current_config = custom_config
DEPTHFM_MODEL_PATH = folder_paths.get_full_path("checkpoints", depthfm_model) DEPTHFM_MODEL_PATH = folder_paths.get_full_path("checkpoints", depthfm_model)
self.model = DepthFM(vae, DEPTHFM_MODEL_PATH) self.model = DepthFM(vae, DEPTHFM_MODEL_PATH, device, offload_device, dtype)
self.model.eval().to(dtype).to(device) self.model.eval().to(dtype).to(device)
images = images.permute(0, 3, 1, 2) images = images.permute(0, 3, 1, 2)