| 31 | return x |
| 32 | |
| 33 | class Detector(nn.Module): |
| 34 | def __init__(self, num_classes, pre_trained=None, det_token_num=100, backbone_name='tiny', init_pe_size=[800,1344], mid_pe_size=None, use_checkpoint=False): |
| 35 | super().__init__() |
| 36 | # import pdb;pdb.set_trace() |
| 37 | if backbone_name == 'tiny': |
| 38 | self.backbone, hidden_dim = tiny(pretrained=pre_trained) |
| 39 | elif backbone_name == 'small': |
| 40 | self.backbone, hidden_dim = small(pretrained=pre_trained) |
| 41 | elif backbone_name == 'base': |
| 42 | self.backbone, hidden_dim = base(pretrained=pre_trained) |
| 43 | elif backbone_name == 'small_dWr': |
| 44 | self.backbone, hidden_dim = small_dWr(pretrained=pre_trained) |
| 45 | else: |
| 46 | raise ValueError(f'backbone {backbone_name} not supported') |
| 47 | |
| 48 | self.backbone.finetune_det(det_token_num=det_token_num, img_size=init_pe_size, mid_pe_size=mid_pe_size, use_checkpoint=use_checkpoint) |
| 49 | |
| 50 | self.class_embed = MLP(hidden_dim, hidden_dim, num_classes + 1, 3) |
| 51 | self.bbox_embed = MLP(hidden_dim, hidden_dim, 4, 3) |
| 52 | |
| 53 | def forward(self, samples: NestedTensor): |
| 54 | # import pdb;pdb.set_trace() |
| 55 | if isinstance(samples, (list, torch.Tensor)): |
| 56 | samples = nested_tensor_from_tensor_list(samples) |
| 57 | x = self.backbone(samples.tensors) |
| 58 | # x = x[:, 1:,:] |
| 59 | outputs_class = self.class_embed(x) |
| 60 | outputs_coord = self.bbox_embed(x).sigmoid() |
| 61 | out = {'pred_logits': outputs_class, 'pred_boxes': outputs_coord} |
| 62 | return out |
| 63 | |
| 64 | def forward_return_attention(self, samples: NestedTensor): |
| 65 | if isinstance(samples, (list, torch.Tensor)): |
| 66 | samples = nested_tensor_from_tensor_list(samples) |
| 67 | attention = self.backbone(samples.tensors, return_attention=True) |
| 68 | return attention |
| 69 | |
| 70 | class SetCriterion(nn.Module): |
| 71 | """ This class computes the loss for DETR. |
no outgoing calls
no test coverage detected