comdoleger commited on
Commit
48a89e2
·
verified ·
1 Parent(s): 744e37c

Upload testing/compare_keys.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. testing/compare_keys.py +99 -0
testing/compare_keys.py ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+
4
+ import torch
5
+ from diffusers.loaders import LoraLoaderMixin
6
+ from safetensors.torch import load_file
7
+ from collections import OrderedDict
8
+ import json
9
+ # this was just used to match the vae keys to the diffusers keys
10
+ # you probably wont need this. Unless they change them.... again... again
11
+ # on second thought, you probably will
12
+
13
+ device = torch.device('cpu')
14
+ dtype = torch.float32
15
+
16
+ parser = argparse.ArgumentParser()
17
+
18
+ # require at lease one config file
19
+ parser.add_argument(
20
+ 'file_1',
21
+ nargs='+',
22
+ type=str,
23
+ help='Path to first safe tensor file'
24
+ )
25
+
26
+ parser.add_argument(
27
+ 'file_2',
28
+ nargs='+',
29
+ type=str,
30
+ help='Path to second safe tensor file'
31
+ )
32
+
33
+ args = parser.parse_args()
34
+
35
+ find_matches = False
36
+
37
+ state_dict_file_1 = load_file(args.file_1[0])
38
+ state_dict_1_keys = list(state_dict_file_1.keys())
39
+
40
+ state_dict_file_2 = load_file(args.file_2[0])
41
+ state_dict_2_keys = list(state_dict_file_2.keys())
42
+ keys_in_both = []
43
+
44
+ keys_not_in_state_dict_2 = []
45
+ for key in state_dict_1_keys:
46
+ if key not in state_dict_2_keys:
47
+ keys_not_in_state_dict_2.append(key)
48
+
49
+ keys_not_in_state_dict_1 = []
50
+ for key in state_dict_2_keys:
51
+ if key not in state_dict_1_keys:
52
+ keys_not_in_state_dict_1.append(key)
53
+
54
+ keys_in_both = []
55
+ for key in state_dict_1_keys:
56
+ if key in state_dict_2_keys:
57
+ keys_in_both.append(key)
58
+
59
+ # sort them
60
+ keys_not_in_state_dict_2.sort()
61
+ keys_not_in_state_dict_1.sort()
62
+ keys_in_both.sort()
63
+
64
+
65
+ json_data = {
66
+ "both": keys_in_both,
67
+ "not_in_state_dict_2": keys_not_in_state_dict_2,
68
+ "not_in_state_dict_1": keys_not_in_state_dict_1
69
+ }
70
+ json_data = json.dumps(json_data, indent=4)
71
+
72
+ remaining_diffusers_values = OrderedDict()
73
+ for key in keys_not_in_state_dict_1:
74
+ remaining_diffusers_values[key] = state_dict_file_2[key]
75
+
76
+ # print(remaining_diffusers_values.keys())
77
+
78
+ remaining_ldm_values = OrderedDict()
79
+ for key in keys_not_in_state_dict_2:
80
+ remaining_ldm_values[key] = state_dict_file_1[key]
81
+
82
+ # print(json_data)
83
+
84
+ project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
85
+ json_save_path = os.path.join(project_root, 'config', 'keys.json')
86
+ json_matched_save_path = os.path.join(project_root, 'config', 'matched.json')
87
+ json_duped_save_path = os.path.join(project_root, 'config', 'duped.json')
88
+ state_dict_1_filename = os.path.basename(args.file_1[0])
89
+ state_dict_2_filename = os.path.basename(args.file_2[0])
90
+ # save key names for each in own file
91
+ with open(os.path.join(project_root, 'config', f'{state_dict_1_filename}.json'), 'w') as f:
92
+ f.write(json.dumps(state_dict_1_keys, indent=4))
93
+
94
+ with open(os.path.join(project_root, 'config', f'{state_dict_2_filename}.json'), 'w') as f:
95
+ f.write(json.dumps(state_dict_2_keys, indent=4))
96
+
97
+
98
+ with open(json_save_path, 'w') as f:
99
+ f.write(json_data)