Input: - src_logits: bs, num_dn, num_classes - tgt_labels: bs, num_dn
(self, outputs, targets, indices, idx, num_boxes,
data_batch)
| 786 | |
| 787 | @torch.no_grad() |
| 788 | def loss_matching_cost(self, outputs, targets, indices, idx, num_boxes, |
| 789 | data_batch): |
| 790 | """ |
| 791 | Input: |
| 792 | - src_logits: bs, num_dn, num_classes |
| 793 | - tgt_labels: bs, num_dn |
| 794 | |
| 795 | """ |
| 796 | cost_mean_dict = indices[1] |
| 797 | losses = {'set_{}'.format(k): v for k, v in cost_mean_dict.items()} |
| 798 | return losses |
| 799 | |
| 800 | def _get_src_permutation_idx(self, indices): |
| 801 | # permute predictions following indices |