| 25 | |
| 26 | |
| 27 | def load_model_from_config_addtoken(config, ckpt, verbose=False): |
| 28 | print(f"Loading model from {ckpt}") |
| 29 | pl_sd = torch.load(ckpt, map_location="cpu") |
| 30 | if "global_step" in pl_sd: |
| 31 | print(f"Global Step: {pl_sd['global_step']}") |
| 32 | sd = pl_sd["state_dict"] |
| 33 | model = instantiate_from_config(config.model) |
| 34 | |
| 35 | token_weights = sd["cond_stage_model.transformer.text_model.embeddings.token_embedding.weight"] |
| 36 | del sd["cond_stage_model.transformer.text_model.embeddings.token_embedding.weight"] |
| 37 | m, u = model.load_state_dict(sd, strict=False) |
| 38 | model.cond_stage_model.transformer.text_model.embeddings.token_embedding.weight.data[:token_weights.shape[0]] = token_weights |
| 39 | if len(m) > 0 and verbose: |
| 40 | print("missing keys:") |
| 41 | print(m) |
| 42 | if len(u) > 0 and verbose: |
| 43 | print("unexpected keys:") |
| 44 | print(u) |
| 45 | |
| 46 | model.cuda() |
| 47 | model.eval() |
| 48 | return model |
| 49 | |
| 50 | |
| 51 | def convert(ckpt, delta_ckpt, sd_version, config, modelname, mode): |