Spaces:
Paused
Paused
Download scripts/train/tool_add_control_sd21.py from NguyenDinhHieu/EquiFashion: direct link, hf CLI and curl.
- Browser
- Download file 1.56 kB
-
https://huggingface.co/spaces/NguyenDinhHieu/EquiFashion/resolve/256371375472d2e0578e963e93eded913dd47de0/scripts/train/tool_add_control_sd21.py
- Command line
-
hf download hf://spaces/NguyenDinhHieu/EquiFashion@256371375472d2e0578e963e93eded913dd47de0/scripts/train/tool_add_control_sd21.py
-
curl -L -o tool_add_control_sd21.py https://huggingface.co/spaces/NguyenDinhHieu/EquiFashion/resolve/256371375472d2e0578e963e93eded913dd47de0/scripts/train/tool_add_control_sd21.py
1.56 kB
| import sys | |
| import os | |
| assert len(sys.argv) == 3, 'Args are wrong.' | |
| input_path = sys.argv[1] | |
| output_path = sys.argv[2] | |
| assert os.path.exists(input_path), 'Input model does not exist.' | |
| assert not os.path.exists(output_path), 'Output filename already exists.' | |
| assert os.path.exists(os.path.dirname(output_path)), 'Output path is not valid.' | |
| import torch | |
| # from share import * | |
| from cldm.model import create_model | |
| """ | |
| python tool_add_control_sd21.py /data/lh/docker/models/my_train/dress_code_30epochs.ckpt /data/lh/docker/models/my_train/control_dress_code_ini.ckpt | |
| """ | |
| def get_node_name(name, parent_name): | |
| if len(name) <= len(parent_name): | |
| return False, '' | |
| p = name[:len(parent_name)] | |
| if p != parent_name: | |
| return False, '' | |
| return True, name[len(parent_name):] | |
| model = create_model(config_path='./configs/control_sd/cldm_v2_orig.yaml') | |
| pretrained_weights = torch.load(input_path) | |
| if 'state_dict' in pretrained_weights: | |
| pretrained_weights = pretrained_weights['state_dict'] | |
| scratch_dict = model.state_dict() | |
| target_dict = {} | |
| for k in scratch_dict.keys(): | |
| is_control, name = get_node_name(k, 'control_') | |
| if is_control: | |
| copy_k = 'model.diffusion_' + name | |
| else: | |
| copy_k = k | |
| if copy_k in pretrained_weights: | |
| target_dict[k] = pretrained_weights[copy_k].clone() | |
| else: | |
| target_dict[k] = scratch_dict[k].clone() | |
| print(f'These weights are newly added: {k}') | |
| model.load_state_dict(target_dict, strict=True) | |
| torch.save(model.state_dict(), output_path) | |
| print('Done.') | |