(self, pretrained_dict=None, pretrained_layers=[], verbose=True)
| 573 | raise TypeError('pretrained must be a str or None') |
| 574 | |
| 575 | def load_weights(self, pretrained_dict=None, pretrained_layers=[], verbose=True): |
| 576 | model_dict = self.state_dict() |
| 577 | |
| 578 | missed_dict = [k for k in model_dict.keys() if k not in pretrained_dict] |
| 579 | logger.info(f'=> Missed keys {missed_dict}') |
| 580 | unexpected_dict = [k for k in pretrained_dict.keys() if k not in model_dict] |
| 581 | logger.info(f'=> Unexpected keys {unexpected_dict}') |
| 582 | |
| 583 | pretrained_dict = { |
| 584 | k: v for k, v in pretrained_dict.items() |
| 585 | if k in model_dict.keys() |
| 586 | } |
| 587 | |
| 588 | need_init_state_dict = {} |
| 589 | for k, v in pretrained_dict.items(): |
| 590 | need_init = ( |
| 591 | ( |
| 592 | k.split('.')[0] in pretrained_layers |
| 593 | or pretrained_layers[0] == '*' |
| 594 | ) |
| 595 | and 'relative_position_index' not in k |
| 596 | and 'attn_mask' not in k |
| 597 | ) |
| 598 | |
| 599 | if need_init: |
| 600 | # if verbose: |
| 601 | # logger.info(f'=> init {k} from {pretrained}') |
| 602 | |
| 603 | if ('pool_layers' in k) or ('focal_layers' in k) and v.size() != model_dict[k].size(): |
| 604 | table_pretrained = v |
| 605 | table_current = model_dict[k] |
| 606 | fsize1 = table_pretrained.shape[2] |
| 607 | fsize2 = table_current.shape[2] |
| 608 | |
| 609 | # NOTE: different from interpolation used in self-attention, we use padding or clipping for focal conv |
| 610 | if fsize1 < fsize2: |
| 611 | table_pretrained_resized = torch.zeros(table_current.shape) |
| 612 | table_pretrained_resized[:, :, (fsize2-fsize1)//2:-(fsize2-fsize1)//2, (fsize2-fsize1)//2:-(fsize2-fsize1)//2] = table_pretrained |
| 613 | v = table_pretrained_resized |
| 614 | elif fsize1 > fsize2: |
| 615 | table_pretrained_resized = table_pretrained[:, :, (fsize1-fsize2)//2:-(fsize1-fsize2)//2, (fsize1-fsize2)//2:-(fsize1-fsize2)//2] |
| 616 | v = table_pretrained_resized |
| 617 | |
| 618 | |
| 619 | if ("modulation.f" in k or "pre_conv" in k): |
| 620 | table_pretrained = v |
| 621 | table_current = model_dict[k] |
| 622 | if table_pretrained.shape != table_current.shape: |
| 623 | if len(table_pretrained.shape) == 2: |
| 624 | dim = table_pretrained.shape[1] |
| 625 | assert table_current.shape[1] == dim |
| 626 | L1 = table_pretrained.shape[0] |
| 627 | L2 = table_current.shape[0] |
| 628 | |
| 629 | if L1 < L2: |
| 630 | table_pretrained_resized = torch.zeros(table_current.shape) |
| 631 | # copy for linear project |
| 632 | table_pretrained_resized[:2*dim] = table_pretrained[:2*dim] |
no test coverage detected