first commit
This commit is contained in:
27
tests/test_model_forward.py
Normal file
27
tests/test_model_forward.py
Normal file
@@ -0,0 +1,27 @@
|
||||
import torch
|
||||
|
||||
from models.laneformer import LaneFormer
|
||||
|
||||
|
||||
def test_forward_shapes():
|
||||
torch.manual_seed(0)
|
||||
model = LaneFormer(pretrained=False, max_lanes=8)
|
||||
x = torch.randn(2, 3, 360, 640)
|
||||
out = model(x)
|
||||
assert out["cls_logits"].shape == (2, 8, 2)
|
||||
assert out["curve_coeffs"].shape == (2, 8, 4)
|
||||
assert out["endpoints"].shape == (2, 8, 2)
|
||||
assert torch.isfinite(out["cls_logits"]).all()
|
||||
assert torch.isfinite(out["curve_coeffs"]).all()
|
||||
assert (out["endpoints"] >= 0).all() and (out["endpoints"] <= 1).all()
|
||||
|
||||
|
||||
def test_backward_runs():
|
||||
torch.manual_seed(0)
|
||||
model = LaneFormer(pretrained=False, max_lanes=8)
|
||||
x = torch.randn(1, 3, 360, 640)
|
||||
out = model(x)
|
||||
loss = out["cls_logits"].sum() + out["curve_coeffs"].sum() + out["endpoints"].sum()
|
||||
loss.backward()
|
||||
grads = [p.grad for p in model.parameters() if p.requires_grad]
|
||||
assert any(g is not None and torch.isfinite(g).all() for g in grads)
|
||||
Reference in New Issue
Block a user