Views
No views yet
| Dev (50 steps) | Dev (4 steps) | Dev + Schnell (4 steps) |
|---|---|---|
![]() |
![]() |
![]() |
1from diffusers import FluxTransformer2DModel
2from huggingface_hub import snapshot_download
3from accelerate import init_empty_weights
4from diffusers.models.model_loading_utils import load_model_dict_into_meta
5import safetensors.torch
6import glob
7import torch
8
9
10with init_empty_weights():
11 config = FluxTransformer2DModel.load_config("black-forest-labs/FLUX.1-dev", subfolder="transformer")
12 model = FluxTransformer2DModel.from_config(config)
13
14dev_ckpt = snapshot_download(repo_id="black-forest-labs/FLUX.1-dev", allow_patterns="transformer/*")
15schnell_ckpt = snapshot_download(repo_id="black-forest-labs/FLUX.1-schnell", allow_patterns="transformer/*")
16
17dev_shards = sorted(glob.glob(f"{dev_ckpt}/transformer/*.safetensors"))
18schnell_shards = sorted(glob.glob(f"{schnell_ckpt}/transformer/*.safetensors"))
19
20merged_state_dict = {}
21guidance_state_dict = {}
22
23for i in range(len((dev_shards))):
24 state_dict_dev_temp = safetensors.torch.load_file(dev_shards[i])
25 state_dict_schnell_temp = safetensors.torch.load_file(schnell_shards[i])
26
27 keys = list(state_dict_dev_temp.keys())
28 for k in keys:
29 if "guidance" not in k:
30 merged_state_dict[k] = (state_dict_dev_temp.pop(k) + state_dict_schnell_temp.pop(k)) / 2
31 else:
32 guidance_state_dict[k] = state_dict_dev_temp.pop(k)
33
34 if len(state_dict_dev_temp) > 0:
35 raise ValueError(f"There should not be any residue but got: {list(state_dict_dev_temp.keys())}.")
36 if len(state_dict_schnell_temp) > 0:
37 raise ValueError(f"There should not be any residue but got: {list(state_dict_dev_temp.keys())}.")
38
39merged_state_dict.update(guidance_state_dict)
40load_model_dict_into_meta(model, merged_state_dict)
41
42model.to(torch.bfloat16).save_pretrained("merged-flux")1from diffusers import FluxPipeline
2import torch
3
4pipeline = FluxPipeline.from_pretrained(
5 "sayakpaul/FLUX.1-merged", torch_dtype=torch.bfloat16
6).to("cuda")
7image = pipeline(
8 prompt="a tiny astronaut hatching from an egg on the moon",
9 guidance_scale=3.5,
10 num_inference_steps=4,
11 height=880,
12 width=1184,
13 max_sequence_length=512,
14 generator=torch.manual_seed(0),
15).images[0]
16image.save("merged_flux.png")