Support bf16 safetensors checkpoints
https://huggingface.co/Kijai/DynamiCrafter_pruned/tree/main
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user