Support bf16 safetensors checkpoints

https://huggingface.co/Kijai/DynamiCrafter_pruned/tree/main
This commit is contained in:
kijai
2024-03-17 16:11:17 +02:00
parent b40302a65c
commit 04b7cf96d2
2 changed files with 11 additions and 7 deletions
+6 -6
View File
@@ -37,14 +37,14 @@ 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())))
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu())))
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu())))
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu())))
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1)))
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)))
# ddim sampling parameters
ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu(),
ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu().to(torch.float32),
ddim_timesteps=self.ddim_timesteps,
eta=ddim_eta,verbose=verbose)
self.register_buffer('ddim_sigmas', ddim_sigmas)
+5 -1
View File
@@ -41,6 +41,7 @@ class DynamiCrafterModelLoader:
[
'fp32',
'fp16',
'bf16',
], {
"default": 'fp16'
}),
@@ -65,7 +66,10 @@ class DynamiCrafterModelLoader:
model_path = folder_paths.get_full_path("checkpoints", ckpt_name)
ckpt_base_name = os.path.basename(ckpt_name)
base_name, _ = os.path.splitext(ckpt_base_name)
config_file=os.path.join(script_directory, "configs", f"{base_name}.yaml")
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:
config_file=os.path.join(script_directory, "configs", "dynamicrafter_1024_v1.yaml")
config = OmegaConf.load(config_file)
model_config = config.pop("model", OmegaConf.create())
model_config['params']['unet_config']['params']['use_checkpoint']=False