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

Method load_model

src/diffusers_model_pipeline.py:472–498  ·  view source on GitHub ↗
(self, save_path, compress=False)

Source from the content-addressed store, hash-verified

470 torch.save(delta_dict, save_path)
471
472 def load_model(self, save_path, compress=False):
473 st = torch.load(save_path)
474 if 'text_encoder' in st:
475 self.text_encoder.load_state_dict(st['text_encoder'])
476 if 'modifier_token' in st:
477 modifier_tokens = list(st['modifier_token'].keys())
478 modifier_token_id = []
479 for modifier_token in modifier_tokens:
480 num_added_tokens = self.tokenizer.add_tokens(modifier_token)
481 if num_added_tokens == 0:
482 raise ValueError(
483 f"The tokenizer already contains the token {modifier_token}. Please pass a different"
484 " `modifier_token` that is not already in the tokenizer."
485 )
486 modifier_token_id.append(self.tokenizer.convert_tokens_to_ids(modifier_token))
487 # Resize the token embeddings as we are adding new special tokens to the tokenizer
488 self.text_encoder.resize_token_embeddings(len(self.tokenizer))
489 token_embeds = self.text_encoder.get_input_embeddings().weight.data
490 for i, id_ in enumerate(modifier_token_id):
491 token_embeds[id_] = st['modifier_token'][modifier_tokens[i]]
492
493 for name, params in self.unet.named_parameters():
494 if 'attn2' in name:
495 if compress and ('to_k' in name or 'to_v' in name):
496 params.data += st['unet'][name]['u']@st['unet'][name]['v']
497 elif name in st['unet']:
498 params.data.copy_(st['unet'][f'{name}'])

Callers 2

sampleFunction · 0.80
convertFunction · 0.80

Calls

no outgoing calls

Tested by

no test coverage detected