allow loading models saved in Comfy
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user