| 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}']) |