File size: 3,626 Bytes
9b92c75
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

from collections.abc import Sequence

import numpy as np
import torch
from scipy.optimize import linear_sum_assignment
from torch import Tensor

from .boxes import box_cxcywh_to_xyxy, generalized_box_iou_batched


def hungarian_match(
    outputs: dict[str, Tensor],
    targets: Sequence[dict[str, Tensor]],
    class_cost: float = 2.0,
    bbox_cost: float = 5.0,
    giou_cost: float = 2.0,
) -> list[tuple[Tensor, Tensor]]:
    """Match predictions to targets with a detached CPU Hungarian solve."""
    return hungarian_match_layers([outputs], targets, class_cost, bbox_cost, giou_cost)[0]


def hungarian_match_layers(
    outputs: Sequence[dict[str, Tensor]],
    targets: Sequence[dict[str, Tensor]],
    class_cost: float = 2.0,
    bbox_cost: float = 5.0,
    giou_cost: float = 2.0,
) -> list[list[tuple[Tensor, Tensor]]]:
    """Match all decoder layers with one GPU-to-CPU synchronization."""
    if not outputs:
        return []

    device = outputs[0]["pred_boxes"].device
    target_sizes = [len(target["labels"]) for target in targets]
    max_targets = max(target_sizes, default=0)
    if max_targets == 0:
        empty = torch.empty(0, dtype=torch.int64, device=device)
        return [[(empty, empty) for _ in targets] for _ in outputs]

    batch_size = len(targets)
    padded_boxes = torch.zeros(batch_size, max_targets, 4, dtype=torch.float32, device=device)
    padded_labels = torch.zeros(batch_size, max_targets, dtype=torch.int64, device=device)
    for batch_index, target in enumerate(targets):
        count = target_sizes[batch_index]
        if count > 0:
            padded_boxes[batch_index, :count] = target["boxes"]
            padded_labels[batch_index, :count] = target["labels"]

    padded_costs = []
    for layer_output in outputs:
        probabilities = layer_output["pred_logits"].sigmoid().detach()
        boxes = layer_output["pred_boxes"].detach()
        num_queries = boxes.shape[1]
        gather_index = padded_labels[:, None, :].expand(-1, num_queries, -1)
        class_term = -probabilities.gather(2, gather_index)
        bbox_term = torch.cdist(boxes, padded_boxes, p=1)
        giou_term = -generalized_box_iou_batched(
            box_cxcywh_to_xyxy(boxes), box_cxcywh_to_xyxy(padded_boxes)
        )
        cost = (class_cost * class_term + bbox_cost * bbox_term + giou_cost * giou_term).float()
        for batch_index in range(batch_size):
            padded_costs.append(cost[batch_index])

    costs = torch.stack(padded_costs).cpu().numpy()

    cpu_matches: list[tuple[np.ndarray, np.ndarray]] = []
    for cost, target_size in zip(costs, target_sizes * len(outputs), strict=True):
        if target_size == 0:
            empty = np.empty(0, dtype=np.int64)
            cpu_matches.append((empty, empty))
        else:
            cpu_matches.append(linear_sum_assignment(cost[:, :target_size]))

    lengths = [len(rows) for rows, _ in cpu_matches]
    all_rows = torch.as_tensor(
        np.concatenate([rows for rows, _ in cpu_matches]),
        dtype=torch.int64,
        device=device,
    )
    all_cols = torch.as_tensor(
        np.concatenate([cols for _, cols in cpu_matches]),
        dtype=torch.int64,
        device=device,
    )
    flat_matches = []
    offset = 0
    for length in lengths:
        flat_matches.append(
            (all_rows[offset : offset + length], all_cols[offset : offset + length])
        )
        offset += length

    batch_size = len(targets)
    return [
        flat_matches[start : start + batch_size]
        for start in range(0, len(flat_matches), batch_size)
    ]