MCPcopy Create free account
hub / github.com/Monalissaa/DisenDiff / main

Function main

src/get_deltas.py:9–40  ·  view source on GitHub ↗
(path, newtoken=0)

Source from the content-addressed store, hash-verified

7
8
9def main(path, newtoken=0):
10 layers = []
11 for files in glob.glob(f'{path}/checkpoints/*'):
12 if ('=' in files or '_' in files) and 'delta' not in files:
13 print(files)
14 if '=' in files:
15 epoch_number = files.split('=')[1].split('.ckpt')[0]
16 elif '_' in files:
17 epoch_number = files.split('/')[-1].split('.ckpt')[0]
18
19 st = torch.load(files)["state_dict"]
20 if len(layers) == 0:
21 for key in list(st.keys()):
22 # if 'attn2.to_k' in key or 'attn2.to_v' in key:
23 if 'attn2.to_k' in key or 'attn2.to_v' in key or 'attn2.to_q' in key:
24 layers.append(key)
25 print(layers)
26 st_delta = {'state_dict': {}}
27 for each in layers:
28 st_delta['state_dict'][each] = st[each].clone()
29 print('/'.join(files.split('/')[:-1]) + f'/delta_epoch={epoch_number}.ckpt')
30
31 num_tokens = st['cond_stage_model.transformer.text_model.embeddings.token_embedding.weight'].shape[0]
32
33 if newtoken > 0:
34 print("saving the optimized embedding")
35 st_delta['state_dict']['embed'] = st['cond_stage_model.transformer.text_model.embeddings.token_embedding.weight'][-newtoken:].clone()
36 print(st_delta['state_dict']['embed'].shape, num_tokens)
37
38
39 torch.save(st_delta, '/'.join(files.split('/')[:-1]) + f'/delta_epoch={epoch_number}.ckpt')
40 os.remove(files)
41
42
43def parse_args():

Callers 1

get_deltas.pyFile · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected