-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathutility.py
More file actions
128 lines (106 loc) · 4.47 KB
/
Copy pathutility.py
File metadata and controls
128 lines (106 loc) · 4.47 KB
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
"""Paired optimizer outcomes used by GATE-A and later utility supervision."""
from __future__ import annotations
from dataclasses import asdict, dataclass
from enum import Enum
class UtilityLabel(str, Enum):
HELPFUL = "helpful"
NEUTRAL = "neutral"
HARMFUL = "harmful"
@dataclass(frozen=True)
class OptimizerOutcome:
"""Metrics from one frozen optimizer run, in the Gate-A evaluation units."""
gt_edge_distance_mm: float
contact_f1_30mm: float
mpjpe_mm: float
pa_mpjpe_mm: float
pve_translation_aligned_mm: float
relative_translation_error_mm: float
penetration_max_mm: float
penetration_mean_mm: float
reprojection_error_px: float
optimizer_success: bool = True
@dataclass(frozen=True)
class FrozenDecisionThresholds:
primary_contact_improvement_mm: float
contact_f1_change: float
mpjpe_noninferiority_mm: float
pa_mpjpe_noninferiority_mm: float
pve_translation_aligned_noninferiority_mm: float
relative_translation_error_noninferiority_mm: float
penetration_max_noninferiority_mm: float
penetration_mean_noninferiority_mm: float
reprojection_noninferiority_px: float
SAFETY_DELTA_TO_MARGIN = {
"mpjpe_mm": "mpjpe_noninferiority_mm",
"pa_mpjpe_mm": "pa_mpjpe_noninferiority_mm",
"pve_translation_aligned_mm": "pve_translation_aligned_noninferiority_mm",
"relative_translation_error_mm": "relative_translation_error_noninferiority_mm",
"penetration_max_mm": "penetration_max_noninferiority_mm",
"penetration_mean_mm": "penetration_mean_noninferiority_mm",
"reprojection_error_px": "reprojection_noninferiority_px",
}
def paired_deltas(base: OptimizerOutcome, intervention: OptimizerOutcome) -> dict[str, float]:
"""Return intervention-minus-base deltas; lower is better except contact F1."""
return {
"gt_edge_distance_mm": intervention.gt_edge_distance_mm - base.gt_edge_distance_mm,
"contact_f1_30mm": intervention.contact_f1_30mm - base.contact_f1_30mm,
"mpjpe_mm": intervention.mpjpe_mm - base.mpjpe_mm,
"pa_mpjpe_mm": intervention.pa_mpjpe_mm - base.pa_mpjpe_mm,
"pve_translation_aligned_mm": (
intervention.pve_translation_aligned_mm - base.pve_translation_aligned_mm
),
"relative_translation_error_mm": (
intervention.relative_translation_error_mm - base.relative_translation_error_mm
),
"penetration_max_mm": intervention.penetration_max_mm - base.penetration_max_mm,
"penetration_mean_mm": intervention.penetration_mean_mm - base.penetration_mean_mm,
"reprojection_error_px": intervention.reprojection_error_px - base.reprojection_error_px,
}
def safety_violations(
deltas: dict[str, float], thresholds: FrozenDecisionThresholds
) -> list[str]:
"""Return safety metrics whose intervention delta exceeds its frozen margin."""
return [
metric
for metric, margin_name in SAFETY_DELTA_TO_MARGIN.items()
if deltas[metric] > getattr(thresholds, margin_name)
]
def classify_utility(
base: OptimizerOutcome,
intervention: OptimizerOutcome,
thresholds: FrozenDecisionThresholds,
) -> tuple[UtilityLabel, dict[str, float]]:
"""Classify exactly as preregistered for the 50-event Gate-A experiment.
Contact F1 is retained as a diagnostic and paper-claim check. It does not
override the primary GT-edge-distance definition used for the utility label.
"""
delta = paired_deltas(base, intervention)
if not intervention.optimizer_success:
return UtilityLabel.HARMFUL, delta
primary_delta = delta["gt_edge_distance_mm"]
primary_help = primary_delta < -thresholds.primary_contact_improvement_mm
primary_harm = primary_delta > thresholds.primary_contact_improvement_mm
if safety_violations(delta, thresholds) or primary_harm:
return UtilityLabel.HARMFUL, delta
if primary_help:
return UtilityLabel.HELPFUL, delta
return UtilityLabel.NEUTRAL, delta
def serialize_pair(
event_id: str,
condition: str,
base: OptimizerOutcome,
intervention: OptimizerOutcome,
thresholds: FrozenDecisionThresholds,
*,
scientific_result: bool = False,
) -> dict:
label, deltas = classify_utility(base, intervention, thresholds)
return {
"event_id": event_id,
"condition": condition,
"scientific_result": scientific_result,
"label": label.value,
"base": asdict(base),
"intervention": asdict(intervention),
"deltas": deltas,
}