File size: 2,863 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
77
78
79
80
81
82
83
84
import torch

from objectmodel_v1.losses import ObjectModelCriterion
from objectmodel_v1.model import build_model
from objectmodel_v1.postprocess import decode_predictions


def tiny_config(dense_aux=True):
    return {
        "model": {
            "num_classes": 5,
            "input_size": 128,
            "stem_channels": 16,
            "backbone_channels": [24, 32, 48, 64],
            "backbone_depths": [1, 1, 1, 1],
            "hidden_dim": 48,
            "fpn_depth": 1,
            "latent_count": 8,
            "latent_pool_sizes": [4, 2, 1],
            "latent_layers": 1,
            "decoder_layers": 2,
            "num_queries": 12,
            "num_heads": 4,
            "local_points": 2,
            "dropout": 0.0,
            "dense_aux": dense_aux,
        },
        "loss": {
            "cost_class": 2.0,
            "cost_bbox": 5.0,
            "cost_giou": 2.0,
            "weight_class": 2.0,
            "weight_bbox": 5.0,
            "weight_giou": 2.0,
            "weight_dense": 1.0,
            "focal_alpha": 0.25,
            "focal_gamma": 2.0,
            "aux_weight": 1.0,
            "dense_topk": 3,
        },
    }


def test_forward_shapes_and_ranges():
    model = build_model(tiny_config()).eval()
    with torch.no_grad():
        output = model(torch.randn(2, 3, 128, 128))
    assert output["pred_logits"].shape == (2, 12, 5)
    assert output["pred_boxes"].shape == (2, 12, 4)
    assert len(output["aux_outputs"]) == 1
    assert "dense_outputs" not in output
    assert torch.all((output["pred_boxes"] >= 0) & (output["pred_boxes"] <= 1))


def test_loss_backward_with_empty_target():
    config = tiny_config()
    model = build_model(config).train()
    criterion = ObjectModelCriterion(config)
    output = model(torch.randn(2, 3, 128, 128))
    targets = [
        {
            "boxes": torch.tensor([[0.5, 0.5, 0.25, 0.3], [0.2, 0.2, 0.1, 0.1]]),
            "labels": torch.tensor([1, 3]),
        },
        {"boxes": torch.empty(0, 4), "labels": torch.empty(0, dtype=torch.long)},
    ]
    losses = criterion(output, targets)
    assert all(torch.isfinite(value) for value in losses.values())
    losses["loss_total"].backward()
    gradients = [parameter.grad for parameter in model.parameters() if parameter.grad is not None]
    assert gradients
    assert all(torch.isfinite(gradient).all() for gradient in gradients)


def test_decode_is_nms_free_top_k_filter():
    model = build_model(tiny_config(dense_aux=False)).eval()
    with torch.no_grad():
        output = model(torch.randn(1, 3, 128, 128))
    result = decode_predictions(output, [(240, 320)], confidence=0.0, top_k=4)[0]
    assert result["boxes"].shape == (4, 4)
    assert result["scores"].shape == (4,)
    assert torch.all(result["boxes"][:, [0, 2]] <= 320)
    assert torch.all(result["boxes"][:, [1, 3]] <= 240)