torchmetrics-ext 0.1.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- torchmetrics_ext-0.1.0/PKG-INFO +7 -0
- torchmetrics_ext-0.1.0/README.md +2 -0
- torchmetrics_ext-0.1.0/setup.cfg +4 -0
- torchmetrics_ext-0.1.0/setup.py +13 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext/__init__.py +0 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext/visual_grounding/__init__.py +3 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext/visual_grounding/scanrefer.py +119 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext.egg-info/PKG-INFO +7 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext.egg-info/SOURCES.txt +10 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext.egg-info/dependency_links.txt +1 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext.egg-info/requires.txt +2 -0
- torchmetrics_ext-0.1.0/src/torchmetrics_ext.egg-info/top_level.txt +1 -0
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
from setuptools import find_packages, setup
|
|
2
|
+
|
|
3
|
+
setup(
|
|
4
|
+
name="torchmetrics_ext",
|
|
5
|
+
version="0.1.0",
|
|
6
|
+
author="Yiming Zhang",
|
|
7
|
+
packages=find_packages(where="src"),
|
|
8
|
+
package_dir={"": "src"},
|
|
9
|
+
download_url="https://github.com/eamonn-zh/torchmetrics_ext",
|
|
10
|
+
install_requires=[
|
|
11
|
+
"torchmetrics", "torch"
|
|
12
|
+
]
|
|
13
|
+
)
|
|
File without changes
|
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
from torchmetrics import Metric
|
|
3
|
+
from typing import Dict, Sequence
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class ScanReferMetric(Metric):
|
|
7
|
+
r"""
|
|
8
|
+
Compute the Acc@kIoU for the ScanRefer 3D visual grounding task.
|
|
9
|
+
Please refer to https://daveredrum.github.io/ScanRefer/ for more details.
|
|
10
|
+
|
|
11
|
+
Example:
|
|
12
|
+
>>> import torch
|
|
13
|
+
>>> from src.evaluation.scanrefer_metric import ScanReferMetric
|
|
14
|
+
>>> metric = ScanReferMetric()
|
|
15
|
+
>>> pred_aabbs = torch.rand(size=(3, 2, 3), dtype=torch.float32)
|
|
16
|
+
>>> gt_aabbs = torch.rand(size=(3, 2, 3), dtype=torch.float32)
|
|
17
|
+
>>> eval_types = ("unique", "multiple", "multiple")
|
|
18
|
+
>>> metric(pred_aabbs, gt_aabbs, eval_types)
|
|
19
|
+
|
|
20
|
+
"""
|
|
21
|
+
def __init__(self, *args, **kwargs):
|
|
22
|
+
super().__init__(*args, **kwargs)
|
|
23
|
+
self.eval_types_mapping = {"unique": 0, "multiple": 1}
|
|
24
|
+
self.add_state("unique_tp_thresh_25", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
25
|
+
self.add_state("unique_tp_thresh_50", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
26
|
+
self.add_state("multiple_tp_thresh_25", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
27
|
+
self.add_state("multiple_tp_thresh_50", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
28
|
+
self.add_state("all_tp_thresh_25", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
29
|
+
self.add_state("all_tp_thresh_50", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
30
|
+
self.add_state("unique_total", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
31
|
+
self.add_state("multiple_total", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
32
|
+
self.add_state("all_total", default=torch.tensor(0), dist_reduce_fx="sum")
|
|
33
|
+
|
|
34
|
+
@staticmethod
|
|
35
|
+
def _get_batch_aabb_pair_ious_optimized(batch_boxes_1_bound: torch.Tensor, batch_boxes_2_bound: torch.Tensor) -> torch.Tensor:
|
|
36
|
+
"""
|
|
37
|
+
:param batch_boxes_1_bound: a batch of axis-aligned bounding boxes (B, 2, 3)
|
|
38
|
+
:param batch_boxes_2_bound: a batch of axis-aligned bounding boxes (B, 2, 3)
|
|
39
|
+
:return: IoU values for each pair of axis-aligned bounding boxes (B, )
|
|
40
|
+
"""
|
|
41
|
+
# directly unpack the min and max without splitting
|
|
42
|
+
box_1_x_min, box_1_y_min, box_1_z_min = batch_boxes_1_bound[:, 0].unbind(dim=1)
|
|
43
|
+
box_1_x_max, box_1_y_max, box_1_z_max = batch_boxes_1_bound[:, 1].unbind(dim=1)
|
|
44
|
+
|
|
45
|
+
box_2_x_min, box_2_y_min, box_2_z_min = batch_boxes_2_bound[:, 0].unbind(dim=1)
|
|
46
|
+
box_2_x_max, box_2_y_max, box_2_z_max = batch_boxes_2_bound[:, 1].unbind(dim=1)
|
|
47
|
+
|
|
48
|
+
# calculate intersections directly
|
|
49
|
+
x_a = torch.maximum(box_1_x_min, box_2_x_min)
|
|
50
|
+
y_a = torch.maximum(box_1_y_min, box_2_y_min)
|
|
51
|
+
z_a = torch.maximum(box_1_z_min, box_2_z_min)
|
|
52
|
+
x_b = torch.minimum(box_1_x_max, box_2_x_max)
|
|
53
|
+
y_b = torch.minimum(box_1_y_max, box_2_y_max)
|
|
54
|
+
z_b = torch.minimum(box_1_z_max, box_2_z_max)
|
|
55
|
+
|
|
56
|
+
# simplify volume calculations
|
|
57
|
+
intersection_volume = torch.clamp((x_b - x_a), min=0) * torch.clamp((y_b - y_a), min=0) * torch.clamp(
|
|
58
|
+
(z_b - z_a), min=0
|
|
59
|
+
)
|
|
60
|
+
box_1_volume = (box_1_x_max - box_1_x_min) * (box_1_y_max - box_1_y_min) * (box_1_z_max - box_1_z_min)
|
|
61
|
+
box_2_volume = (box_2_x_max - box_2_x_min) * (box_2_y_max - box_2_y_min) * (box_2_z_max - box_2_z_min)
|
|
62
|
+
|
|
63
|
+
# IoU calculation with epsilon to prevent division by zero
|
|
64
|
+
ious = intersection_volume / (box_1_volume + box_2_volume - intersection_volume + torch.finfo(torch.float32).eps)
|
|
65
|
+
return ious.flatten()
|
|
66
|
+
|
|
67
|
+
def _convert_eval_types_to_idx(self, eval_types: Sequence[str], device) -> torch.Tensor:
|
|
68
|
+
eval_types_tensor = torch.empty(size=(len(eval_types), ), dtype=torch.bool, device=device)
|
|
69
|
+
for i, eval_type in enumerate(eval_types):
|
|
70
|
+
eval_types_tensor[i] = self.eval_types_mapping[eval_type]
|
|
71
|
+
return eval_types_tensor
|
|
72
|
+
|
|
73
|
+
def update(self, preds: torch.Tensor, targets: torch.Tensor, eval_types: Sequence[str]) -> None:
|
|
74
|
+
"""
|
|
75
|
+
:param preds: predicted axis-aligned bounding boxes (B, 2, 3)
|
|
76
|
+
:param targets: ground truth axis-aligned bounding boxes (B, 2, 3)
|
|
77
|
+
:param eval_types: a sequence of "unique" or "multiple" labels (B, )
|
|
78
|
+
"""
|
|
79
|
+
|
|
80
|
+
# check input sizes
|
|
81
|
+
if preds.shape != targets.shape or preds.shape[0] != len(eval_types):
|
|
82
|
+
raise ValueError("preds, targets and eval_types must have the same length")
|
|
83
|
+
|
|
84
|
+
# convert evaluation types to numerical values for convenience
|
|
85
|
+
eval_types_tensor = self._convert_eval_types_to_idx(eval_types, preds.device)
|
|
86
|
+
|
|
87
|
+
# calculate axis-aligned bounding boxes between predictions and GTs
|
|
88
|
+
ious = self._get_batch_aabb_pair_ious_optimized(preds, targets)
|
|
89
|
+
|
|
90
|
+
# count true positives above the IoU thresholds
|
|
91
|
+
tp_thresh_25_mask = ious >= 0.25
|
|
92
|
+
tp_thresh_50_mask = ious >= 0.50
|
|
93
|
+
|
|
94
|
+
eval_type_unique_mask = eval_types_tensor == self.eval_types_mapping["unique"]
|
|
95
|
+
eval_type_multiple_mask = eval_types_tensor == self.eval_types_mapping["multiple"]
|
|
96
|
+
|
|
97
|
+
# update metrics
|
|
98
|
+
self.all_total += targets.shape[0]
|
|
99
|
+
self.unique_total += torch.count_nonzero(eval_type_unique_mask)
|
|
100
|
+
self.multiple_total += torch.count_nonzero(eval_type_multiple_mask)
|
|
101
|
+
|
|
102
|
+
self.all_tp_thresh_25 += torch.count_nonzero(tp_thresh_25_mask)
|
|
103
|
+
self.all_tp_thresh_50 += torch.count_nonzero(tp_thresh_50_mask)
|
|
104
|
+
|
|
105
|
+
self.unique_tp_thresh_25 += torch.count_nonzero(tp_thresh_25_mask & eval_type_unique_mask)
|
|
106
|
+
self.unique_tp_thresh_50 += torch.count_nonzero(tp_thresh_50_mask & eval_type_unique_mask)
|
|
107
|
+
|
|
108
|
+
self.multiple_tp_thresh_25 += torch.count_nonzero(tp_thresh_25_mask & eval_type_multiple_mask)
|
|
109
|
+
self.multiple_tp_thresh_50 += torch.count_nonzero(tp_thresh_50_mask & eval_type_multiple_mask)
|
|
110
|
+
|
|
111
|
+
def compute(self) -> Dict[str, torch.Tensor]:
|
|
112
|
+
return {
|
|
113
|
+
"unique_0.25": self.unique_tp_thresh_25 / self.unique_total,
|
|
114
|
+
"unique_0.5": self.unique_tp_thresh_50 / self.unique_total,
|
|
115
|
+
"multiple_0.25": self.multiple_tp_thresh_25 / self.multiple_total,
|
|
116
|
+
"multiple_0.5": self.multiple_tp_thresh_50 / self.multiple_total,
|
|
117
|
+
"all_0.25": self.all_tp_thresh_25 / self.all_total,
|
|
118
|
+
"all_0.5": self.all_tp_thresh_50 / self.all_total
|
|
119
|
+
}
|
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
README.md
|
|
2
|
+
setup.py
|
|
3
|
+
src/torchmetrics_ext/__init__.py
|
|
4
|
+
src/torchmetrics_ext.egg-info/PKG-INFO
|
|
5
|
+
src/torchmetrics_ext.egg-info/SOURCES.txt
|
|
6
|
+
src/torchmetrics_ext.egg-info/dependency_links.txt
|
|
7
|
+
src/torchmetrics_ext.egg-info/requires.txt
|
|
8
|
+
src/torchmetrics_ext.egg-info/top_level.txt
|
|
9
|
+
src/torchmetrics_ext/visual_grounding/__init__.py
|
|
10
|
+
src/torchmetrics_ext/visual_grounding/scanrefer.py
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
torchmetrics_ext
|