File size: 5,664 Bytes
d667566
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
"""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