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