(ckpt, delta_ckpt, sd_version, config, modelname, mode)
| 49 | |
| 50 | |
| 51 | def 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()): |
no test coverage detected