allow loading models saved in Comfy

This commit is contained in:
kijai
2024-09-08 15:57:45 +03:00
parent 280f2557b9
commit 95de0cbcc9
2 changed files with 18 additions and 0 deletions
+9
View File
@@ -50,6 +50,15 @@ def load_flow_model(
# load_sft doesn't support torch.device
logger.info(f"Loading state dict from {ckpt_path}")
sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype)
# Check the first key to see if it contains the prefix
first_key = next(iter(sd))
if first_key.startswith("model.diffusion_model."):
# Remove the 'model.diffusion_model.' prefix from keys if it exists
sd = {
key.replace("model.diffusion_model.", ""): value
for key, value in sd.items()
}
info = model.load_state_dict(sd, strict=False, assign=True)
logger.info(f"Loaded Flux: {info}")
return model
+9
View File
@@ -1004,6 +1004,15 @@ def load_models_from_stable_diffusion_checkpoint(v2, ckpt_path, device="cpu", dt
unet_config = create_unet_diffusers_config(v2, unet_use_linear_projection_in_v2)
converted_unet_checkpoint = convert_ldm_unet_checkpoint(v2, state_dict, unet_config)
# convert keys of comfy saved models
first_key = next(iter(converted_unet_checkpoint))
if first_key.startswith("model.diffusion_model."):
# Remove the 'model.diffusion_model.' prefix from keys if it exists
converted_unet_checkpoint = {
key.replace("model.diffusion_model.", ""): value
for key, value in converted_unet_checkpoint.items()
}
unet = UNet2DConditionModel(**unet_config).to(device)
info = unet.load_state_dict(converted_unet_checkpoint)
logger.info(f"loading u-net: {info}")