comdoleger commited on
Commit
6e01e99
·
verified ·
1 Parent(s): 48a89e2

Upload testing/generate_lora_mapping.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. testing/generate_lora_mapping.py +130 -0
testing/generate_lora_mapping.py ADDED
@@ -0,0 +1,130 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from collections import OrderedDict
2
+
3
+ import torch
4
+ from safetensors.torch import load_file
5
+ import argparse
6
+ import os
7
+ import json
8
+
9
+ PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
10
+
11
+ keymap_path = os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps', 'stable_diffusion_sdxl.json')
12
+
13
+ # load keymap
14
+ with open(keymap_path, 'r') as f:
15
+ keymap = json.load(f)
16
+
17
+ lora_keymap = OrderedDict()
18
+
19
+ # convert keymap to lora key naming
20
+ for ldm_key, diffusers_key in keymap['ldm_diffusers_keymap'].items():
21
+ if ldm_key.endswith('.bias') or diffusers_key.endswith('.bias'):
22
+ # skip it
23
+ continue
24
+ # sdxl has same te for locon with kohya and ours
25
+ if ldm_key.startswith('conditioner'):
26
+ #skip it
27
+ continue
28
+ # ignore vae
29
+ if ldm_key.startswith('first_stage_model'):
30
+ continue
31
+ ldm_key = ldm_key.replace('model.diffusion_model.', 'lora_unet_')
32
+ ldm_key = ldm_key.replace('.weight', '')
33
+ ldm_key = ldm_key.replace('.', '_')
34
+
35
+ diffusers_key = diffusers_key.replace('unet_', 'lora_unet_')
36
+ diffusers_key = diffusers_key.replace('.weight', '')
37
+ diffusers_key = diffusers_key.replace('.', '_')
38
+
39
+ lora_keymap[f"{ldm_key}.alpha"] = f"{diffusers_key}.alpha"
40
+ lora_keymap[f"{ldm_key}.lora_down.weight"] = f"{diffusers_key}.lora_down.weight"
41
+ lora_keymap[f"{ldm_key}.lora_up.weight"] = f"{diffusers_key}.lora_up.weight"
42
+
43
+
44
+ parser = argparse.ArgumentParser()
45
+ parser.add_argument("input", help="input file")
46
+ parser.add_argument("input2", help="input2 file")
47
+
48
+ args = parser.parse_args()
49
+
50
+ # name = args.name
51
+ # if args.sdxl:
52
+ # name += '_sdxl'
53
+ # elif args.sd2:
54
+ # name += '_sd2'
55
+ # else:
56
+ # name += '_sd1'
57
+ name = 'stable_diffusion_locon_sdxl'
58
+
59
+ locon_save = load_file(args.input)
60
+ our_save = load_file(args.input2)
61
+
62
+ our_extra_keys = list(set(our_save.keys()) - set(locon_save.keys()))
63
+ locon_extra_keys = list(set(locon_save.keys()) - set(our_save.keys()))
64
+
65
+ print(f"we have {len(our_extra_keys)} extra keys")
66
+ print(f"locon has {len(locon_extra_keys)} extra keys")
67
+
68
+ save_dtype = torch.float16
69
+ print(f"our extra keys: {our_extra_keys}")
70
+ print(f"locon extra keys: {locon_extra_keys}")
71
+
72
+
73
+ def export_state_dict(our_save):
74
+ converted_state_dict = OrderedDict()
75
+ for key, value in our_save.items():
76
+ # test encoders share keys for some reason
77
+ if key.startswith('lora_te'):
78
+ converted_state_dict[key] = value.detach().to('cpu', dtype=save_dtype)
79
+ else:
80
+ converted_key = key
81
+ for ldm_key, diffusers_key in lora_keymap.items():
82
+ if converted_key == diffusers_key:
83
+ converted_key = ldm_key
84
+
85
+ converted_state_dict[converted_key] = value.detach().to('cpu', dtype=save_dtype)
86
+ return converted_state_dict
87
+
88
+ def import_state_dict(loaded_state_dict):
89
+ converted_state_dict = OrderedDict()
90
+ for key, value in loaded_state_dict.items():
91
+ if key.startswith('lora_te'):
92
+ converted_state_dict[key] = value.detach().to('cpu', dtype=save_dtype)
93
+ else:
94
+ converted_key = key
95
+ for ldm_key, diffusers_key in lora_keymap.items():
96
+ if converted_key == ldm_key:
97
+ converted_key = diffusers_key
98
+
99
+ converted_state_dict[converted_key] = value.detach().to('cpu', dtype=save_dtype)
100
+ return converted_state_dict
101
+
102
+
103
+ # check it again
104
+ converted_state_dict = export_state_dict(our_save)
105
+ converted_extra_keys = list(set(converted_state_dict.keys()) - set(locon_save.keys()))
106
+ locon_extra_keys = list(set(locon_save.keys()) - set(converted_state_dict.keys()))
107
+
108
+
109
+ print(f"we have {len(converted_extra_keys)} extra keys")
110
+ print(f"locon has {len(locon_extra_keys)} extra keys")
111
+
112
+ print(f"our extra keys: {converted_extra_keys}")
113
+
114
+ # convert back
115
+ cycle_state_dict = import_state_dict(converted_state_dict)
116
+ cycle_extra_keys = list(set(cycle_state_dict.keys()) - set(our_save.keys()))
117
+ our_extra_keys = list(set(our_save.keys()) - set(cycle_state_dict.keys()))
118
+
119
+ print(f"we have {len(our_extra_keys)} extra keys")
120
+ print(f"cycle has {len(cycle_extra_keys)} extra keys")
121
+
122
+ # save keymap
123
+ to_save = OrderedDict()
124
+ to_save['ldm_diffusers_keymap'] = lora_keymap
125
+
126
+ with open(os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps', f'{name}.json'), 'w') as f:
127
+ json.dump(to_save, f, indent=4)
128
+
129
+
130
+