Spaces:
Running on Zero
Running on Zero
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}")
|