Views
No views yet
black-forest-labs/FLUX.1-dev
and black-forest-labs/FLUX.1-schnell originally provided by @sayakpaul.| Dev (50 steps) | Dev (4 steps) | Dev + Schnell Merge (4 steps) | This Model (6-8 steps recommended) |
|---|---|---|---|
![]() |
![]() |
![]() |
![]() |
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")