File size: 4,953 Bytes
bc4c433
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
"""Occupancy classification metrics from logits vs labels.

The model emits raw logits (no sigmoid in ``forward``). Metrics apply
sigmoid only at decision time so they stay consistent with
``BCEWithLogitsLoss`` and with inference (threshold ``0.5``).

Inside (label ``1``) is the product-relevant class: points that belong
in the interior. Precision / recall are therefore reported for that
class only. No extra packages (sklearn, etc.).
"""

from __future__ import annotations

from dataclasses import dataclass

import torch
from torch import Tensor


@dataclass(frozen=True)
class OccupancyMetrics:
    """Pointwise occupancy scores for one batch or a full set of queries."""

    accuracy: float
    inside_precision: float
    inside_recall: float
    inside_iou: float
    inside_f1: float


def _safe_div(numerator: float, denominator: float) -> float:
    """Return ``0.0`` when the count in the denominator is zero (no sklearn)."""
    if denominator <= 0.0:
        return 0.0
    return numerator / denominator


def _flatten_pair(logits: Tensor, labels: Tensor) -> tuple[Tensor, Tensor]:
    """
    Collapse ``(B, 1)`` or ``(B,)`` logits / labels to a shared 1-D view.

    Last dim of logits is 1 when it comes from ``OccupancyMLP``; labels from
    the Dataset match that. A 1-D label vector is accepted so callers do not
    have to unsqueeze.
    """
    if logits.numel() != labels.numel():
        raise ValueError(
            f"logits and labels must have the same number of elements, "
            f"got logits={tuple(logits.shape)} labels={tuple(labels.shape)}"
        )
    return logits.reshape(-1), labels.reshape(-1)


def occupancy_metrics(
    logits: Tensor,
    labels: Tensor,
    *,
    threshold: float = 0.5,
) -> OccupancyMetrics:
    """
    Accuracy plus inside precision / recall from occupancy logits.

    Parameters
    ----------
    logits:
        Unnormalized scores, shape ``(B, 1)`` or ``(B,)``. Positive → inside.
    labels:
        Float ``{0, 1}`` with the same number of elements as ``logits``.
    threshold:
        Decision cut on ``sigmoid(logit)``. Default ``0.5`` matches inference.

    Returns
    -------
    OccupancyMetrics
        Scalar floats on CPU (safe to print or average across batches).
    """
    # Sigmoid here only: training still uses BCE-with-logits on raw logits.
    tp, fp, fn, correct, n = occupancy_counts(logits, labels, threshold=threshold)
    return occupancy_metrics_from_counts(tp=tp, fp=fp, fn=fn, correct=correct, n=n)


def occupancy_counts(
    logits: Tensor,
    labels: Tensor,
    *,
    threshold: float = 0.5,
) -> tuple[float, float, float, float, float]:
    """Return ``(tp, fp, fn, correct, n)`` for a micro-average over points."""
    logits_flat, labels_flat = _flatten_pair(logits, labels)
    pred_inside = logits_flat.sigmoid() >= threshold
    true_inside = labels_flat > 0.5
    pred_f = pred_inside.to(dtype=torch.float32)
    true_f = true_inside.to(dtype=torch.float32)
    tp = float((pred_f * true_f).sum().item())
    fp = float((pred_f * (1.0 - true_f)).sum().item())
    fn = float(((1.0 - pred_f) * true_f).sum().item())
    correct = float((pred_inside == true_inside).to(dtype=torch.float32).sum().item())
    n = float(pred_inside.numel())
    return tp, fp, fn, correct, n


def occupancy_metrics_from_counts(
    *,
    tp: float,
    fp: float,
    fn: float,
    correct: float,
    n: float,
) -> OccupancyMetrics:
    """Build metrics from accumulated confusion counts (point micro-average)."""
    precision = _safe_div(tp, tp + fp)
    recall = _safe_div(tp, tp + fn)
    return OccupancyMetrics(
        accuracy=_safe_div(correct, n),
        inside_precision=precision,
        inside_recall=recall,
        inside_iou=_safe_div(tp, tp + fp + fn),
        inside_f1=_safe_div(2.0 * precision * recall, precision + recall),
    )


def accuracy_from_logits(
    logits: Tensor,
    labels: Tensor,
    *,
    threshold: float = 0.5,
) -> float:
    """Fraction of points whose thresholded sigmoid matches the 0/1 label.

    Parameters
    ----------
    logits, labels, threshold:
        Same meaning as :func:`occupancy_metrics`.

    Returns
    -------
    float
        Accuracy in ``[0, 1]``.
    """
    return occupancy_metrics(logits, labels, threshold=threshold).accuracy


if __name__ == "__main__":
    # Deterministic 8-point batch: 4 true insides, 4 true outsides, all correct.
    demo_logits = torch.tensor(
        [[4.0], [3.0], [2.0], [1.0], [-1.0], [-2.0], [-3.0], [-4.0]]
    )
    demo_labels = torch.tensor(
        [[1.0], [1.0], [1.0], [1.0], [0.0], [0.0], [0.0], [0.0]]
    )
    scores = occupancy_metrics(demo_logits, demo_labels)
    print(f"accuracy={scores.accuracy:.4f}")
    print(f"inside_precision={scores.inside_precision:.4f}")
    print(f"inside_recall={scores.inside_recall:.4f}")
    print(f"inside_iou={scores.inside_iou:.4f}")
    print(f"inside_f1={scores.inside_f1:.4f}")