| 40 | self.tracker = BYTETracker() |
| 41 | |
| 42 | def forward_train(self, |
| 43 | img, |
| 44 | img_metas, |
| 45 | gt_bboxes, |
| 46 | gt_labels, |
| 47 | gt_match_indices, |
| 48 | ref_img, |
| 49 | ref_img_metas, |
| 50 | ref_gt_bboxes, |
| 51 | ref_gt_labels, |
| 52 | ref_gt_match_indices, |
| 53 | gt_bboxes_ignore=None, |
| 54 | gt_masks=None, |
| 55 | ref_gt_bboxes_ignore=None, |
| 56 | ref_gt_masks=None, |
| 57 | **kwargs): |
| 58 | x = self.extract_feat(img) |
| 59 | |
| 60 | losses = dict() |
| 61 | |
| 62 | # RPN forward and loss |
| 63 | proposal_cfg = self.train_cfg.get('rpn_proposal', self.test_cfg.rpn) |
| 64 | rpn_losses, proposal_list = self.rpn_head.forward_train( |
| 65 | x, |
| 66 | img_metas, |
| 67 | gt_bboxes, |
| 68 | gt_labels=None, |
| 69 | gt_bboxes_ignore=gt_bboxes_ignore, |
| 70 | proposal_cfg=proposal_cfg) |
| 71 | losses.update(rpn_losses) |
| 72 | |
| 73 | ref_x = self.extract_feat(ref_img) |
| 74 | ref_proposals = self.rpn_head.simple_test_rpn(ref_x, ref_img_metas) |
| 75 | |
| 76 | roi_losses = self.roi_head.forward_train( |
| 77 | x, img_metas, proposal_list, gt_bboxes, gt_labels, |
| 78 | gt_match_indices, ref_x, ref_img_metas, ref_proposals, |
| 79 | ref_gt_bboxes, ref_gt_labels, gt_bboxes_ignore, gt_masks, |
| 80 | ref_gt_bboxes_ignore, **kwargs) |
| 81 | losses.update(roi_losses) |
| 82 | |
| 83 | return losses |
| 84 | |
| 85 | def simple_test(self, img, img_metas, rescale=False): |
| 86 | # TODO inherit from a base tracker |