"""COCO val2017 mAP eval for split-tower detection heads on frozen EUPE-ViT-B. Default output is hard NMS (the standard reported in the FCOS / YOLO / DETR literature). Pass `--soft-nms` to additionally report a soft-NMS comparison pass on the same forward outputs. Auto-detects architecture flags from the checkpoint's state_dict: hidden width, n_std / n_dw tower depths, per_scale_bias, reg_on_raw. Usage: python eval_coco_map.py --picker python eval_coco_map.py --picker --soft-nms """ import argparse import os import re import sys import time os.environ.setdefault("COCO_TEXT_EMBED_PATH", "/mnt/d/detection-heads/_diag_text_embed_vitb32.pt") import torch sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from train_split_tower_5scale import SplitTowerHead, make_locations, STRIDES DEVICE = "cuda" COCO_ROOT = os.environ["ARENA_COCO_ROOT"] VAL_CACHE = os.environ["ARENA_VAL_CACHE"] RESOLUTION = 640 H = RESOLUTION // 16 NUM_CLASSES = 80 SCORE_THRESH = 0.05 MAX_PER_IMAGE = 100 def hard_nms_per_class(boxes, scores, labels, iou_thresh=0.5): """Per-class hard NMS via torchvision.ops.nms. Returns indices into input.""" from torchvision.ops import nms keep = [] for c in labels.unique(): m = labels == c idx = m.nonzero(as_tuple=True)[0] k = nms(boxes[m], scores[m], iou_thresh) keep.append(idx[k]) return torch.cat(keep) if keep else torch.tensor([], dtype=torch.long, device=boxes.device) def soft_nms_linear(boxes, scores, iou_thresh=0.5, min_score=0.001): """GPU-vectorized linear-decay soft NMS (Bodla et al. 2017). Iteratively picks the max-score box, multiplicatively decays overlapping scores by (1 - IoU), repeats until the max falls below min_score.""" if boxes.numel() == 0: return torch.tensor([], dtype=torch.long, device=boxes.device), scores.clone() N = boxes.shape[0] scores = scores.clone() areas = (boxes[:, 2] - boxes[:, 0]).clamp(min=0) * (boxes[:, 3] - boxes[:, 1]).clamp(min=0) active = torch.ones(N, dtype=torch.bool, device=boxes.device) kept_idx, kept_scores = [], [] while active.any(): masked = torch.where(active, scores, torch.full_like(scores, -float("inf"))) i = int(masked.argmax()) if scores[i] < min_score: break kept_idx.append(i); kept_scores.append(scores[i].item()) active[i] = False box_i = boxes[i] x1 = torch.maximum(box_i[0], boxes[:, 0]); y1 = torch.maximum(box_i[1], boxes[:, 1]) x2 = torch.minimum(box_i[2], boxes[:, 2]); y2 = torch.minimum(box_i[3], boxes[:, 3]) inter = (x2 - x1).clamp(min=0) * (y2 - y1).clamp(min=0) iou = inter / (areas[i] + areas - inter).clamp(min=1e-6) decay = torch.where(iou > iou_thresh, 1 - iou, torch.ones_like(iou)) scores = scores * decay if not kept_idx: return (torch.tensor([], dtype=torch.long, device=boxes.device), torch.tensor([], device=boxes.device)) return (torch.tensor(kept_idx, dtype=torch.long, device=boxes.device), torch.tensor(kept_scores, device=boxes.device)) def soft_nms_per_class(boxes, scores, labels, iou_thresh=0.5): keep_idx, keep_scores = [], [] for c in labels.unique(): m = labels == c idx = m.nonzero(as_tuple=True)[0] k, rescored = soft_nms_linear(boxes[m], scores[m], iou_thresh) keep_idx.append(idx[k]); keep_scores.append(rescored) if not keep_idx: return (torch.tensor([], dtype=torch.long, device=boxes.device), torch.tensor([], device=boxes.device)) return torch.cat(keep_idx), torch.cat(keep_scores) def autodetect_and_load(ckpt_path): ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False) sd = ckpt["head"] if isinstance(ckpt, dict) and "head" in ckpt else ckpt per_scale_bias = (sd["cls_bias"].shape == (5, 80)) reg_on_raw = any("reg_stem" in k for k in sd.keys()) hidden = sd["stem.weight"].shape[0] tower_idxs = sorted({int(m.group(1)) for k in sd if (m := re.match(r"cls_tower\.(\d+)\.", k))}) n_std = sum(1 for i in tower_idxs if f"cls_tower.{i}.conv.weight" in sd) n_dw = len(tower_idxs) - n_std return sd, hidden, n_std, n_dw, per_scale_bias, reg_on_raw def main(): parser = argparse.ArgumentParser() parser.add_argument("--picker", required=True, help="Path to picker .pth checkpoint") parser.add_argument("--soft-nms", action="store_true", help="Also run a soft-NMS pass and report both. Default is hard NMS only.") args = parser.parse_args() t_start = time.time() print(f"[{time.time()-t_start:6.1f}s] START pid={os.getpid()}", flush=True) print(f"[{time.time()-t_start:6.1f}s] picker={args.picker} soft_nms={args.soft_nms}", flush=True) print(f"[{time.time()-t_start:6.1f}s] Inspecting checkpoint ...", flush=True) sd, hidden, n_std, n_dw, per_scale_bias, reg_on_raw = autodetect_and_load(args.picker) print(f"[{time.time()-t_start:6.1f}s] Arch: hidden={hidden} n_std={n_std} n_dw={n_dw} " f"per_scale_bias={per_scale_bias} reg_on_raw={reg_on_raw}", flush=True) print(f"[{time.time()-t_start:6.1f}s] Constructing head ...", flush=True) head = SplitTowerHead(hidden=hidden, n_std_layers=n_std, n_dw_layers=n_dw, n_scales=4, reg_on_raw=reg_on_raw, per_scale_bias=per_scale_bias).to(DEVICE) head.load_state_dict(sd); head.eval() print(f"[{time.time()-t_start:6.1f}s] Head loaded: {sum(p.numel() for p in head.parameters()):,} params", flush=True) print(f"[{time.time()-t_start:6.1f}s] Loading val cache ...", flush=True) val = torch.load(VAL_CACHE, map_location="cpu", weights_only=False) from pycocotools.coco import COCO from pycocotools.cocoeval import COCOeval coco_gt = COCO(os.path.join(COCO_ROOT, "annotations", "instances_val2017.json")) cat_ids = sorted(coco_gt.getCatIds()) idx_to_cat = {i: c for i, c in enumerate(cat_ids)} print(f"[{time.time()-t_start:6.1f}s] Val: {len(val)} items, {len(cat_ids)} categories", flush=True) feat_sizes = [(H*2, H*2), (H, H), (H//2, H//2), (H//4, H//4), (H//8, H//8)] all_locs = torch.cat(make_locations(feat_sizes, STRIDES, torch.device(DEVICE))) print(f"[{time.time()-t_start:6.1f}s] ==================== INFERENCE START ====================", flush=True) results_hard = [] results_soft = [] if args.soft_nms else None t0 = time.time() PRINT_EVERY = 500 with torch.no_grad(): for idx, item in enumerate(val): spatial = item["spatial"].unsqueeze(0).float().to(DEVICE) img_id = int(item["img_id"]); scale = item["scale"] cls_l, reg_l, ctr_l = head(spatial) cls_s = torch.cat([c.permute(0,2,3,1).reshape(-1, NUM_CLASSES) for c in cls_l]).sigmoid() reg_s = torch.cat([r.permute(0,2,3,1).reshape(-1, 4) for r in reg_l]) ctr_s = torch.cat([c.permute(0,2,3,1).reshape(-1) for c in ctr_l]).sigmoid() scores_all = cls_s * ctr_s.unsqueeze(1) mask = scores_all > SCORE_THRESH if not mask.any(): if (idx+1) % PRINT_EVERY == 0: elapsed = time.time()-t0; rate = (idx+1)/elapsed print(f"[{time.time()-t_start:6.1f}s] img {idx+1}/{len(val)} | " f"hard={len(results_hard):,}" + (f" soft={len(results_soft):,}" if args.soft_nms else "") + f" | {rate:.1f} img/s | ETA {(len(val)-idx-1)/rate:.0f}s", flush=True) continue loc_idx, cls_idx = mask.nonzero(as_tuple=True) sc = scores_all[loc_idx, cls_idx] xy = all_locs[loc_idx]; rr = reg_s[loc_idx] x1 = (xy[:,0]-rr[:,0]).clamp(0, RESOLUTION); y1 = (xy[:,1]-rr[:,1]).clamp(0, RESOLUTION) x2 = (xy[:,0]+rr[:,2]).clamp(0, RESOLUTION); y2 = (xy[:,1]+rr[:,3]).clamp(0, RESOLUTION) boxes = torch.stack([x1, y1, x2, y2], dim=-1) keep_hard = hard_nms_per_class(boxes, sc, cls_idx, iou_thresh=0.5) _append_dets(results_hard, boxes, sc[keep_hard], keep_hard, cls_idx, scale, idx_to_cat, img_id) if args.soft_nms: keep_soft, sc_soft = soft_nms_per_class(boxes, sc, cls_idx, iou_thresh=0.5) _append_dets(results_soft, boxes, sc_soft, keep_soft, cls_idx, scale, idx_to_cat, img_id) if (idx+1) % PRINT_EVERY == 0: elapsed = time.time()-t0; rate = (idx+1)/elapsed msg = (f"[{time.time()-t_start:6.1f}s] img {idx+1}/{len(val)} | " f"hard={len(results_hard):,}") if args.soft_nms: msg += f" soft={len(results_soft):,}" msg += f" | n_pos={mask.sum().item()} | {rate:.1f} img/s | ETA {(len(val)-idx-1)/rate:.0f}s" print(msg, flush=True) print(f"\n[{time.time()-t_start:6.1f}s] Inference done in {time.time()-t0:.0f}s", flush=True) passes = [("HARD NMS", results_hard)] if args.soft_nms: passes.append(("SOFT NMS", results_soft)) for label, results in passes: print(f"\n{'='*40}\n{label}\n{'='*40}", flush=True) coco_dt = coco_gt.loadRes(results) ev = COCOeval(coco_gt, coco_dt, "bbox") ev.params.imgIds = sorted(coco_gt.getImgIds())[:len(val)] ev.evaluate(); ev.accumulate(); ev.summarize() def _append_dets(sink, boxes, scores_kept, keep, cls_idx, scale, idx_to_cat, img_id): if keep.numel() == 0: return if scores_kept.numel() > MAX_PER_IMAGE: top = scores_kept.topk(MAX_PER_IMAGE) scores_kept = top.values; keep = keep[top.indices] bx = boxes[keep] / scale w = (bx[:,2]-bx[:,0]).clamp(min=0); h = (bx[:,3]-bx[:,1]).clamp(min=0) for i in range(keep.numel()): s = scores_kept[i].item() if s < SCORE_THRESH: continue sink.append({"image_id": img_id, "category_id": idx_to_cat[cls_idx[keep[i]].item()], "bbox": [bx[i,0].item(), bx[i,1].item(), w[i].item(), h[i].item()], "score": s}) if __name__ == "__main__": main()