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