TobiasLogic's picture
Upload ObjectModel-v1: code, README, assets (v1 checkpoint)
9b92c75 verified
Raw
History Blame Contribute Delete
3.63 kB
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)
]