first commit
This commit is contained in:
96
engine/evaluate.py
Normal file
96
engine/evaluate.py
Normal file
@@ -0,0 +1,96 @@
|
||||
"""TuSimple-style accuracy/FP/FN evaluation (paper Eq. 10).
|
||||
|
||||
Simplification flagged explicitly: the official TuSimple metric uses a 25px
|
||||
tolerance defined in original-resolution pixel space with per-clip point sampling.
|
||||
We don't retain clip structure (our data is single frames, not clips) or the
|
||||
per-sample letterbox scale at collation time, so we evaluate in normalized [0,1]
|
||||
canonical (640-wide) space with an equivalent threshold (25/640 ~= 0.039). This is a
|
||||
consistent proxy metric for comparing our own checkpoints, not a pixel-exact
|
||||
reproduction of the official script -- consistent with the val-split caveat already
|
||||
flagged in the project plan (no official test_label.json available).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from scipy.optimize import linear_sum_assignment
|
||||
|
||||
from models.head import eval_curve
|
||||
|
||||
X_THRESHOLD_NORM = 25.0 / 640.0
|
||||
CLS_PROB_THRESHOLD = 0.5
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def evaluate(model, dataloader, device, sample_ys: torch.Tensor) -> dict:
|
||||
model.eval()
|
||||
total_correct_points = 0
|
||||
total_gt_points = 0
|
||||
total_fp_lanes = 0
|
||||
total_fn_lanes = 0
|
||||
total_pred_lanes = 0
|
||||
total_gt_lanes = 0
|
||||
|
||||
for batch in dataloader:
|
||||
images = batch["images"].to(device)
|
||||
target_xs = batch["target_xs"].to(device)
|
||||
target_valid_mask = batch["target_valid_mask"].to(device)
|
||||
target_lane_valid = batch["target_lane_valid"].to(device)
|
||||
sys_dev = sample_ys.to(device)
|
||||
|
||||
out = model(images)
|
||||
probs = F.softmax(out["cls_logits"], dim=-1)[..., 1] # (B, N)
|
||||
pred_xs_all = eval_curve(out["curve_coeffs"], sys_dev.view(1, 1, -1).expand(*probs.shape, -1)) # (B, N, S)
|
||||
|
||||
B = images.shape[0]
|
||||
for b in range(B):
|
||||
pos_idx = (probs[b] > CLS_PROB_THRESHOLD).nonzero(as_tuple=True)[0]
|
||||
gt_idx = target_lane_valid[b].nonzero(as_tuple=True)[0]
|
||||
|
||||
total_pred_lanes += pos_idx.numel()
|
||||
total_gt_lanes += gt_idx.numel()
|
||||
|
||||
if gt_idx.numel() == 0:
|
||||
total_fp_lanes += pos_idx.numel()
|
||||
continue
|
||||
if pos_idx.numel() == 0:
|
||||
total_fn_lanes += gt_idx.numel()
|
||||
gt_mask = target_valid_mask[b, gt_idx]
|
||||
total_gt_points += int(gt_mask.sum().item())
|
||||
continue
|
||||
|
||||
pred_xs = pred_xs_all[b, pos_idx] # (P, S)
|
||||
gt_xs = target_xs[b, gt_idx] # (M, S)
|
||||
gt_mask = target_valid_mask[b, gt_idx] # (M, S)
|
||||
|
||||
diff = (pred_xs.unsqueeze(1) - gt_xs.unsqueeze(0)).abs() # (P, M, S)
|
||||
mask = gt_mask.unsqueeze(0).float()
|
||||
denom = mask.sum(dim=-1).clamp(min=1.0)
|
||||
cost = (diff * mask).sum(dim=-1) / denom # (P, M)
|
||||
|
||||
pred_local, gt_local = linear_sum_assignment(cost.cpu().numpy())
|
||||
|
||||
matched_pred = set(pred_local.tolist())
|
||||
matched_gt = set(gt_local.tolist())
|
||||
total_fp_lanes += pos_idx.numel() - len(matched_pred)
|
||||
total_fn_lanes += gt_idx.numel() - len(matched_gt)
|
||||
|
||||
for p_local, g_local in zip(pred_local, gt_local):
|
||||
m = gt_mask[g_local]
|
||||
if not m.any():
|
||||
continue
|
||||
correct = (diff[p_local, g_local][m] < X_THRESHOLD_NORM).sum().item()
|
||||
total_correct_points += correct
|
||||
total_gt_points += int(m.sum().item())
|
||||
|
||||
unmatched_gt = set(range(gt_idx.numel())) - matched_gt
|
||||
for g_local in unmatched_gt:
|
||||
m = gt_mask[g_local]
|
||||
total_gt_points += int(m.sum().item())
|
||||
|
||||
accuracy = total_correct_points / total_gt_points if total_gt_points > 0 else 0.0
|
||||
fp = total_fp_lanes / total_pred_lanes if total_pred_lanes > 0 else 0.0
|
||||
fn = total_fn_lanes / total_gt_lanes if total_gt_lanes > 0 else 0.0
|
||||
|
||||
model.train()
|
||||
return {"accuracy": accuracy, "fp": fp, "fn": fn}
|
||||
Reference in New Issue
Block a user