From 6630d27dcd43df30ab8a97312da759113b483968 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 17 Mar 2024 17:18:45 +0200 Subject: [PATCH] dtype fixes --- lvdm/models/samplers/ddim.py | 10 +++++----- nodes.py | 32 +++++++++++++++++++++++++------- 2 files changed, 30 insertions(+), 12 deletions(-) diff --git a/lvdm/models/samplers/ddim.py b/lvdm/models/samplers/ddim.py index 3133382..7e8237a 100644 --- a/lvdm/models/samplers/ddim.py +++ b/lvdm/models/samplers/ddim.py @@ -37,11 +37,11 @@ class DDIMSampler(object): self.register_buffer('alphas_cumprod_prev', to_torch(self.model.alphas_cumprod_prev)) # calculations for diffusion q(x_t | x_{t-1}) and others - self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod.cpu().to(torch.float32)))) - self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu().to(torch.float32)))) - self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu().to(torch.float32)))) - self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu().to(torch.float32)))) - self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu().to(torch.float32) - 1))) + self.register_buffer('sqrt_alphas_cumprod', to_torch(torch.sqrt(alphas_cumprod))) + self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(torch.sqrt(1. - alphas_cumprod))) + self.register_buffer('log_one_minus_alphas_cumprod', to_torch(torch.log(1. - alphas_cumprod))) + self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(torch.sqrt(1. / alphas_cumprod))) + self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(torch.sqrt(1. / alphas_cumprod))) # ddim sampling parameters ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu().to(torch.float32), diff --git a/nodes.py b/nodes.py index d8e2f96..f7a4278 100644 --- a/nodes.py +++ b/nodes.py @@ -42,8 +42,9 @@ class DynamiCrafterModelLoader: 'fp32', 'fp16', 'bf16', + 'auto' ], { - "default": 'fp16' + "default": 'auto' }), }, } @@ -55,12 +56,10 @@ class DynamiCrafterModelLoader: def loadmodel(self, dtype, ckpt_name): mm.soft_empty_cache() - device = mm.get_torch_device() custom_config = { 'dtype': dtype, 'ckpt_name': ckpt_name, } - dtype = convert_dtype(dtype) if not hasattr(self, 'model') or self.model == None or custom_config != self.current_config: self.current_config = custom_config model_path = folder_paths.get_full_path("checkpoints", ckpt_name) @@ -68,14 +67,33 @@ class DynamiCrafterModelLoader: base_name, _ = os.path.splitext(ckpt_base_name) if 'interp' in base_name and '512' in base_name: config_file=os.path.join(script_directory, "configs", "dynamicrafter_512_interp_v1.yaml") - if '1024' in base_name: + elif '1024' in base_name: config_file=os.path.join(script_directory, "configs", "dynamicrafter_1024_v1.yaml") + elif '512' in base_name: + config_file=os.path.join(script_directory, "configs", "dynamicrafter_512_v1.yaml") + elif '256' in base_name: + config_file=os.path.join(script_directory, "configs", "dynamicrafter_256_v1.yaml") + else: + print(f"No matching config for model: {ckpt_name}") config = OmegaConf.load(config_file) model_config = config.pop("model", OmegaConf.create()) model_config['params']['unet_config']['params']['use_checkpoint']=False self.model = instantiate_from_config(model_config) self.model = load_model_checkpoint(self.model, model_path) - self.model.eval().to(dtype) + self.model.eval() + if dtype == "auto": + try: + if mm.should_use_bf16(): + self.model.to(convert_dtype('bf16')) + elif mm.should_use_fp16(): + self.model.to(convert_dtype('fp16')) + else: + self.model.to(convert_dtype('fp32')) + except: + raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.") + else: + self.model.to(convert_dtype(dtype)) + print(f"Model using dtype: {self.model.dtype}") return (self.model,) class DynamiCrafterI2V: @@ -132,7 +150,7 @@ class DynamiCrafterI2V: raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.") else: model.first_stage_model.to(convert_dtype(vae_dtype)) - print(f"Using {model.first_stage_model.dtype} VAE") + print(f"VAE using dtype: {model.first_stage_model.dtype}") self.model = model self.model.to(device) @@ -311,7 +329,7 @@ class DynamiCrafterBatchInterpolation: raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.") else: model.first_stage_model.to(convert_dtype(vae_dtype)) - print(f"Using {model.first_stage_model.dtype} VAE") + print(f"VAE using dtype: {model.first_stage_model.dtype}") self.model = model self.model.to(device)