Improve memory management
This commit is contained in:
+29
-15
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user