Input: - src_logits: bs, num_dn, num_classes - tgt_labels: bs, num_dn
(self, outputs, targets, indices, idx, num_boxes,
data_batch)
| 1641 | |
| 1642 | @torch.no_grad() |
| 1643 | def loss_matching_cost(self, outputs, targets, indices, idx, num_boxes, |
| 1644 | data_batch): |
| 1645 | """ |
| 1646 | Input: |
| 1647 | - src_logits: bs, num_dn, num_classes |
| 1648 | - tgt_labels: bs, num_dn |
| 1649 | |
| 1650 | """ |
| 1651 | cost_mean_dict = indices[1] |
| 1652 | losses = {'set_{}'.format(k): v for k, v in cost_mean_dict.items()} |
| 1653 | return losses |
| 1654 | |
| 1655 | def _get_src_permutation_idx(self, indices): |
| 1656 | # permute predictions following indices |