Download src/pmdm/decode.py from siddhant20/task1: direct link, hf CLI and curl.
- Browser
- Download file 5.66 kB
-
https://huggingface.co/siddhant20/task1/resolve/main/src/pmdm/decode.py
- Command line
-
hf download hf://siddhant20/task1/src/pmdm/decode.py
-
curl -L -o decode.py https://huggingface.co/siddhant20/task1/resolve/main/src/pmdm/decode.py
5.66 kB
| """Heatmap decoding, box fusion and the metric-aware post-processing. | |
| Two dataset facts are exploited here: | |
| * no ground-truth box is a pure deletion (0 of 1404), so a candidate whose | |
| template side has ink and whose photo side does not is dropped; | |
| * the ink-blob box sits within (0, 0, +1, +1) of the ground-truth box, so a | |
| constant margin is added after snapping. | |
| """ | |
| from __future__ import annotations | |
| import cv2 | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from .config import BOX_MARGIN, MAX_BOXES_PER_IMAGE, OUT_STRIDE, TOPK_PER_TILE, WBF_IOU | |
| from .metric import iou_matrix | |
| def decode_heatmap(hm_logits: torch.Tensor, wh: torch.Tensor, off: torch.Tensor, | |
| score_thr: float = 0.05, topk: int = TOPK_PER_TILE, | |
| stride: int = OUT_STRIDE): | |
| """Single-image decode. Returns (boxes (N,4) in tile pixels, scores (N,)).""" | |
| hm = torch.sigmoid(hm_logits) | |
| keep = (F.max_pool2d(hm, 3, stride=1, padding=1) == hm).float() | |
| hm = hm * keep | |
| scores, flat = hm.reshape(-1).topk(min(topk, hm.numel())) | |
| size = hm.shape[-1] | |
| ys = (flat // size).float() | |
| xs = (flat % size).float() | |
| off_flat = off.reshape(2, -1)[:, flat] | |
| wh_flat = wh.reshape(2, -1)[:, flat] | |
| cx = (xs + off_flat[0]) * stride | |
| cy = (ys + off_flat[1]) * stride | |
| w, h = wh_flat[0].clamp(min=1.0), wh_flat[1].clamp(min=1.0) | |
| boxes = torch.stack([cx - w / 2, cy - h / 2, cx + w / 2, cy + h / 2], 1) | |
| m = scores >= score_thr | |
| return boxes[m].cpu().numpy(), scores[m].cpu().numpy() | |
| def wbf(boxes: np.ndarray, scores: np.ndarray, iou_thr: float = WBF_IOU): | |
| """Weighted Boxes Fusion. Averages coordinates instead of discarding them, | |
| which matters at IoU 0.5 on 8x8 boxes.""" | |
| if len(boxes) == 0: | |
| return boxes, scores | |
| order = np.argsort(-scores) | |
| boxes, scores = boxes[order], scores[order] | |
| clusters: list[list[int]] = [] | |
| fused: list[np.ndarray] = [] | |
| fused_scores: list[float] = [] | |
| for i in range(len(boxes)): | |
| placed = False | |
| if fused: | |
| ious = iou_matrix(boxes[i:i + 1], np.asarray(fused, np.float32))[0] | |
| j = int(np.argmax(ious)) | |
| if ious[j] >= iou_thr: | |
| clusters[j].append(i) | |
| members = clusters[j] | |
| w = scores[members] | |
| fused[j] = (boxes[members] * w[:, None]).sum(0) / w.sum() | |
| fused_scores[j] = float(w.sum() / min(len(members) + 1, 3)) | |
| placed = True | |
| if not placed: | |
| clusters.append([i]) | |
| fused.append(boxes[i].copy()) | |
| fused_scores.append(float(scores[i])) | |
| f = np.asarray(fused, np.float32) | |
| s = np.clip(np.asarray(fused_scores, np.float32), 0, 1) | |
| order = np.argsort(-s) | |
| return f[order], s[order] | |
| def polarity_filter(boxes: np.ndarray, scores: np.ndarray, tn: np.ndarray, pn: np.ndarray, | |
| ink_thr: float = 3.0): | |
| """Drop candidates that look like a deletion (template inked, photo blank).""" | |
| if len(boxes) == 0: | |
| return boxes, scores | |
| keep = np.ones(len(boxes), bool) | |
| for i, (x1, y1, x2, y2) in enumerate(boxes.astype(int)): | |
| x1, y1 = max(0, x1), max(0, y1) | |
| x2, y2 = min(tn.shape[1], x2), min(tn.shape[0], y2) | |
| if x2 <= x1 or y2 <= y1: | |
| keep[i] = False | |
| continue | |
| t_ink = float(tn[y1:y2, x1:x2].mean()) | |
| p_ink = float(pn[y1:y2, x1:x2].mean()) | |
| if t_ink > ink_thr and p_ink < ink_thr: | |
| keep[i] = False | |
| return boxes[keep], scores[keep] | |
| def snap_to_ink(boxes: np.ndarray, tn: np.ndarray, pn: np.ndarray, pad: int = 6, | |
| diff_thr: float = 25.0, margin: tuple[int, int, int, int] = BOX_MARGIN, | |
| max_shift: int = 6) -> np.ndarray: | |
| """Refit each box to the local ink-difference blob, then apply the measured margin.""" | |
| if len(boxes) == 0: | |
| return boxes | |
| diff = np.abs(pn.astype(np.float32) - tn.astype(np.float32)) | |
| out = boxes.copy() | |
| h, w = diff.shape | |
| for i, (x1, y1, x2, y2) in enumerate(boxes): | |
| a, b = int(max(0, y1 - pad)), int(min(h, y2 + pad)) | |
| c, d = int(max(0, x1 - pad)), int(min(w, x2 + pad)) | |
| if b <= a or d <= c: | |
| continue | |
| m = diff[a:b, c:d] > diff_thr | |
| if m.sum() < 3: | |
| continue | |
| ys, xs = np.nonzero(m) | |
| nx1, ny1 = xs.min() + c, ys.min() + a | |
| nx2, ny2 = xs.max() + 1 + c, ys.max() + 1 + a | |
| cand = np.array([nx1 - margin[0], ny1 - margin[1], nx2 + margin[2], ny2 + margin[3]], | |
| np.float32) | |
| if np.abs(cand - boxes[i]).max() <= max_shift: | |
| out[i] = cand | |
| return out | |
| def clip_boxes(boxes: np.ndarray, h: int, w: int) -> np.ndarray: | |
| if len(boxes) == 0: | |
| return boxes | |
| b = boxes.copy() | |
| b[:, [0, 2]] = np.clip(b[:, [0, 2]], 0, w) | |
| b[:, [1, 3]] = np.clip(b[:, [1, 3]], 0, h) | |
| return b | |
| def postprocess(boxes: np.ndarray, scores: np.ndarray, tn: np.ndarray, pn: np.ndarray, | |
| use_snap: bool = True, use_polarity: bool = True): | |
| if len(boxes) == 0: | |
| return boxes, scores | |
| h, w = tn.shape[:2] | |
| boxes = clip_boxes(boxes, h, w) | |
| if use_polarity: | |
| boxes, scores = polarity_filter(boxes, scores, tn, pn) | |
| if use_snap and len(boxes): | |
| boxes = snap_to_ink(boxes, tn, pn) | |
| keep = (boxes[:, 2] - boxes[:, 0] > 2) & (boxes[:, 3] - boxes[:, 1] > 2) | |
| boxes, scores = boxes[keep], scores[keep] | |
| if len(boxes) > MAX_BOXES_PER_IMAGE: | |
| order = np.argsort(-scores)[:MAX_BOXES_PER_IMAGE] | |
| boxes, scores = boxes[order], scores[order] | |
| return boxes, scores | |