File size: 1,040 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 | from __future__ import annotations
import torch
from torch import Tensor
from .boxes import box_cxcywh_to_xyxy
@torch.no_grad()
def decode_predictions(
outputs: dict[str, Tensor],
image_sizes: list[tuple[int, int]],
confidence: float = 0.25,
top_k: int = 300,
) -> list[dict[str, Tensor]]:
logits = outputs["pred_logits"].sigmoid()
boxes = box_cxcywh_to_xyxy(outputs["pred_boxes"]).clamp(0.0, 1.0)
results = []
for index, (height, width) in enumerate(image_sizes):
scores, labels = logits[index].max(dim=-1)
keep = scores >= confidence
if keep.sum() > top_k:
selected = scores.masked_fill(~keep, -1).topk(top_k).indices
else:
selected = torch.where(keep)[0]
selected_boxes = boxes[index, selected].clone()
selected_boxes[:, [0, 2]] *= width
selected_boxes[:, [1, 3]] *= height
results.append(
{"scores": scores[selected], "labels": labels[selected], "boxes": selected_boxes}
)
return results
|