(self, pretrained_dict=None, pretrained_layers=[], verbose=True)
| 661 | |
| 662 | |
| 663 | def load_weights(self, pretrained_dict=None, pretrained_layers=[], verbose=True): |
| 664 | model_dict = self.state_dict() |
| 665 | pretrained_dict = { |
| 666 | k: v for k, v in pretrained_dict.items() |
| 667 | if k in model_dict.keys() |
| 668 | } |
| 669 | need_init_state_dict = {} |
| 670 | for k, v in pretrained_dict.items(): |
| 671 | need_init = ( |
| 672 | ( |
| 673 | k.split('.')[0] in pretrained_layers |
| 674 | or pretrained_layers[0] == '*' |
| 675 | ) |
| 676 | and 'relative_position_index' not in k |
| 677 | and 'attn_mask' not in k |
| 678 | ) |
| 679 | |
| 680 | if need_init: |
| 681 | # if verbose: |
| 682 | # logger.info(f'=> init {k} from {pretrained}') |
| 683 | |
| 684 | if 'relative_position_bias_table' in k and v.size() != model_dict[k].size(): |
| 685 | relative_position_bias_table_pretrained = v |
| 686 | relative_position_bias_table_current = model_dict[k] |
| 687 | L1, nH1 = relative_position_bias_table_pretrained.size() |
| 688 | L2, nH2 = relative_position_bias_table_current.size() |
| 689 | if nH1 != nH2: |
| 690 | logger.info(f"Error in loading {k}, passing") |
| 691 | else: |
| 692 | if L1 != L2: |
| 693 | logger.info( |
| 694 | '=> load_pretrained: resized variant: {} to {}' |
| 695 | .format((L1, nH1), (L2, nH2)) |
| 696 | ) |
| 697 | S1 = int(L1 ** 0.5) |
| 698 | S2 = int(L2 ** 0.5) |
| 699 | relative_position_bias_table_pretrained_resized = torch.nn.functional.interpolate( |
| 700 | relative_position_bias_table_pretrained.permute(1, 0).view(1, nH1, S1, S1), |
| 701 | size=(S2, S2), |
| 702 | mode='bicubic') |
| 703 | v = relative_position_bias_table_pretrained_resized.view(nH2, L2).permute(1, 0) |
| 704 | |
| 705 | if 'absolute_pos_embed' in k and v.size() != model_dict[k].size(): |
| 706 | absolute_pos_embed_pretrained = v |
| 707 | absolute_pos_embed_current = model_dict[k] |
| 708 | _, L1, C1 = absolute_pos_embed_pretrained.size() |
| 709 | _, L2, C2 = absolute_pos_embed_current.size() |
| 710 | if C1 != C1: |
| 711 | logger.info(f"Error in loading {k}, passing") |
| 712 | else: |
| 713 | if L1 != L2: |
| 714 | logger.info( |
| 715 | '=> load_pretrained: resized variant: {} to {}' |
| 716 | .format((1, L1, C1), (1, L2, C2)) |
| 717 | ) |
| 718 | S1 = int(L1 ** 0.5) |
| 719 | S2 = int(L2 ** 0.5) |
| 720 | absolute_pos_embed_pretrained = absolute_pos_embed_pretrained.reshape(-1, S1, S1, C1) |
no test coverage detected