File size: 3,284 Bytes
9b92c75
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
import torch
from test_model import tiny_config

from objectmodel_v1.losses import ObjectModelCriterion


def reference_dense_targets(criterion, logits, targets, level_index):
    batch, _, height, width = logits.shape
    target_logits = torch.zeros_like(logits)
    target_boxes = torch.zeros(batch, 4, height, width, device=logits.device)
    positive = torch.zeros(batch, height, width, device=logits.device)
    for batch_index, target in enumerate(targets):
        if target["labels"].numel() == 0:
            continue
        centers = target["boxes"][:, :2]
        grid = (centers * torch.tensor([width, height], device=centers.device)).long()
        grid[:, 0].clamp_(0, width - 1)
        grid[:, 1].clamp_(0, height - 1)
        areas = target["boxes"][:, 2] * target["boxes"][:, 3]
        target_levels = torch.where(areas < 0.02, 0, torch.where(areas < 0.15, 1, 2))
        for target_index in torch.where(target_levels == level_index)[0]:
            center_x, center_y = grid[target_index]
            label = target["labels"][target_index]
            candidates = []
            for dy in (-1, 0, 1):
                for dx in (-1, 0, 1):
                    x = int((center_x + dx).clamp(0, width - 1))
                    y = int((center_y + dy).clamp(0, height - 1))
                    candidates.append((dx * dx + dy * dy, x, y))
            for _, x, y in sorted(candidates)[: criterion.dense_topk]:
                target_logits[batch_index, label, y, x] = 1.0
                target_boxes[batch_index, :, y, x] = target["boxes"][target_index]
                positive[batch_index, y, x] = 1.0
    return target_logits, target_boxes, positive


def test_dense_targets_match_reference_with_boundaries_and_collisions():
    criterion = ObjectModelCriterion(tiny_config())
    logits = torch.randn(2, 5, 8, 8)
    targets = [
        {
            "boxes": torch.tensor(
                [
                    [0.01, 0.01, 0.10, 0.10],
                    [0.02, 0.02, 0.11, 0.11],
                    [0.99, 0.99, 0.50, 0.50],
                ]
            ),
            "labels": torch.tensor([1, 2, 3]),
        },
        {"boxes": torch.empty(0, 4), "labels": torch.empty(0, dtype=torch.long)},
    ]

    for level_index in range(3):
        expected = reference_dense_targets(criterion, logits, targets, level_index)
        actual = criterion._dense_targets(logits, targets, level_index)
        for expected_tensor, actual_tensor in zip(expected, actual, strict=True):
            assert torch.equal(expected_tensor, actual_tensor)


def test_dense_targets_match_reference_for_production_topk():
    config = tiny_config()
    config["loss"]["dense_topk"] = 5
    criterion = ObjectModelCriterion(config)
    logits = torch.randn(2, 5, 12, 10)
    targets = [
        {"boxes": torch.rand(9, 4), "labels": torch.randint(0, 5, (9,))},
        {"boxes": torch.rand(4, 4), "labels": torch.randint(0, 5, (4,))},
    ]

    for level_index in range(3):
        expected = reference_dense_targets(criterion, logits, targets, level_index)
        actual = criterion._dense_targets(logits, targets, level_index)
        for expected_tensor, actual_tensor in zip(expected, actual, strict=True):
            assert torch.equal(expected_tensor, actual_tensor)