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