| 429 | |
| 430 | |
| 431 | def inference(self, |
| 432 | retina_box_cls, retina_box_delta, retina_anchors, |
| 433 | small_det_logits, small_det_delta, small_det_anchors, |
| 434 | image_sizes |
| 435 | ): |
| 436 | results = [] |
| 437 | |
| 438 | N, _, _, _ = retina_box_cls[0].size() |
| 439 | retina_box_cls = [permute_to_N_HWA_K(x, self.num_classes) for x in retina_box_cls] |
| 440 | retina_box_delta = [permute_to_N_HWA_K(x, 4) for x in retina_box_delta] |
| 441 | small_det_logits = [x.view(N, -1, self.num_classes) for x in small_det_logits] |
| 442 | small_det_delta = [x.view(N, -1, 4) for x in small_det_delta] |
| 443 | |
| 444 | for img_idx, image_size in enumerate(image_sizes): |
| 445 | |
| 446 | retina_box_cls_per_image = [box_cls_per_level[img_idx] for box_cls_per_level in retina_box_cls] |
| 447 | retina_box_reg_per_image = [box_reg_per_level[img_idx] for box_reg_per_level in retina_box_delta] |
| 448 | small_det_logits_per_image = [small_det_cls_per_level[img_idx] for small_det_cls_per_level in small_det_logits] |
| 449 | small_det_reg_per_image = [small_det_reg_per_level[img_idx] for small_det_reg_per_level in small_det_delta] |
| 450 | |
| 451 | if len(small_det_anchors) == 0 or type(small_det_anchors[0]) == torch.Tensor: |
| 452 | small_det_anchor_per_image = [small_det_anchor_per_level[img_idx] for small_det_anchor_per_level in small_det_anchors] |
| 453 | else: |
| 454 | small_det_anchor_per_image = small_det_anchors |
| 455 | |
| 456 | results_per_img = self.inference_single_image( |
| 457 | retina_box_cls_per_image, retina_box_reg_per_image, retina_anchors, |
| 458 | small_det_logits_per_image, small_det_reg_per_image, small_det_anchor_per_image, |
| 459 | tuple(image_size)) |
| 460 | results.append(results_per_img) |
| 461 | |
| 462 | return results |
| 463 | |
| 464 | |
| 465 | def inference_single_image(self, |