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.
@@ -0,0 +1,7 @@
1
+ Metadata-Version: 2.1
2
+ Name: torchmetrics_ext
3
+ Version: 0.1.0
4
+ Download-URL: https://github.com/eamonn-zh/torchmetrics_ext
5
+ Author: Yiming Zhang
6
+ Requires-Dist: torchmetrics
7
+ Requires-Dist: torch
@@ -0,0 +1,2 @@
1
+ # TorchMetrics Extensions
2
+ Extensions of TorchMetrics
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -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
+ )
@@ -0,0 +1,3 @@
1
+ from torchmetrics_ext.visual_grounding.scanrefer import ScanReferMetric
2
+
3
+ __all__ = ["ScanReferMetric"]
@@ -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,7 @@
1
+ Metadata-Version: 2.1
2
+ Name: torchmetrics-ext
3
+ Version: 0.1.0
4
+ Download-URL: https://github.com/eamonn-zh/torchmetrics_ext
5
+ Author: Yiming Zhang
6
+ Requires-Dist: torchmetrics
7
+ Requires-Dist: torch
@@ -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,2 @@
1
+ torchmetrics
2
+ torch
@@ -0,0 +1 @@
1
+ torchmetrics_ext