MCPcopy Create free account
hub / github.com/YesianRohn/TextSSR / main

Function main

diffusers/scripts/convert_hunyuandit_to_diffusers.py:8–237  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

6
7
8def main(args):
9 state_dict = torch.load(args.pt_checkpoint_path, map_location="cpu")
10
11 if args.load_key != "none":
12 try:
13 state_dict = state_dict[args.load_key]
14 except KeyError:
15 raise KeyError(
16 f"{args.load_key} not found in the checkpoint."
17 f"Please load from the following keys:{state_dict.keys()}"
18 )
19
20 device = "cuda"
21 model_config = HunyuanDiT2DModel.load_config("Tencent-Hunyuan/HunyuanDiT-Diffusers", subfolder="transformer")
22 model_config[
23 "use_style_cond_and_image_meta_size"
24 ] = args.use_style_cond_and_image_meta_size ### version <= v1.1: True; version >= v1.2: False
25
26 # input_size -> sample_size, text_dim -> cross_attention_dim
27 for key in state_dict:
28 print("local:", key)
29
30 model = HunyuanDiT2DModel.from_config(model_config).to(device)
31
32 for key in model.state_dict():
33 print("diffusers:", key)
34
35 num_layers = 40
36 for i in range(num_layers):
37 # attn1
38 # Wkqv -> to_q, to_k, to_v
39 q, k, v = torch.chunk(state_dict[f"blocks.{i}.attn1.Wqkv.weight"], 3, dim=0)
40 q_bias, k_bias, v_bias = torch.chunk(state_dict[f"blocks.{i}.attn1.Wqkv.bias"], 3, dim=0)
41 state_dict[f"blocks.{i}.attn1.to_q.weight"] = q
42 state_dict[f"blocks.{i}.attn1.to_q.bias"] = q_bias
43 state_dict[f"blocks.{i}.attn1.to_k.weight"] = k
44 state_dict[f"blocks.{i}.attn1.to_k.bias"] = k_bias
45 state_dict[f"blocks.{i}.attn1.to_v.weight"] = v
46 state_dict[f"blocks.{i}.attn1.to_v.bias"] = v_bias
47 state_dict.pop(f"blocks.{i}.attn1.Wqkv.weight")
48 state_dict.pop(f"blocks.{i}.attn1.Wqkv.bias")
49
50 # q_norm, k_norm -> norm_q, norm_k
51 state_dict[f"blocks.{i}.attn1.norm_q.weight"] = state_dict[f"blocks.{i}.attn1.q_norm.weight"]
52 state_dict[f"blocks.{i}.attn1.norm_q.bias"] = state_dict[f"blocks.{i}.attn1.q_norm.bias"]
53 state_dict[f"blocks.{i}.attn1.norm_k.weight"] = state_dict[f"blocks.{i}.attn1.k_norm.weight"]
54 state_dict[f"blocks.{i}.attn1.norm_k.bias"] = state_dict[f"blocks.{i}.attn1.k_norm.bias"]
55
56 state_dict.pop(f"blocks.{i}.attn1.q_norm.weight")
57 state_dict.pop(f"blocks.{i}.attn1.q_norm.bias")
58 state_dict.pop(f"blocks.{i}.attn1.k_norm.weight")
59 state_dict.pop(f"blocks.{i}.attn1.k_norm.bias")
60
61 # out_proj -> to_out
62 state_dict[f"blocks.{i}.attn1.to_out.0.weight"] = state_dict[f"blocks.{i}.attn1.out_proj.weight"]
63 state_dict[f"blocks.{i}.attn1.to_out.0.bias"] = state_dict[f"blocks.{i}.attn1.out_proj.bias"]
64 state_dict.pop(f"blocks.{i}.attn1.out_proj.weight")
65 state_dict.pop(f"blocks.{i}.attn1.out_proj.bias")

Calls 10

load_configMethod · 0.80
load_state_dictMethod · 0.80
swap_scale_shiftFunction · 0.70
loadMethod · 0.45
toMethod · 0.45
from_configMethod · 0.45
state_dictMethod · 0.45
popMethod · 0.45
from_pretrainedMethod · 0.45
save_pretrainedMethod · 0.45

Tested by

no test coverage detected