task1 / src /pmdm /verifier.py
siddhant20's picture
Add files using upload-large-folder tool
d667566 verified
Raw History Blame Contribute Delete
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)}
@torch.no_grad()
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()