diff --git a/library/flux_utils.py b/library/flux_utils.py index 10d0e7e..f8a4c0c 100644 --- a/library/flux_utils.py +++ b/library/flux_utils.py @@ -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 diff --git a/library/model_util.py b/library/model_util.py index d8d79f8..47c5674 100644 --- a/library/model_util.py +++ b/library/model_util.py @@ -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}")