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