MCPcopy Create free account
hub / github.com/adobe-research/custom-diffusion / convert

Function convert

src/convert.py:51–138  ·  view source on GitHub ↗
(ckpt, delta_ckpt, sd_version, config, modelname, mode)

Source from the content-addressed store, hash-verified

49
50
51def convert(ckpt, delta_ckpt, sd_version, config, modelname, mode):
52 config = OmegaConf.load(config)
53 model = load_model_from_config(config, f"{ckpt}")
54 # get the mapping of layer names between diffuser and CompVis checkpoints
55 mapping_compvis_to_diffuser = {}
56 mapping_compvis_to_diffuser_rev = {}
57 for key in list(model.state_dict().keys()):
58 if 'attn2' in key:
59 diffuser_key = key.replace('model.diffusion_model.', '')
60 if 'input_blocks' in key:
61 i, j = [int(x) for x in key.split('.')[3:5]]
62 i_, j_ = max(0, i // 3), 0 if i in [1, 4, 7] else 1
63 diffuser_key = diffuser_key.replace(f'input_blocks.{i}.{j}', f'down_blocks.{i_}.attentions.{j_}')
64 if 'output_blocks' in key:
65 i, j = [int(x) for x in key.split('.')[3:5]]
66 i_, j_ = max(0, i // 3), 0 if i % 3 == 0 else 1 if i % 3 == 1 else 2
67 diffuser_key = diffuser_key.replace(f'output_blocks.{i}.{j}', f'up_blocks.{i_}.attentions.{j_}')
68 diffuser_key = diffuser_key.replace('middle_block.1', 'mid_block.attentions.0')
69 mapping_compvis_to_diffuser[key] = diffuser_key
70 mapping_compvis_to_diffuser_rev[diffuser_key] = key
71
72 # convert checkpoint to webui
73 if mode in ['diffuser-to-webui' or 'compvis-to-webui']:
74 outpath = f'{os.path.dirname(delta_ckpt)}/webui'
75 os.makedirs(outpath, exist_ok=True)
76 if mode == 'diffuser-to-webui':
77 st = torch.load(delta_ckpt)
78 compvis_st = {}
79 compvis_st['state_dict'] = {}
80 for key in list(st['unet'].keys()):
81 compvis_st['state_dict'][mapping_compvis_to_diffuser_rev[key]] = st['unet'][key]
82
83 model.load_state_dict(compvis_st['state_dict'], strict=False)
84 torch.save({'state_dict': model.state_dict()}, f'{outpath}/{modelname}')
85
86 if 'modifier_token' in st:
87 os.makedirs(f'{outpath}/embeddings/', exist_ok=True)
88 for word, feat in st['modifier_token'].items():
89 torch.save({word: feat}, f'{outpath}/embeddings/{word}.pt')
90 else:
91 compvis_st = torch.load(delta_ckpt)["state_dict"]
92 model.load_state_dict(compvis_st['state_dict'], strict=False)
93 torch.save({'state_dict': model.state_dict()}, f'{outpath}/{modelname}')
94
95 if 'embed' in st:
96 os.makedirs(f'{outpath}/embeddings/', exist_ok=True)
97 for i, feat in enumerate(st['embed']):
98 torch.save({f'<new{i}>': feat}, f'{outpath}/embeddings/<new{i}>.pt')
99 # convert checkpoint from CompVis to diffuser
100 elif mode == 'compvis-to-diffuser':
101 st = torch.load(delta_ckpt)["state_dict"]
102 diffuser_st = {'unet': {}}
103 if 'embed' in st:
104 diffuser_st['modifier_token'] = {}
105 for i in range(st['embed'].size(0)):
106 diffuser_st['modifier_token'][f'<new{i+1}>'] = st['embed'][i].clone()
107 del st['embed']
108 for key in list(st.keys()):

Callers 1

convert.pyFile · 0.85

Calls 4

load_model_from_configFunction · 0.70
load_modelMethod · 0.45
save_pretrainedMethod · 0.45

Tested by

no test coverage detected