MCPcopy Create free account
hub / github.com/UX-Decoder/Semantic-SAM / load_weights

Method load_weights

semantic_sam/backbone/focal_dw.py:575–660  ·  view source on GitHub ↗
(self, pretrained_dict=None, pretrained_layers=[], verbose=True)

Source from the content-addressed store, hash-verified

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]

Callers 1

get_focal_backboneFunction · 0.45

Calls 1

itemsMethod · 0.80

Tested by

no test coverage detected