dtype fixes

This commit is contained in:
kijai
2024-03-17 17:18:45 +02:00
parent 794a70c912
commit 6630d27dcd
2 changed files with 30 additions and 12 deletions
+5 -5
View File
@@ -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),
+25 -7
View File
@@ -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)