MCPcopy Create free account
hub / github.com/Monalissaa/DisenDiff / load_model_from_config_addtoken

Function load_model_from_config_addtoken

src/convert.py:27–48  ·  view source on GitHub ↗
(config, ckpt, verbose=False)

Source from the content-addressed store, hash-verified

25
26
27def 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
51def convert(ckpt, delta_ckpt, sd_version, config, modelname, mode):

Callers 1

convertFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected