Spaces:
Running on Zero
Running on Zero
| from __future__ import annotations | |
| import torch | |
| from torch import Tensor | |
| from .boxes import box_cxcywh_to_xyxy | |
| 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 | |