| 7 | |
| 8 | |
| 9 | def 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 | |
| 43 | def parse_args(): |