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])