| 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 |
|
|