Download src/pmdm/verifier.py from siddhant20/task1: direct link, hf CLI and curl.
- Browser
- Download file 9.04 kB
-
https://huggingface.co/siddhant20/task1/resolve/main/src/pmdm/verifier.py
- Command line
-
hf download hf://siddhant20/task1/src/pmdm/verifier.py
-
curl -L -o verifier.py https://huggingface.co/siddhant20/task1/resolve/main/src/pmdm/verifier.py
9.04 kB
| """Stage 2: patch verifier and box refiner. | |
| Stage 1 is tuned for recall; most of its output is print-artefact noise. This | |
| model looks at a single candidate at high resolution and answers two questions: | |
| is this a real content difference, and where exactly are its corners. The second | |
| answer is what converts a 0.4-IoU candidate into a true positive at the 0.5 | |
| threshold, which matters because the median box is 22 px wide. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import timm | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader, Dataset | |
| from .config import CKPT, PREP | |
| from .metric import iou_matrix | |
| from .preprocess import load_prepared | |
| CROP = 96 | |
| CONTEXT = 2.5 # crop side = CONTEXT * max(box side), clamped | |
| POS_IOU = 0.35 | |
| NEG_IOU = 0.20 | |
| def crop_window(box: np.ndarray, h: int, w: int) -> tuple[int, int, int, int]: | |
| x1, y1, x2, y2 = box | |
| cx, cy = (x1 + x2) / 2, (y1 + y2) / 2 | |
| side = float(np.clip(max(x2 - x1, y2 - y1) * CONTEXT, 32, 256)) | |
| x0 = int(round(np.clip(cx - side / 2, 0, max(0, w - side)))) | |
| y0 = int(round(np.clip(cy - side / 2, 0, max(0, h - side)))) | |
| return x0, y0, int(round(side)), int(round(side)) | |
| def make_crop(t, p, tn, pn, box: np.ndarray): | |
| """8-channel crop plus the box in normalized crop coordinates.""" | |
| import cv2 | |
| h, w = t.shape[:2] | |
| x0, y0, side, _ = crop_window(box, h, w) | |
| sl = (slice(y0, y0 + side), slice(x0, x0 + side)) | |
| stack = np.concatenate([t[sl], p[sl], tn[sl][..., None], pn[sl][..., None]], -1) | |
| stack = cv2.resize(stack, (CROP, CROP), interpolation=cv2.INTER_LINEAR) | |
| scale = CROP / side | |
| local = np.array([(box[0] - x0) * scale, (box[1] - y0) * scale, | |
| (box[2] - x0) * scale, (box[3] - y0) * scale], np.float32) | |
| return stack.astype(np.float32).transpose(2, 0, 1) / 255.0, local, (x0, y0, side) | |
| class Verifier(nn.Module): | |
| def __init__(self, backbone: str = "resnet18", pretrained: bool = True): | |
| super().__init__() | |
| self.backbone = timm.create_model(backbone, pretrained=pretrained, in_chans=8, | |
| num_classes=0) | |
| feat = self.backbone.num_features | |
| self.head_cls = nn.Linear(feat + 4, 1) | |
| self.head_box = nn.Linear(feat + 4, 4) | |
| def forward(self, x: torch.Tensor, box: torch.Tensor): | |
| f = torch.cat([self.backbone(x), box / CROP], 1) | |
| return self.head_cls(f).squeeze(1), self.head_box(f) | |
| class CandidateDataset(Dataset): | |
| """Candidates produced by stage 1, labelled against ground truth.""" | |
| def __init__(self, records: list[dict], prep_root: Path | None = None, jitter: bool = True): | |
| self.records = records | |
| self.prep_root = prep_root or PREP | |
| self.jitter = jitter | |
| self._cache: dict = {} | |
| def __len__(self) -> int: | |
| return len(self.records) | |
| def _pair(self, split: str, idx: int): | |
| key = (split, idx) | |
| if key not in self._cache: | |
| if len(self._cache) > 8: | |
| self._cache.clear() | |
| self._cache[key] = load_prepared(split, idx, self.prep_root) | |
| return self._cache[key] | |
| def __getitem__(self, i: int): | |
| r = self.records[i] | |
| t, p, tn, pn = self._pair(r["split"], r["idx"]) | |
| box = np.asarray(r["box"], np.float32) | |
| if self.jitter: | |
| box = box + np.random.uniform(-2, 2, 4).astype(np.float32) | |
| stack, local, _ = make_crop(t, p, tn, pn, box) | |
| target = np.zeros(4, np.float32) | |
| if r["label"] > 0: | |
| gt = np.asarray(r["gt"], np.float32) | |
| x0, y0, side, _ = crop_window(box, t.shape[0], t.shape[1]) | |
| scale = CROP / side | |
| gt_local = np.array([(gt[0] - x0) * scale, (gt[1] - y0) * scale, | |
| (gt[2] - x0) * scale, (gt[3] - y0) * scale], np.float32) | |
| target = gt_local - local | |
| return { | |
| "x": torch.from_numpy(stack), | |
| "box": torch.from_numpy(local), | |
| "label": torch.tensor(float(r["label"])), | |
| "delta": torch.from_numpy(target), | |
| } | |
| def build_records(candidates: dict, gt: dict) -> list[dict]: | |
| """candidates: key -> (split, idx, boxes, scores); gt: key -> boxes.""" | |
| records = [] | |
| for key, (split, idx, boxes, _scores) in candidates.items(): | |
| g = gt.get(key, np.zeros((0, 4), np.float32)) | |
| ious = iou_matrix(boxes, g) if len(boxes) and len(g) else np.zeros((len(boxes), len(g))) | |
| for i in range(len(boxes)): | |
| best = float(ious[i].max()) if ious.size else 0.0 | |
| j = int(ious[i].argmax()) if ious.size else -1 | |
| if best >= POS_IOU: | |
| records.append({"split": split, "idx": idx, "box": boxes[i].tolist(), | |
| "label": 1, "gt": g[j].tolist()}) | |
| elif best < NEG_IOU: | |
| records.append({"split": split, "idx": idx, "box": boxes[i].tolist(), | |
| "label": 0, "gt": boxes[i].tolist()}) | |
| return records | |
| def train_verifier(records: list[dict], epochs: int = 8, batch: int = 64, lr: float = 3e-4, | |
| device: str = "cuda", out_dir: Path | None = None) -> dict: | |
| out_dir = Path(out_dir) if out_dir is not None else Path(CKPT) / "stage2" | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| (out_dir / "records.json").write_text(json.dumps(records[:50000])) | |
| print(f"[verifier] {len(records)} records, " | |
| f"{sum(r['label'] for r in records)} positive", flush=True) | |
| ds = CandidateDataset(records) | |
| batch = min(batch, max(1, len(records))) | |
| drop_last = len(records) >= 2 * batch # never leave the loader empty | |
| dl = DataLoader(ds, batch_size=batch, shuffle=True, num_workers=4, drop_last=drop_last) | |
| if len(dl) == 0: | |
| raise ValueError(f"verifier loader is empty for {len(records)} records") | |
| model = Verifier().to(device) | |
| opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) | |
| sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, total_steps=epochs * len(dl), | |
| pct_start=0.1) | |
| n_pos = sum(r["label"] for r in records) | |
| pos_weight = torch.tensor([max(1.0, (len(records) - n_pos) / max(1, n_pos))]).to(device) | |
| for epoch in range(epochs): | |
| running = 0.0 | |
| for step, b in enumerate(dl): | |
| x = b["x"].to(device) | |
| box = b["box"].to(device) | |
| logit, delta = model(x, box) | |
| loss_cls = F.binary_cross_entropy_with_logits(logit, b["label"].to(device), | |
| pos_weight=pos_weight) | |
| mask = b["label"].to(device) > 0 | |
| loss_box = (F.smooth_l1_loss(delta[mask], b["delta"].to(device)[mask]) | |
| if mask.any() else delta.sum() * 0) | |
| loss = loss_cls + 0.5 * loss_box | |
| opt.zero_grad(set_to_none=True) | |
| loss.backward() | |
| opt.step() | |
| sched.step() | |
| running += float(loss.detach()) | |
| if step % 100 == 0: | |
| print(f"[verifier] epoch {epoch} step {step}/{len(dl)} " | |
| f"loss={running / (step + 1):.4f}", flush=True) | |
| torch.save({"model": model.state_dict(), "epoch": epoch}, out_dir / "last.pt") | |
| return {"records": len(records), "positives": int(n_pos)} | |
| def apply_verifier(model, split: str, idx: int, boxes: np.ndarray, scores: np.ndarray, | |
| device: str = "cuda", prep_root: Path | None = None, batch: int = 128, | |
| refine: bool = True, blend: float = 0.5): | |
| if len(boxes) == 0: | |
| return boxes, scores | |
| t, p, tn, pn = load_prepared(split, idx, prep_root or PREP) | |
| crops, locals_, windows = [], [], [] | |
| for box in boxes: | |
| stack, local, win = make_crop(t, p, tn, pn, box) | |
| crops.append(stack) | |
| locals_.append(local) | |
| windows.append(win) | |
| out_scores, out_boxes = [], [] | |
| for i in range(0, len(crops), batch): | |
| x = torch.from_numpy(np.stack(crops[i:i + batch])).to(device) | |
| b = torch.from_numpy(np.stack(locals_[i:i + batch])).to(device) | |
| logit, delta = model(x, b) | |
| probs = torch.sigmoid(logit).cpu().numpy() | |
| new_local = (b + delta).cpu().numpy() | |
| for k, prob in enumerate(probs): | |
| x0, y0, side = windows[i + k] | |
| scale = side / CROP | |
| nb = new_local[k] * scale + np.array([x0, y0, x0, y0], np.float32) | |
| out_boxes.append(nb if refine else boxes[i + k]) | |
| out_scores.append(prob) | |
| fused = (np.asarray(out_scores, np.float32) ** blend) * (scores ** (1 - blend)) | |
| return np.asarray(out_boxes, np.float32), fused | |
| def load_verifier(path: Path, device: str = "cuda"): | |
| model = Verifier(pretrained=False) | |
| state = torch.load(path, map_location="cpu") | |
| model.load_state_dict(state["model"] if "model" in state else state) | |
| return model.to(device).eval() | |