objectmodel-v1 / tests /test_matching.py
TobiasLogic's picture
Upload ObjectModel-v1: code, README, assets (v1 checkpoint)
9b92c75 verified
Raw
History Blame Contribute Delete
1.51 kB
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])