diff --git a/depthfm/dfm.py b/depthfm/dfm.py index 63899cb..9862f28 100644 --- a/depthfm/dfm.py +++ b/depthfm/dfm.py @@ -5,8 +5,16 @@ import torch.nn as nn from torch import Tensor from functools import partial 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 def exists(val): @@ -14,19 +22,24 @@ def exists(val): class DepthFM(nn.Module): - def __init__(self, vae, ckpt_path: str): + def __init__(self, vae, ckpt_path: str, device, offload_device,dtype): super().__init__() - #vae_id = "runwayml/stable-diffusion-v1-5" - #self.vae = AutoencoderKL.from_pretrained(vae_id, subfolder="vae") self.vae = vae self.scale_factor = 0.18215 + self.device = device + self.offload_device = offload_device # set with checkpoint - ckpt = torch.load(ckpt_path, map_location="cpu") - self.noising_step = ckpt['noising_step'] - self.empty_text_embed = ckpt['empty_text_embedding'] - self.model = UNetModel(**ckpt['ldm_hparams']) - self.model.load_state_dict(ckpt['state_dict']) + state_dict = torch.load(ckpt_path) + self.noising_step = state_dict['noising_step'] + self.empty_text_embed = state_dict['empty_text_embedding'] + with (init_empty_weights() if is_accelerate_available else nullcontext()): + 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): if t.numel() == 1: @@ -67,7 +80,9 @@ class DepthFM(nn.Module): 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) + 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) context = ims_z @@ -78,9 +93,14 @@ class DepthFM(nn.Module): x_source = q_sample(x_source, self.noising_step) # solve ODE + self.model.to(self.device) 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) + self.vae.first_stage_model = self.vae.first_stage_model.to(self.offload_device) + depth = depth.mean(dim=1, keepdim=True) if ensemble_size > 1: @@ -98,13 +118,7 @@ class DepthFM(nn.Module): @torch.no_grad() 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) - # if sample_posterior: - # z = posterior.latent_dist.sample() - # else: - # z = posterior.latent_dist.mode() - # normalize latent code + z = self.vae.first_stage_model.encode(x) z = z * self.scale_factor return z diff --git a/depthfm/unet/attention.py b/depthfm/unet/attention.py index 68d1a48..0bb1b8d 100644 --- a/depthfm/unet/attention.py +++ b/depthfm/unet/attention.py @@ -9,6 +9,9 @@ from typing import Optional, Any from .util import checkpoint import model_management +if model_management.XFORMERS_IS_AVAILABLE: + import xformers + class Conv2d(torch.nn.Conv2d): def reset_parameters(self): return None diff --git a/nodes.py b/nodes.py index 97c5fb0..21ef1fa 100644 --- a/nodes.py +++ b/nodes.py @@ -1,4 +1,3 @@ -import os import torch import torch.nn.functional as F 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): device = model_management.get_torch_device() + offload_device = model_management.unet_offload_device() dtype = convert_dtype(dtype) custom_config = { @@ -55,7 +55,7 @@ class Depth_fm: if not hasattr(self, "model") or custom_config != self.current_config: self.current_config = custom_config 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) images = images.permute(0, 3, 1, 2)