diff --git a/docs/source/en/api/pipelines/krea2.md b/docs/source/en/api/pipelines/krea2.md index 4c7107425d86..19c8a1cb5184 100644 --- a/docs/source/en/api/pipelines/krea2.md +++ b/docs/source/en/api/pipelines/krea2.md @@ -70,6 +70,17 @@ image = pipe( image.save("krea2_turbo.png") ``` +## Loading single-file checkpoints + +```python +import torch +from diffusers import Krea2Pipeline, Krea2Transformer2DModel + +transformer = Krea2Transformer2DModel.from_single_file( + "https://huggingface.co/krea/Krea-2-Turbo/blob/main/turbo.safetensors", dtype=torch.bfloat16 +) +pipe = Krea2Pipeline.from_pretrained("krea/Krea-2-Turbo", transformer=transformer, dtype=torch.bfloat16).to("cuda") +``` ## Krea2Pipeline diff --git a/src/diffusers/loaders/single_file_model.py b/src/diffusers/loaders/single_file_model.py index 57ac5f71fb19..1556227673f4 100644 --- a/src/diffusers/loaders/single_file_model.py +++ b/src/diffusers/loaders/single_file_model.py @@ -42,6 +42,7 @@ convert_flux_transformer_checkpoint_to_diffusers, convert_hidream_transformer_to_diffusers, convert_hunyuan_video_transformer_to_diffusers, + convert_krea2_transformer_checkpoint_to_diffusers, convert_ldm_unet_checkpoint, convert_ldm_vae_checkpoint, convert_ltx2_audio_vae_to_diffusers, @@ -199,6 +200,10 @@ "checkpoint_mapping_fn": convert_qwen_image21_transformer_checkpoint_to_diffusers, "default_subfolder": "transformer", }, + "Krea2Transformer2DModel": { + "checkpoint_mapping_fn": convert_krea2_transformer_checkpoint_to_diffusers, + "default_subfolder": "transformer", + }, "Flux2Transformer2DModel": { "checkpoint_mapping_fn": convert_flux2_transformer_checkpoint_to_diffusers, "default_subfolder": "transformer", diff --git a/src/diffusers/loaders/single_file_utils.py b/src/diffusers/loaders/single_file_utils.py index 35d75d678087..3e160be02b4b 100644 --- a/src/diffusers/loaders/single_file_utils.py +++ b/src/diffusers/loaders/single_file_utils.py @@ -160,6 +160,7 @@ "audio_vae.per_channel_statistics.mean-of-means", ], "qwen-image-2.1": ["model.diffusion_model.txt_in.text_norm.weight", "txt_in.text_norm.weight"], + "krea2": ["model.diffusion_model.txtfusion.projector.weight", "txtfusion.projector.weight"], } DIFFUSERS_DEFAULT_PIPELINE_PATHS = { @@ -245,6 +246,7 @@ "z-image-turbo-controlnet-2.1": {"pretrained_model_name_or_path": "hlky/Z-Image-Turbo-Fun-Controlnet-Union-2.1"}, "ltx2-dev": {"pretrained_model_name_or_path": "Lightricks/LTX-2"}, "qwen-image-2.1": {"pretrained_model_name_or_path": "Qwen/Qwen-Image-2.1"}, + "krea2": {"pretrained_model_name_or_path": "krea/Krea-2-Raw"}, "minimax-h3": {"pretrained_model_name_or_path": "MiniMaxAI/MiniMax-H3"}, } @@ -791,6 +793,9 @@ def infer_diffusers_model_type(checkpoint): elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["qwen-image-2.1"]): model_type = "qwen-image-2.1" + elif any(key in checkpoint for key in CHECKPOINT_KEY_NAMES["krea2"]): + model_type = "krea2" + elif CHECKPOINT_KEY_NAMES["wan_vae"] in checkpoint: # All Wan models use the same VAE so we can use the same default model repo to fetch the config model_type = "wan-t2v-14B" @@ -4335,3 +4340,51 @@ def convert_qwen_image21_transformer_checkpoint_to_diffusers(checkpoint, **kwarg converted_state_dict[new_key] = value return converted_state_dict + + +def convert_krea2_transformer_checkpoint_to_diffusers(checkpoint, **kwargs): + prefix_rename_dict = { + "first.": "img_in.", + "tmlp.0.": "time_embed.linear_1.", + "tmlp.2.": "time_embed.linear_2.", + "tproj.1.": "time_mod_proj.", + "txtmlp.0.scale": "txt_in.norm.weight", + "txtmlp.1.": "txt_in.linear_1.", + "txtmlp.3.": "txt_in.linear_2.", + "txtfusion.": "text_fusion.", + "blocks.": "transformer_blocks.", + "last.linear.": "final_layer.linear.", + "last.norm.scale": "final_layer.norm.weight", + "last.modulation.lin": "final_layer.scale_shift_table", + } + block_rename_dict = { + ".attn.wq.": ".attn.to_q.", + ".attn.wk.": ".attn.to_k.", + ".attn.wv.": ".attn.to_v.", + ".attn.wo.": ".attn.to_out.0.", + ".attn.gate.": ".attn.to_gate.", + ".attn.qknorm.qnorm.scale": ".attn.norm_q.weight", + ".attn.qknorm.knorm.scale": ".attn.norm_k.weight", + ".mlp.": ".ff.", + ".prenorm.scale": ".norm1.weight", + ".postnorm.scale": ".norm2.weight", + ".mod.lin": ".scale_shift_table", + } + + converted_state_dict = {} + for key in list(checkpoint.keys()): + new_key = key.replace("model.diffusion_model.", "") + for old, new in prefix_rename_dict.items(): + if new_key.startswith(old): + new_key = new + new_key[len(old) :] + break + for old, new in block_rename_dict.items(): + new_key = new_key.replace(old, new) + + value = checkpoint.pop(key) + # The original checkpoint stores each block's six modulation vectors flattened into one. + if new_key.startswith("transformer_blocks.") and new_key.endswith(".scale_shift_table"): + value = value.reshape(6, -1) + converted_state_dict[new_key] = value + + return converted_state_dict diff --git a/src/diffusers/models/transformers/transformer_krea2.py b/src/diffusers/models/transformers/transformer_krea2.py index 55d275e5dca7..980d942831e6 100644 --- a/src/diffusers/models/transformers/transformer_krea2.py +++ b/src/diffusers/models/transformers/transformer_krea2.py @@ -21,7 +21,7 @@ import torch.nn.functional as F from ...configuration_utils import ConfigMixin, register_to_config -from ...loaders import PeftAdapterMixin +from ...loaders import FromOriginalModelMixin, PeftAdapterMixin from ...utils import apply_lora_scale, logging from ...utils.torch_utils import maybe_adjust_dtype_for_device from ..attention import AttentionMixin, AttentionModuleMixin @@ -336,7 +336,7 @@ def forward(self, ids: torch.Tensor) -> torch.Tensor: return freqs_cos, freqs_sin -class Krea2Transformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin): +class Krea2Transformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin, FromOriginalModelMixin): r""" The single-stream MMDiT flow-matching backbone used by the Krea 2 pipeline. diff --git a/tests/models/transformers/test_models_transformer_krea2.py b/tests/models/transformers/test_models_transformer_krea2.py index 261bc13e77b9..3b5fb4f37ac2 100644 --- a/tests/models/transformers/test_models_transformer_krea2.py +++ b/tests/models/transformers/test_models_transformer_krea2.py @@ -25,6 +25,7 @@ LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, + SingleFileTesterMixin, TorchCompileTesterMixin, TrainingTesterMixin, ) @@ -159,3 +160,21 @@ class TestKrea2TransformerAttention(Krea2TransformerTesterConfig, AttentionTeste class TestKrea2TransformerLoRA(Krea2TransformerTesterConfig, LoraTesterMixin): pass + + +class TestKrea2TransformerSingleFile(Krea2TransformerTesterConfig, SingleFileTesterMixin): + @property + def ckpt_path(self): + return "https://huggingface.co/krea/Krea-2-Raw/blob/main/raw.safetensors" + + @property + def pretrained_model_name_or_path(self): + return "krea/Krea-2-Raw" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} + + @property + def torch_dtype(self): + return torch.bfloat16