| import torch | |
| from objectmodel_v1.matching import hungarian_match, hungarian_match_layers | |
| def test_layer_matcher_matches_individual_calls(): | |
| torch.manual_seed(7) | |
| outputs = [ | |
| { | |
| "pred_logits": torch.randn(3, 12, 5), | |
| "pred_boxes": torch.rand(3, 12, 4), | |
| } | |
| for _ in range(3) | |
| ] | |
| targets = [ | |
| {"labels": torch.tensor([1, 3]), "boxes": torch.rand(2, 4)}, | |
| {"labels": torch.empty(0, dtype=torch.long), "boxes": torch.empty(0, 4)}, | |
| {"labels": torch.tensor([0, 2, 4]), "boxes": torch.rand(3, 4)}, | |
| ] | |
| expected = [hungarian_match(output, targets) for output in outputs] | |
| actual = hungarian_match_layers(outputs, targets) | |
| for expected_layer, actual_layer in zip(expected, actual, strict=True): | |
| for expected_match, actual_match in zip(expected_layer, actual_layer, strict=True): | |
| assert torch.equal(expected_match[0], actual_match[0]) | |
| assert torch.equal(expected_match[1], actual_match[1]) | |
| def test_layer_matcher_handles_all_empty_targets(): | |
| outputs = [{"pred_logits": torch.randn(2, 4, 3), "pred_boxes": torch.rand(2, 4, 4)}] | |
| targets = [ | |
| {"labels": torch.empty(0, dtype=torch.long), "boxes": torch.empty(0, 4)}, | |
| {"labels": torch.empty(0, dtype=torch.long), "boxes": torch.empty(0, 4)}, | |
| ] | |
| matches = hungarian_match_layers(outputs, targets) | |
| assert len(matches) == 1 | |
| assert all(rows.numel() == 0 and cols.numel() == 0 for rows, cols in matches[0]) | |