MaziyarPanahi commited on
Commit
d5b398e
·
0 Parent(s):

ModernJEV-Decide-Preview: model, results, compute costs and ML Intern workflow

Browse files
.gitattributes ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz filter=lfs diff=lfs merge=lfs -text
33
+ *.zip filter=lfs diff=lfs merge=lfs -text
34
+ *.zst filter=lfs diff=lfs merge=lfs -text
35
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
COST-AUDIT.json ADDED
@@ -0,0 +1,207 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "source": "Live job total durations (including scheduling conservatively) and current published per-minute hardware rates; not an invoice",
3
+ "invoice_verified": false,
4
+ "previous_jobs_compute_upper_bound_usd": 2.208502,
5
+ "jobs": [
6
+ {
7
+ "id": "6abce8d0031314b6963447a4",
8
+ "flavor": "cpu-basic",
9
+ "startedAt": "2026-09-30T10:47:47.400Z",
10
+ "finishedAt": "2026-09-30T10:47:50.469Z",
11
+ "durations": {
12
+ "schedulingSecs": 3,
13
+ "runningSecs": 3,
14
+ "totalSecs": 6
15
+ },
16
+ "status": {
17
+ "stage": "ERROR",
18
+ "message": "Job failed with exit code: 3. Reason: Error",
19
+ "failureCount": 0
20
+ },
21
+ "conservative_billable_minutes": 1,
22
+ "pricing_unit_usd": 0.000167,
23
+ "estimated_upper_bound_usd": 0.000167
24
+ },
25
+ {
26
+ "id": "6abcf03c031314b69634487d",
27
+ "flavor": "a100-large",
28
+ "startedAt": "2026-09-30T11:19:38.618Z",
29
+ "finishedAt": "2026-09-30T11:26:45.280Z",
30
+ "durations": {
31
+ "schedulingSecs": 13,
32
+ "runningSecs": 426,
33
+ "totalSecs": 440
34
+ },
35
+ "status": {
36
+ "stage": "ERROR",
37
+ "message": "Job failed with exit code: 1. Reason: Error",
38
+ "failureCount": 0
39
+ },
40
+ "conservative_billable_minutes": 8,
41
+ "pricing_unit_usd": 0.041667,
42
+ "estimated_upper_bound_usd": 0.333336
43
+ },
44
+ {
45
+ "id": "6abcf3be4c46ef198703619a",
46
+ "flavor": "a100-large",
47
+ "startedAt": "2026-09-30T11:34:28.605Z",
48
+ "finishedAt": "2026-09-30T11:37:23.342Z",
49
+ "durations": {
50
+ "schedulingSecs": 5,
51
+ "runningSecs": 174,
52
+ "totalSecs": 180
53
+ },
54
+ "status": {
55
+ "stage": "ERROR",
56
+ "message": "Job failed with exit code: 1. Reason: Error",
57
+ "failureCount": 0
58
+ },
59
+ "conservative_billable_minutes": 3,
60
+ "pricing_unit_usd": 0.041667,
61
+ "estimated_upper_bound_usd": 0.125001
62
+ },
63
+ {
64
+ "id": "6abcfa794c46ef1987036338",
65
+ "flavor": "a100-large",
66
+ "startedAt": "2026-09-30T12:03:11.619Z",
67
+ "finishedAt": "2026-09-30T12:05:30.241Z",
68
+ "durations": {
69
+ "schedulingSecs": 6,
70
+ "runningSecs": 138,
71
+ "totalSecs": 144
72
+ },
73
+ "status": {
74
+ "stage": "ERROR",
75
+ "message": "Job failed with exit code: 1. Reason: Error",
76
+ "failureCount": 0
77
+ },
78
+ "conservative_billable_minutes": 3,
79
+ "pricing_unit_usd": 0.041667,
80
+ "estimated_upper_bound_usd": 0.125001
81
+ },
82
+ {
83
+ "id": "6abcff404c46ef1987036446",
84
+ "flavor": "a100-large",
85
+ "startedAt": "2026-09-30T12:23:44.354Z",
86
+ "finishedAt": "2026-09-30T12:25:32.810Z",
87
+ "durations": {
88
+ "schedulingSecs": 15,
89
+ "runningSecs": 108,
90
+ "totalSecs": 124
91
+ },
92
+ "status": {
93
+ "stage": "ERROR",
94
+ "message": "Job failed with exit code: 1. Reason: Error",
95
+ "failureCount": 0
96
+ },
97
+ "conservative_billable_minutes": 3,
98
+ "pricing_unit_usd": 0.041667,
99
+ "estimated_upper_bound_usd": 0.125001
100
+ },
101
+ {
102
+ "id": "6abd00604c46ef1987036497",
103
+ "flavor": "a100-large",
104
+ "startedAt": "2026-09-30T12:28:24.196Z",
105
+ "finishedAt": "2026-09-30T12:31:55.114Z",
106
+ "durations": {
107
+ "schedulingSecs": 7,
108
+ "runningSecs": 210,
109
+ "totalSecs": 218
110
+ },
111
+ "status": {
112
+ "stage": "COMPLETED",
113
+ "message": null,
114
+ "failureCount": 0
115
+ },
116
+ "conservative_billable_minutes": 4,
117
+ "pricing_unit_usd": 0.041667,
118
+ "estimated_upper_bound_usd": 0.166668
119
+ },
120
+ {
121
+ "id": "6abd028f031314b696344aed",
122
+ "flavor": "h200",
123
+ "startedAt": "2026-09-30T12:37:43.301Z",
124
+ "finishedAt": "2026-09-30T12:41:20.011Z",
125
+ "durations": {
126
+ "schedulingSecs": 7,
127
+ "runningSecs": 216,
128
+ "totalSecs": 224
129
+ },
130
+ "status": {
131
+ "stage": "COMPLETED",
132
+ "message": null,
133
+ "failureCount": 0
134
+ },
135
+ "conservative_billable_minutes": 4,
136
+ "pricing_unit_usd": 0.083333,
137
+ "estimated_upper_bound_usd": 0.333332
138
+ },
139
+ {
140
+ "id": "6abd053cfbc85ba68235adf3",
141
+ "flavor": "h200",
142
+ "startedAt": "2026-09-30T12:49:08.277Z",
143
+ "finishedAt": "2026-09-30T13:00:31.350Z",
144
+ "durations": {
145
+ "schedulingSecs": 7,
146
+ "runningSecs": 683,
147
+ "totalSecs": 690
148
+ },
149
+ "status": {
150
+ "stage": "ERROR",
151
+ "message": "Job failed with exit code: 1. Reason: Error",
152
+ "failureCount": 0
153
+ },
154
+ "conservative_billable_minutes": 12,
155
+ "pricing_unit_usd": 0.083333,
156
+ "estimated_upper_bound_usd": 0.9999960000000001
157
+ }
158
+ ],
159
+ "next_job_hardware": "a100-large",
160
+ "next_job_rate_per_hour_usd": 2.50002,
161
+ "next_job_timeout_minutes": 165,
162
+ "conservative_combined_upper_bound_usd": 9.375055,
163
+ "final_job": {
164
+ "id": "6abd0d31404719ba37613271",
165
+ "flavor": "a100-large",
166
+ "startedAt": "2026-09-30T13:23:03.853Z",
167
+ "finishedAt": "2026-09-30T16:00:04.270Z",
168
+ "durations": {
169
+ "schedulingSecs": 6,
170
+ "runningSecs": 9420,
171
+ "totalSecs": 9426
172
+ },
173
+ "status": {
174
+ "stage": "COMPLETED",
175
+ "message": null,
176
+ "failureCount": 0
177
+ }
178
+ },
179
+ "final_job_estimated_upper_bound_usd": 6.5834,
180
+ "all_jobs_estimated_upper_bound_usd": 8.7919,
181
+ "completed_under_budget": true,
182
+ "successful_run_breakdown": {
183
+ "applies_to": "Original 60000-decision checkpoint, not separate 6000-decision ML Intern replay",
184
+ "job_id": "6abd0d31404719ba37613271",
185
+ "training_rows": 60000,
186
+ "optimizer_steps": 3720,
187
+ "training_phase_seconds": 7786.5,
188
+ "training_phase_minutes": 129.775,
189
+ "published_a100_hourly_rate_usd": 2.5,
190
+ "training_phase_estimate_usd": 5.4073,
191
+ "full_successful_job_conservative_estimate_usd": 6.5834,
192
+ "guide_successful_run_approximate_compute_usd": 6.6,
193
+ "earlier_failed_attempts_and_pilots_estimate_usd": 2.208502,
194
+ "historical_development_total_estimate_usd": 8.7919,
195
+ "pricing_source": "https://huggingface.co/docs/hub/jobs-pricing",
196
+ "pricing_verified_date": "2026-10-01",
197
+ "training_duration_source": "train_done event in the successful job logs; 7786.5 seconds, 3720 steps, 60000 rows",
198
+ "invoice_verified": false,
199
+ "excluded": [
200
+ "HuggingChat inference",
201
+ "local CPU work",
202
+ "operator time",
203
+ "separate later 6k replay"
204
+ ],
205
+ "guide_note": "Runtime-based estimate for one successful run; not a fixed price guarantee."
206
+ }
207
+ }
JOB-LEDGER.json ADDED
@@ -0,0 +1,106 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "checked_at": "2026-09-30T16:07:03.089274+00:00",
3
+ "authorized_total_budget_usd": 10,
4
+ "prior_jobs_conservative_reserve_usd": 2.5,
5
+ "current_job_hardware_ceiling_usd": 6.875,
6
+ "combined_upper_bound_usd": 9.375,
7
+ "invoice_verified": false,
8
+ "billing_account": "OpenMed",
9
+ "chat_tokens_billing_account": "MaziyarPanahi; accepted separately by user",
10
+ "full_60000_training_started": true,
11
+ "jobs": [
12
+ {
13
+ "id": "6abce8d0031314b6963447a4",
14
+ "hardware": "cpu-basic",
15
+ "status": "ERROR",
16
+ "phase": "credential probe",
17
+ "optimizer_steps": 0,
18
+ "failure": "Connector lacked private repository write/bucket mount permission; user authorized existing local credential."
19
+ },
20
+ {
21
+ "id": "6abcf03c031314b69634487d",
22
+ "hardware": "a100-large",
23
+ "status": "ERROR",
24
+ "phase": "first pilot",
25
+ "optimizer_steps": 0,
26
+ "failure": "Datasets5 lazy Column used as NumPy array; bulk materialization fixed"
27
+ },
28
+ {
29
+ "id": "6abcf3be4c46ef198703619a",
30
+ "hardware": "a100-large",
31
+ "status": "ERROR",
32
+ "phase": "repaired pilot before prototype",
33
+ "optimizer_steps": 25,
34
+ "failure": "Reload assertion compared Trainer BF16-wrapped output against plain FP32 output; saved checkpoints retained",
35
+ "checkpoint_retained": true
36
+ },
37
+ {
38
+ "id": "6abcfa794c46ef1987036338",
39
+ "hardware": "a100-large",
40
+ "status": "ERROR",
41
+ "phase": "FlashAttention dependency check",
42
+ "optimizer_steps": 0,
43
+ "failure": "kernels0.17.1 outside Transformers5.17 allowed range [0.16.0,0.17.0)"
44
+ },
45
+ {
46
+ "id": "6abcff404c46ef1987036446",
47
+ "hardware": "a100-large",
48
+ "status": "ERROR",
49
+ "phase": "CUDA kernel resolution check",
50
+ "optimizer_steps": 0,
51
+ "failure": "Pinned model-repository SHA was looked up in kernel-repository type by kernels0.16; repositories have distinct SHAs"
52
+ },
53
+ {
54
+ "id": "6abd00604c46ef1987036497",
55
+ "hardware": "a100-large",
56
+ "status": "COMPLETED",
57
+ "phase": "corrected kernel check then measured pilot and conditional prototype",
58
+ "timeout_seconds": 7200,
59
+ "hardware_ceiling_usd": 5.0,
60
+ "optimizer_steps": 25,
61
+ "kernel_forward_backward_verified": true,
62
+ "checkpoint_reload_verified": true,
63
+ "full_prototype_started": false,
64
+ "stop_reason": "Measured unbucketed A100 throughput did not fit the two-hour deadline."
65
+ },
66
+ {
67
+ "id": "6abd028f031314b696344aed",
68
+ "hardware": "h200",
69
+ "status": "COMPLETED",
70
+ "phase": "matching-batch pilot then conditional 60,000-decision prototype",
71
+ "timeout_seconds": 4800,
72
+ "hardware_ceiling_usd": 6.66664,
73
+ "optimizer_steps": 40,
74
+ "checkpoint_reload_verified": true,
75
+ "full_prototype_started": false,
76
+ "measured_rows_per_second": 18.72,
77
+ "stop_reason": "Pilot estimate plus conservative setup/evaluation reservations exceeded the 80-minute wrapper allowance."
78
+ },
79
+ {
80
+ "id": "6abd053cfbc85ba68235adf3",
81
+ "hardware": "h200",
82
+ "status": "ERROR",
83
+ "phase": "full 60,000-decision prototype using frozen successful H200 pilot",
84
+ "timeout_seconds": 5340,
85
+ "hardware_ceiling_usd": 7.41664,
86
+ "pilot_results_revision": "ef6c274e502adb27814f80e492ff5904c4116bfd",
87
+ "optimizer_steps": 0,
88
+ "failure": "TrainingArguments5.17 removed warmup_ratio; prototype-only branch was not included in pilot"
89
+ },
90
+ {
91
+ "id": "6abd0d31404719ba37613271",
92
+ "hardware": "a100-large",
93
+ "status": "COMPLETED",
94
+ "phase": "repaired actual prototype; lazy batch tokenization; warmup_steps0.03",
95
+ "timeout_seconds": 9900,
96
+ "hardware_ceiling_usd": 6.875,
97
+ "recipe_revision": "b402b938f096ec3bd1bcc25c5655b6c4b16e64e1",
98
+ "optimizer_steps_verified": 3720,
99
+ "rows_seen_verified": 60000,
100
+ "first_optimizer_seconds_after_train_process_start": 44.8,
101
+ "finished_at": "2026-09-30T16:00:04.270Z",
102
+ "estimated_compute_upper_bound_usd": 6.5834
103
+ }
104
+ ],
105
+ "all_jobs_estimated_cost_usd": 8.7919
106
+ }
PUBLIC-PROMPT-SHORT.md ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Short HuggingChat ML Intern prompt
2
+
3
+ Before sending: select the account/org that should pay for chat-model tokens. Set the same intended namespace for GPU Jobs.
4
+
5
+ ```text
6
+ Help me train my own small Jev-style decision model with HuggingChat ML Intern.
7
+
8
+ Use answerdotai/ModernBERT-base and MaziyarPanahi/AgentToolDecisions-180K. Train on 60,000 frozen training decisions from next-action and tool-selection, with 4,096-token input pairs. Rank each decision's declared choices using a shared encoder and scalar scoring head.
9
+
10
+ My Job billing namespace: <YOUR_ACCOUNT_OR_ORG>
11
+ New PRIVATE model repo: <YOUR_ACCOUNT_OR_ORG>/ModernJEV-Decide-Preview
12
+ Total compute budget, including failed jobs: $10.
13
+
14
+ First show me a read-only plan and budget; wait for my authorization before launching anything. After approval, validate dependencies and CUDA, measure a short A100 pilot, and proceed only if the full prototype fits the budget. Keep the official splits. Evaluate all 1,700 in-scope test decisions and all 3,652 unseen-task When2Call test rows separately, without training or tuning on When2Call. Report per-task scores against constant-majority and allowed-choice frequency baselines; never report an overall accuracy. Test arbitrary answer labels and lists, including candidate order and label renaming. Save weights before evaluation, plus a useful README, tested prediction helper, actual coverage, costs, and the reproducible recipe.
15
+
16
+ Use JSONL metrics. No Spaces, no visibility changes, no automatic retries or budget extensions. Keep the model private until I authorize release. Do not claim Jev parity.
17
+ ```
18
+
19
+ For implementation details and lessons from this run, use PUBLIC-PROMPT.md. This short prompt is a starting request; review the generated executable plan before granting compute.
PUBLIC-PROMPT.md ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Train your own small agent decision model with HuggingChat
2
+
3
+ Open https://huggingface.co/chat/, enable **ML Intern**, replace the three settings below, and paste the prompt. Start without granting compute. After reviewing the proposal, grant the stated budget to authorize the pilot and bounded prototype. Builders need their own Hugging Face compute credits and permission to write the destination model repository.
4
+
5
+ **Check both billing accounts:** select the intended user or organization inside HuggingChat before using the coding model, and pass that namespace on every Hugging Face Job. A Job's organization does not automatically change the account charged for chat-model tokens.
6
+
7
+ ```text
8
+ NAMESPACE = <your HF username or organization>
9
+ MODEL_REPO = <NAMESPACE>/ModernJEV-Decide-Preview
10
+ TRAINING_DECISIONS_TARGET = 60000
11
+ MAX_SEQUENCE_LENGTH = 4096
12
+ TOTAL_COMPUTE_BUDGET_USD = 10
13
+
14
+ Help me train a small open decision model from
15
+ MaziyarPanahi/AgentToolDecisions-180K using HuggingChat ML Intern.
16
+ Do a read-only preflight first, show your executable plan and budget,
17
+ then STOP for my explicit authorization before any compute job,
18
+ sandbox, repository creation, training or Hub write.
19
+ Use read-only Hub metadata and dataset previews for that preflight.
20
+ Any executable dataset audit, model-load or GPU check belongs in the
21
+ approved pilot; do not provision compute to satisfy a zero-spend step.
22
+
23
+ The model should read an agent's state, a question, and the allowed
24
+ choices, then return one of those choices. It is a Jev-style choice
25
+ scorer, not a reproduction of Jev and not a chatbot.
26
+
27
+ DATA
28
+ Pin dataset revision f2fb14e4ec977c420f376c08785664cd38763d7e.
29
+ Keep the published train/validation/test splits; never resplit.
30
+ Train ONLY agent_next_action_type and tool_selection.
31
+ Their combined train/validation/test counts are 112973/1798/1700.
32
+ Inspect actual rows, label balance and input lengths before training.
33
+ Exclude the single-label tool_or_text_action,
34
+ tool_argument_completeness and tool_response_preference families,
35
+ and exclude the test-only when_to_call_tool family from training.
36
+ After freezing the final checkpoint, evaluate all 3,652 When2Call
37
+ test rows separately as unseen-task transfer, without tuning on them.
38
+ Report only per-task scores, never overall or macro accuracy.
39
+ Reverify constant-majority references: 616/1158 next-action,
40
+ 79/542 tool-selection, and 1295/3652 When2Call. Distinguish these
41
+ from the stronger allowed-choice training-frequency baseline.
42
+ When2Call majority is a descriptive test-set reference, not fitted.
43
+ Next-action has three declared choices but only two gold labels;
44
+ tool_and_response is never gold. Do not imply three-class coverage.
45
+ Assert no group_id crosses splits. For source-row leakage, use the
46
+ full tuple (source_dataset, source_revision, source_config,
47
+ source_split, source_row_id), never bare source_row_id.
48
+
49
+ MODEL AND LOSS
50
+ Start with answerdotai/ModernBERT-base at revision
51
+ 8949b909ec900327062f0ebf497f51aef5e6f0c8.
52
+ Use AutoModelForSequenceClassification(num_labels=1,
53
+ attn_implementation='sdpa'). Score each (question + state,
54
+ candidate label + criterion) pair. Apply softmax cross-entropy
55
+ within each decision's row_id; group_id is an episode, NOT the
56
+ loss grouping key. Score only the row's declared choices.
57
+ Support arbitrary caller-supplied labels and unique answer-string lists;
58
+ test 2/3/7/20 choices, reordered candidates and renamed labels.
59
+ Report API acceptance separately from semantic correctness.
60
+ Do not feed gold labels, gold JSON, gold scores or provenance into
61
+ model inputs. Do not include candidate indices in candidate text.
62
+ Verify trainer and dependency APIs against installed versions.
63
+ In Transformers 5.17, use warmup_steps=0.03 for 3% warmup;
64
+ warmup_ratio was removed. Instantiate the complete production
65
+ TrainingArguments branch before any expensive data preparation.
66
+ Test the full trainer/validation/checkpoint path locally on a tiny
67
+ subset, not only the separate pilot branch.
68
+ Prefer lazy batch tokenization with prefetching so optimizer steps
69
+ begin without an upfront full-dataset mapping pass.
70
+ Do not substitute a different ranking objective without approval.
71
+ For long 4096-token inputs, benchmark ModernBERT's documented
72
+ FlashAttention backend as well as SDPA. Match the installed Torch/CUDA
73
+ ABI to an available prebuilt kernel; do not assume the newest Torch
74
+ release has a matching kernel. Pin and record the kernel revision. For the verified environment in this
75
+ run: torch 2.12.0 + cu126, transformers 5.17.0, kernels 0.16.0;
76
+ Transformers 5.17 rejects kernels 0.17.x. The pinned FlashAttention
77
+ revision is f50dc99ed079b35990bc895d43fd353ea0cb376d.
78
+ Inspect the kernel repository type used by the installed loader. The
79
+ model repository and kernel repository may have different revisions.
80
+ Before data loading, run a small CUDA forward/backward kernel smoke.
81
+ Datasets 5 columns are lazy: bulk-materialize metadata and NumPy arrays
82
+ before numerical comparisons or sorting. Tokenize only sampled training
83
+ candidates if using a sampled pool; evaluation must rank all candidates.
84
+ Verify checkpoint parameter equality and compare predictions at the
85
+ same precision. Do not compare a Trainer BF16-wrapped forward to FP32
86
+ inference with an unrealistically strict tolerance.
87
+
88
+ INPUTS
89
+ The audited 192-row sample had 2221 candidate pairs: full-input
90
+ p50 2460 tokens, p95 3935, max 5179. At 2048, 1643 pairs truncate.
91
+ Start the pilot at 4096 tokens. Preserve question and candidate;
92
+ document any state truncation and measure its frequency. Remove
93
+ duplicated policy only when exactly equal to the first system
94
+ message. Batch by candidate count/token cost and measure GPU
95
+ memory. Never claim all input was read when it was truncated.
96
+
97
+ RUN
98
+ Pass NAMESPACE on EVERY job so billing uses the intended account.
99
+ Check current hardware price and private Hub write permissions.
100
+ If write access is missing, ask me to reconnect with the appropriate
101
+ repository permissions; never request a token pasted into the chat.
102
+ Use one A100 80GB (a100-large) if available at $2.50/hour:
103
+ first a 25-step pilot with a 20-minute timeout, then at most one
104
+ prototype job with a 165-minute timeout. Set a total $10 compute budget.
105
+ Select exactly TRAINING_DECISIONS_TARGET rows from the training split,
106
+ stratified by task family with seed 42; freeze and hash their row IDs.
107
+ Use measured pilot throughput to check whether one complete epoch fits,
108
+ leaving time for validation, final evaluation and upload.
109
+ Record the actual rows, steps, elapsed time and cost. No guarantee
110
+ of a full epoch, no parallel jobs, no automatic retries or extensions.
111
+ If the pilot makes the deadline infeasible, stop and report it.
112
+
113
+ Create a NEW PRIVATE model repository only after authorization.
114
+ Save checkpoints, tokenizer, input template, prediction code,
115
+ dependency pins, metrics and a model card there. Do not overwrite
116
+ an existing repository. Persist the pilot before its job exits.
117
+ Never create a Hugging Face Space or change any Space visibility.
118
+ Use local JSONL metrics; optionally use Trackio only after verifying
119
+ space_id=None performs local logging without creating a Space.
120
+ Never call create_trackio or provision a dashboard Space.
121
+ Never paste tokens into chat, code, logs or artifacts; use secrets.
122
+
123
+ EVALUATION AND DELIVERY
124
+ Use a fixed one-epoch recipe and its final checkpoint for the first
125
+ prototype; monitor validation without selecting on the test set.
126
+ If the time limit stops training early, disclose the actual row coverage.
127
+ Evaluate the final checkpoint
128
+ once on the 1700 in-scope official test decisions, per family.
129
+ Compare with uniform-choice and training-derived majority/frequency
130
+ baselines restricted to the declared candidates. Check shuffled
131
+ candidate order and report label agreement, with a stated tie rule.
132
+ Return typed-choice prediction code and a reproducible recipe.
133
+ Preserve source attribution and inspect licenses before release.
134
+ Report failures and unmeasured results honestly. Do not claim Jev
135
+ parity, general agent competence or support for other primitives.
136
+ Keep weights private until I separately authorize publication.
137
+ ```
README.md ADDED
@@ -0,0 +1,230 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: transformers
4
+ pipeline_tag: text-classification
5
+ base_model: answerdotai/ModernBERT-base
6
+ base_model_relation: finetune
7
+ datasets:
8
+ - MaziyarPanahi/AgentToolDecisions-180K
9
+ tags:
10
+ - modernbert
11
+ - encoder
12
+ - decision-model
13
+ - tool-routing
14
+ - agentic
15
+ - preview
16
+ language:
17
+ - en
18
+ ---
19
+
20
+ <div align="center">
21
+
22
+ # ModernJEV-Decide-Preview
23
+
24
+ ### Give your agent a next move.
25
+
26
+ **A small encoder for choosing an action or a tool from the options you provide.** We trained this 60,000-decision prototype for **about $6.60 in A100 compute, including setup and evaluation**. The training phase itself accounted for **about $5.40**. Both figures are estimates from the measured runtime.
27
+
28
+ 149.6M parameters · 4,096-token inputs · Typed choices · Built on ModernBERT
29
+
30
+ </div>
31
+
32
+ > **Experimental preview.** The final checkpoint processed **60,000/60,000 selected training decisions**. Held-out results cover **1,158 next-action**, **542 tool-selection**, and **3,652 unseen When2Call** decisions, reported separately. This is a small choice-ranking prototype, not a Jev reproduction.
33
+
34
+ An agent doesn't need to write a paragraph every time it makes a decision. Sometimes the useful output is simply **answer**, **call a tool**, or **choose this tool**.
35
+
36
+ ModernJEV-Decide-Preview reads a conversation, its policy and available tools, then scores the choices supplied with your question. Your application receives one of those labels, along with the scores for the alternatives. The model ranks choices; your application executes the selected action.
37
+
38
+ The architecture is a **ModernBERT cross-encoder with one scalar scoring head**. Choice labels and descriptions are input text, so the output is not limited to a fixed vocabulary of tool names. This is a Jev-style choice model built from public agent decisions, with its own training and evaluation; it is not a reproduction of Jev.
39
+
40
+ ![How decisions are scored](architecture.svg)
41
+
42
+ ## Where it fits
43
+
44
+ | Use case | Supply | Receive |
45
+ |---|---|---|
46
+ | **Next-action routing** | Current conversation, policy, tool list | A declared action such as `text_response` or `tool_call` |
47
+ | **Tool selection** | The task and each available tool's purpose | The label of the selected tool |
48
+ | **Workflow branching** | A text state and clearly described alternatives | One branch label; evaluate on your own workflow before adoption |
49
+ | **Decision-model experiments** | Your choices and held-out tasks | Candidate rankings you can inspect and compare |
50
+
51
+ Useful places to start: support agents choosing between lookup and escalation, assistants selecting an API, and workflows choosing between a direct answer and a tool-backed step. The model does not generate tool arguments or execute tools.
52
+
53
+ ## Train a decision model with HuggingChat ML Intern
54
+
55
+ We tested this workflow with **one execution message, an attached tested recipe, and a budget setup click**. In a fresh HuggingChat conversation, ML Intern prepared the code, ran its own checks, launched an OpenMed A100 job, trained **6,000 decisions in 1,500 optimizer steps**, and uploaded a new private checkpoint. No follow-up implementation prompt, Codex code repair or Codex job launch was needed in that conversation. Training took **996.9 seconds (16 minutes 37 seconds)**; that excludes chat preparation, setup and subsequent evaluation.
56
+
57
+ **These are two separate runs:**
58
+
59
+ | | Weights in this repository | ML Intern recipe replay |
60
+ |---|---|---|
61
+ | Selected training decisions | 60,000 | 6,000 |
62
+ | Optimizer steps | 3,720 | 1,500 |
63
+ | Execution | Started with ML Intern; Codex debugged, launched and completed the run | ML Intern executed from one message with supplied source files |
64
+ | Evidence | Per-task model results below | Exact row coverage and checkpoint upload independently verified |
65
+ | Status at this documentation revision | Training and evaluation complete | Training and fresh checkpoint verified; operator stopped after workflow proof |
66
+
67
+ The first replay attempt stopped **before GPU spend** because HuggingChat's tools could not read the existing private source repository, although the local HF login could. We preserved that failed attempt and opened a fresh conversation with the source files attached. The successful training therefore demonstrates **one-prompt execution of a supplied recipe after setup**, not first-attempt success, code invention from scratch, or a guarantee for arbitrary training tasks. ML Intern made its own pre-launch code corrections. After the first GPU run had trained and verified checkpoint reload, its finalizer failed because stdin execution did not define `__file__`. It detected that failure and a reversed When2Call baseline unpack itself, then used the authorized corrective retry. The retry also trained all 6,000 rows and saved a checkpoint. The operator then stopped the remaining evaluations and CPU sandbox because the requested workflow proof was already established. Full replay packaging was not completed. These autonomous repairs, failed compute and the operator termination are part of the recorded run.
68
+
69
+ ### Try the same workflow
70
+
71
+ 1. Open [HuggingChat](https://huggingface.co/chat/) and enable **ML Intern**. Our replay used **GLM-5.3-Flash**.
72
+ 2. Select your intended billing account or organization in HuggingChat settings. Ensure you can create/write a private model repository and submit Jobs in that namespace.
73
+ 3. Download [the execution prompt](workflow/PROMPT-R2.txt) and [the tested source bundle](workflow/recipe-bundle.txt). Change the billing namespace and destination to your own **new private model repository** throughout the prompt. Keep the pinned source revisions. Do not reuse our destination name.
74
+ 4. Attach both files to one message and send the instruction below. Set the conversation's compute-budget control to **$10**; the written budget alone does not replace that UI setting.
75
+ 5. Let ML Intern do its checks, execution and monitoring. A submitted job is not proof of a trained model: check row coverage, saved weights, reload, costs and per-task evaluations.
76
+
77
+ This was the actual execution message accompanying our attachments:
78
+
79
+ ```text
80
+ Execute PROMPT-R2.txt now in ML Intern using the supplied recipe-bundle.txt.
81
+ OpenMed is the billing namespace. This single message authorizes the bounded
82
+ experiment described in the prompt. The earlier attempt stopped before compute
83
+ because the private source lookup failed; use these supplied source files,
84
+ and preserve that fact in the final report.
85
+ ```
86
+
87
+ For your run, replace `OpenMed` in that message with your billing namespace as well. The full attached prompt is part of the instruction, not optional background. It specifies the 6,000-row selection, 4,096-token inputs, rehearsal, privacy, one A100, timeout, total compute cap, at most one bounded retry, and per-task evaluation requirements. Do not claim our measured runtime or results for a new run before measuring them.
88
+
89
+ [Exact prompt](workflow/PROMPT-R2.txt) · [Source bundle and hashes](workflow/recipe-manifest.json) · [Run provenance](workflow/ONE-PROMPT-REPLAY.md) · [Replay job](https://huggingface.co/jobs/OpenMed/6abe2b6ffbc85ba6823612cc)
90
+
91
+ ## Quick start
92
+
93
+ Install PyTorch and Transformers, then use the included inference helper so the input formatting matches training:
94
+
95
+ ```bash
96
+ pip install "torch==2.12.0" "transformers==5.17.0" "huggingface-hub==1.33.0"
97
+ ```
98
+
99
+ Download the supplied helper, then run it alongside your application. During the private preview, your Hugging Face account must have access to the repository.
100
+
101
+ ```bash
102
+ hf download MaziyarPanahi/ModernJEV-Decide-Preview predict.py --local-dir modernjev
103
+ cd modernjev
104
+ ```
105
+
106
+ The included helper was executed against this checkpoint. See [the observed example output](example-output.json); the snippet below prints the actual prediction rather than promising a fixed answer.
107
+
108
+ ```python
109
+ from predict import DecisionModel
110
+
111
+ model = DecisionModel("MaziyarPanahi/ModernJEV-Decide-Preview")
112
+ decision = model.decide(
113
+ state={
114
+ "policy": "Use the order lookup tool when a customer asks about an order.",
115
+ "conversation": [
116
+ {"role": "user", "content": "Where is order A123?"}
117
+ ],
118
+ "available_tools": [
119
+ {"name": "lookup_order", "description": "Retrieve an order's delivery status."},
120
+ {"name": "search_catalog", "description": "Find products in the catalog."},
121
+ ],
122
+ },
123
+ question="Which tool should the assistant call next?",
124
+ criteria={
125
+ "lookup_order": "Retrieve the delivery status of the customer's order.",
126
+ "search_catalog": "Search for products that match a shopping request.",
127
+ },
128
+ )
129
+ print(decision["predicted_label"])
130
+ print(decision["candidates"])
131
+ ```
132
+
133
+ You can also supply `criteria=["Escalate to a human", "Search the product catalog"]` when each answer is its own description. The result contains the selected label, all candidate scores, and whether input truncation occurred. Scores are normalized **within the supplied choice set**; they are not calibrated confidence estimates and cannot be compared directly across unrelated requests.
134
+
135
+ ## Training recipe
136
+
137
+ | Setting | Recorded configuration |
138
+ |---|---|
139
+ | Base | `answerdotai/ModernBERT-base` |
140
+ | Base revision | `8949b909ec900327062f0ebf497f51aef5e6f0c8` |
141
+ | Parameters with scalar head | 149,605,633 |
142
+ | Input limit | 4,096 tokens including both sequences and special tokens |
143
+ | Dataset | `MaziyarPanahi/AgentToolDecisions-180K` |
144
+ | Dataset revision | `f2fb14e4ec977c420f376c08785664cd38763d7e` |
145
+ | Selected training target | Exactly 60,000 train-split decisions, stratified by family, seed 42 |
146
+ | Actual training coverage | 60,000 decisions; 3,720 optimizer steps |
147
+ | Attention backend | kernels-community/flash-attn2@f50dc99ed079b35990bc895d43fd353ea0cb376d |
148
+ | Training elapsed | 129.8 minutes |
149
+ | GPU | NVIDIA A100-SXM4-80GB |
150
+ | Task families | `agent_next_action_type`, `tool_selection` |
151
+ | Objective | Per-decision softmax cross-entropy over candidate scores |
152
+ | Candidate grouping | `row_id`, preserving `group_id` only for episode/split boundaries |
153
+ | Candidate pool | Gold plus up to three declared negatives per training decision; evaluation ranks every declared choice |
154
+ | Evaluation | Monitor official validation; evaluate the final fixed-epoch checkpoint on 1,700 in-scope official test decisions |
155
+
156
+ The dataset has 180,000 total rows. This prototype targets 60,000 of its 112,973 in-scope training decisions; **180,000 is not the number used for training**. Exact selected row IDs, coverage, steps and dependency versions accompany this checkpoint. This repository contains model artifacts; it does not expose an inference endpoint.
157
+
158
+ ## Evaluation
159
+
160
+ Every task is reported separately. The checkpoint was fixed after one pass over the selected training rows; no When2Call labels were used for training, checkpoint selection or hyperparameter tuning. The three single-label task families are excluded from quality claims.
161
+
162
+ | Task | Correct / cases | Model accuracy | Constant-majority reference | Lift |
163
+ |---|---:|---:|---:|---:|
164
+ | Next action | 850 / 1,158 | 73.40% | 53.20% | +20.21 pp |
165
+ | Tool selection | 337 / 542 | 62.18% | 14.58% | +47.60 pp |
166
+ | When2Call — unseen task family | 1,256 / 3,652 | 34.39% | 35.46% | -1.07 pp |
167
+
168
+ **Transfer limitation:** When2Call is below its majority reference. Accepting a new question and new answer labels does not establish that the model understands that task reliably. This preview demonstrates gains on its trained task families; general decision-making and clinical use remain unvalidated.
169
+
170
+ The next-action majority label is `text_response`; the tool-selection majority label is `transfer_to_human_agent`. Both are also the most frequent labels in the selected training subset. The fixed tool label is absent from 235 of the 542 choice lists, so it is a descriptive constant-label reference, not a valid per-row router. When2Call has no training rows: its reference is the test-set majority (a tie between `cannot_answer` and `tool_call`), not a fitted training baseline.
171
+
172
+ The stronger **allowed-choice training-frequency baseline** always chooses among the supplied answers. It is distinct from the constant-majority reference above.
173
+
174
+ | Comparison | Next action (1,158 cases) | Tool selection (542 cases) |
175
+ |---|---:|---:|
176
+ | ModernJEV-Decide-Preview | 73.40% | 62.18% |
177
+ | ModernBERT + untrained scalar head, seed 42 | 42.83% | 0.74% |
178
+ | Allowed-choice training-frequency baseline | 53.20% | 22.88% |
179
+ | Uniform over supplied choices, expected accuracy | 33.33% | 5.19% |
180
+
181
+ Next-action rows have three choices. Tool-selection rows have a median of 20 choices (range 11–32). When2Call rows have four choices; uniform expected accuracy is 25%. Evaluation ranks all declared candidates, including choices not sampled during training.
182
+
183
+ Untrained ModernBERT is a pretrained encoder with a randomly initialized scalar head, not a trained decision classifier. No frozen-backbone trained-head comparison or Jev API comparison was completed in this bounded run.
184
+
185
+ ### Open answer labels
186
+
187
+ The helper accepts arbitrary string labels with descriptions, or a list of unique answer strings. It has no fixed output-class vocabulary. Final-checkpoint interface checks passed with 2, 3, 7 and 20 supplied answers, including reversed order and fresh opaque labels. These are interface checks, not a broad capability benchmark.
188
+
189
+ For one illustrative routing scenario at four answer-list sizes, the intended description was selected at 2, 3 and 7 choices, but not at 20. Renaming labels while retaining descriptions changed the selected description in the 3- and 7-choice variants. **Label wording matters:** use descriptive labels and validate your own question/answer sets. The small illustrative checks are published in full, including the misses.
190
+
191
+ Candidate-order label agreement on the official test subset was **100% across five permutations of 299 decisions**. Reordering a fixed set of label strings is a different test from renaming those strings. Ties use lexical label order.
192
+
193
+ GPU latency for **one candidate forward pass**, excluding tokenization: median **23.6 ms**, p95 **24.2 ms**. A full decision scores multiple candidates; these are not end-to-end decision timings. The two trained tasks were evaluated with GPU BF16 inference. Supplemental When2Call and open-label checks used local CPU FP32 inference on the same saved weights.
194
+
195
+ See [per-task results](per-task-results.json), [all When2Call predictions](evaluation/when2call-predictions.jsonl), [open-label checks](evaluation/open-label-checks.json), [baseline definitions](evaluation/per-task-baselines.json), [training coverage](training_coverage.json), [runtime versions](runtime-versions.json) and [the pinned training recipe](recipe/train.py). The split audit found no cross-split episode or full upstream source conflicts. The 4,096-token limit truncated 414 candidate pairs in the two trained task tests; no When2Call rows were truncated.
196
+
197
+ ### Compute accounting
198
+
199
+ **For a successful run of the 60k recipe, allow about $6.60 in GPU compute at our measured runtime.** We used one A100 80GB at the published [Hugging Face Jobs rate of $2.50/hour](https://huggingface.co/docs/hub/jobs-pricing).
200
+
201
+ | Scope | Measured time | Estimated GPU compute |
202
+ |---|---:|---:|
203
+ | Training phase: 60,000 decisions, 3,720 optimizer steps | 129.8 minutes | **$5.41** |
204
+ | Full successful job: setup, training, evaluation and packaging | 157 minutes | **$6.58**, conservatively rounded; about **$6.60** |
205
+ | Earlier development attempts and pilots | Across multiple jobs | **$2.21** |
206
+ | Original development total, including the successful job | Across all original jobs | **$8.79** |
207
+
208
+ The **$8.79** figure includes our earlier failed attempts and pilots; it is not the cost of the final successful run. The $10 authorization was a spending cap, not a quoted training price. These figures refer to the original **60k checkpoint in this repository**, separately from the later 6k ML Intern workflow test.
209
+
210
+ Costs are runtime/rate estimates, not reconciled invoice amounts. They exclude HuggingChat inference, local CPU work and operator time; another run's runtime may differ. See [cost accounting](COST-AUDIT.json) for the job-level receipts and training-phase calculation.
211
+
212
+ ## Input design matters
213
+
214
+ - Give every choice a concrete, distinct description. Include the same available tool information the application actually has.
215
+ - Use the supplied serialization helper. Long state may need truncation at the 4,096-token limit; the helper reports it.
216
+ - The training data is English agent conversations and tool decisions. New domains and arbitrary workflow labels need their own evaluation.
217
+ - `tool_and_response` is a declared next-action choice but never a gold label in this dataset. That action is outside demonstrated positive training coverage.
218
+ - This preview implements **choice ranking**. It does not implement Jev's other primitives, vision inputs, generated explanations, or autonomous tool execution.
219
+
220
+ ## Reproduce the 60k model or the 6k workflow
221
+
222
+ Use [the original 60k training recipe](recipe/train.py) and [its detailed training prompt](PUBLIC-PROMPT.md) to inspect how these weights were built. Use the **6k ML Intern replay instructions above** to test the one-message execution workflow. Their scope, checkpoints, runtimes and metrics are different. All model-quality results on this page belong to the original 60k checkpoint.
223
+
224
+ ## Attribution
225
+
226
+ Base model: [Answer.AI ModernBERT](https://huggingface.co/answerdotai/ModernBERT-base), Apache 2.0.
227
+
228
+ Decision dataset: [MaziyarPanahi/AgentToolDecisions-180K](https://huggingface.co/datasets/MaziyarPanahi/AgentToolDecisions-180K), transformed from the upstream agentic sources documented in its card. Keep the pinned revisions and source attribution with derivative work, and follow the upstream dataset licenses when redistributing data.
229
+
230
+ **Original 60k workflow:** HuggingChat ML Intern prepared the proposal, training code, helper and model-card draft. Codex audited the data and corrected candidate targets, padding budgets, coverage accounting and checkpoint verification. Hugging Face Jobs was launched with the operator's approved local HF credential after the connector's write permissions returned 403. This is an assisted training workflow; HuggingChat did not autonomously execute the training job.
architecture.svg ADDED
config.json ADDED
@@ -0,0 +1,84 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "ModernBertForSequenceClassification"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "bos_token_id": 50281,
8
+ "classifier_activation": "gelu",
9
+ "classifier_bias": false,
10
+ "classifier_dropout": 0.0,
11
+ "classifier_pooling": "mean",
12
+ "cls_token_id": 50281,
13
+ "decoder_bias": true,
14
+ "deterministic_flash_attn": false,
15
+ "dtype": "float32",
16
+ "embedding_dropout": 0.0,
17
+ "eos_token_id": 50282,
18
+ "global_attn_every_n_layers": 3,
19
+ "gradient_checkpointing": false,
20
+ "hidden_activation": "gelu",
21
+ "hidden_size": 768,
22
+ "id2label": {
23
+ "0": "LABEL_0"
24
+ },
25
+ "initializer_cutoff_factor": 2.0,
26
+ "initializer_range": 0.02,
27
+ "intermediate_size": 1152,
28
+ "label2id": {
29
+ "LABEL_0": 0
30
+ },
31
+ "layer_norm_eps": 1e-05,
32
+ "layer_types": [
33
+ "full_attention",
34
+ "sliding_attention",
35
+ "sliding_attention",
36
+ "full_attention",
37
+ "sliding_attention",
38
+ "sliding_attention",
39
+ "full_attention",
40
+ "sliding_attention",
41
+ "sliding_attention",
42
+ "full_attention",
43
+ "sliding_attention",
44
+ "sliding_attention",
45
+ "full_attention",
46
+ "sliding_attention",
47
+ "sliding_attention",
48
+ "full_attention",
49
+ "sliding_attention",
50
+ "sliding_attention",
51
+ "full_attention",
52
+ "sliding_attention",
53
+ "sliding_attention",
54
+ "full_attention"
55
+ ],
56
+ "local_attention": 128,
57
+ "max_position_embeddings": 8192,
58
+ "mlp_bias": false,
59
+ "mlp_dropout": 0.0,
60
+ "model_type": "modernbert",
61
+ "norm_bias": false,
62
+ "norm_eps": 1e-05,
63
+ "num_attention_heads": 12,
64
+ "num_hidden_layers": 22,
65
+ "pad_token_id": 50283,
66
+ "position_embedding_type": "absolute",
67
+ "rope_parameters": {
68
+ "full_attention": {
69
+ "rope_theta": 160000.0,
70
+ "rope_type": "default"
71
+ },
72
+ "sliding_attention": {
73
+ "rope_theta": 10000.0,
74
+ "rope_type": "default"
75
+ }
76
+ },
77
+ "sep_token_id": 50282,
78
+ "sparse_pred_ignore_index": -100,
79
+ "sparse_prediction": false,
80
+ "tie_word_embeddings": true,
81
+ "transformers_version": "5.17.0",
82
+ "use_cache": false,
83
+ "vocab_size": 50368
84
+ }
evaluation/evaluate_open_labels.py ADDED
@@ -0,0 +1,169 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Supplemental frozen-checkpoint evaluation; no training or Hub writes.
2
+
3
+ Run --audit-only first. After final checkpoint is saved, supply --model and,
4
+ for a Hub model, --revision. CPU is default; JSONL predictions resume safely.
5
+ """
6
+ import argparse
7
+ import hashlib
8
+ import importlib.util
9
+ import json
10
+ import statistics
11
+ from collections import Counter
12
+ from pathlib import Path
13
+
14
+ ROOT = Path(__file__).resolve().parent
15
+ DATA_REV = 'f2fb14e4ec977c420f376c08785664cd38763d7e'
16
+ CACHE = Path('/private/tmp/atd-ml-intern-hub-cache/MaziyarPanahi___agent_tool_decisions-180_k/default/0.0.0') / DATA_REV
17
+ FAMILIES = ('agent_next_action_type', 'tool_selection', 'when_to_call_tool')
18
+ EXPECTED = {FAMILIES[0]: (1158, 616), FAMILIES[1]: (542, 79), FAMILIES[2]: (3652, 1295)}
19
+
20
+
21
+ def write_json(path, value):
22
+ temp = path.with_suffix('.tmp')
23
+ temp.write_text(json.dumps(value, indent=2) + '\n')
24
+ temp.replace(path)
25
+
26
+
27
+ def load_test():
28
+ from datasets import Dataset
29
+ return Dataset.from_file(str(CACHE / 'agent_tool_decisions-180_k-test.arrow'))
30
+
31
+
32
+ def audit(test):
33
+ report = {'dataset_revision': DATA_REV, 'families': {}}
34
+ for family in FAMILIES:
35
+ rows = [r for r in test if r['task_family'] == family]
36
+ labels = Counter(r['gold_label'] for r in rows)
37
+ size, correct = EXPECTED[family]
38
+ assert len(rows) == size and max(labels.values()) == correct
39
+ candidates = [len(json.loads(r['criteria_json'])) for r in rows]
40
+ assert all(r['gold_label'] in json.loads(r['criteria_json']) for r in rows)
41
+ report['families'][family] = {
42
+ 'n': size, 'majority_correct': correct, 'majority_accuracy': correct / size,
43
+ 'majority_labels': sorted(k for k,v in labels.items() if v == correct),
44
+ 'candidate_count': {'median': statistics.median(candidates),
45
+ 'min': min(candidates), 'max': max(candidates)},
46
+ 'uniform_expected_accuracy': statistics.mean(1/n for n in candidates),
47
+ 'reference_definition': 'Descriptive test-set constant-label majority; no model fitting',
48
+ 'unseen_task_family': family == 'when_to_call_tool',
49
+ }
50
+ # Verify When2Call is absent from every training shard, not just the subset.
51
+ from datasets import Dataset
52
+ train_families = Counter()
53
+ selected_ids = set((ROOT / 'selected-row-ids.txt').read_text().splitlines())
54
+ train_gold = {f: Counter() for f in FAMILIES[:2]}
55
+ for path in sorted(CACHE.glob('*-train-*.arrow')):
56
+ columns = Dataset.from_file(str(path)).select_columns(['task_family','row_id','gold_label'])[:]
57
+ for family, row_id, gold in zip(columns['task_family'], columns['row_id'], columns['gold_label']):
58
+ train_families[family] += 1
59
+ if row_id in selected_ids and family in train_gold:
60
+ train_gold[family][gold] += 1
61
+ assert sum(train_families.values()) == 171056
62
+ assert train_families['when_to_call_tool'] == 0
63
+ assert sum(sum(c.values()) for c in train_gold.values()) == 60000
64
+ report['when2call_training_rows'] = 0
65
+ for family, counts in train_gold.items():
66
+ label = min(counts, key=lambda k: (-counts[k], k))
67
+ rows = [r for r in test if r['task_family'] == family]
68
+ report['families'][family]['training_fixed_majority'] = {
69
+ 'label': label, 'correct': sum(r['gold_label'] == label for r in rows),
70
+ 'n': len(rows),
71
+ 'majority_label_absent_from_choices': sum(label not in json.loads(r['criteria_json']) for r in rows),
72
+ }
73
+ return report
74
+
75
+
76
+ def open_labels(client):
77
+ checks = []
78
+ for n in (2, 3, 7, 20):
79
+ descriptions = ['Escalate this damaged parcel to a human support agent.',
80
+ 'Search the product catalog for a new item.']
81
+ descriptions += [f'Route to unrelated department number {i} for a different request.' for i in range(n-2)]
82
+ criteria = {f'fresh_choice_{n}_{i}': d for i,d in enumerate(descriptions)}
83
+ state = {'policy': 'Escalate damaged parcels to a human support agent.',
84
+ 'conversation': [{'role':'user','content':'My parcel arrived broken. I need help.'}]}
85
+ question = 'Which of these custom workflow branches should handle this request?'
86
+ outputs = []
87
+ variants = (criteria, dict(reversed(list(criteria.items()))),
88
+ {f'opaque_{n}_{i}': d for i,d in enumerate(descriptions)}, descriptions)
89
+ for variant in variants:
90
+ answer = client.decide(state=state, question=question, criteria=variant, candidate_batch_size=4)
91
+ allowed = list(variant)
92
+ assert answer['predicted_label'] in allowed
93
+ assert sorted(c['label'] for c in answer['candidates']) == sorted(allowed)
94
+ assert len(answer['candidates']) == n
95
+ outputs.append(answer)
96
+ original_desc = criteria[outputs[0]['predicted_label']]
97
+ renamed_desc = variants[2][outputs[2]['predicted_label']]
98
+ checks.append({'answer_count': n, 'interface_pass': True,
99
+ 'reorder_same_label': outputs[0]['predicted_label'] == outputs[1]['predicted_label'],
100
+ 'rename_same_description': original_desc == renamed_desc,
101
+ 'selected_expected_description': original_desc == descriptions[0],
102
+ 'outputs': outputs})
103
+ return {'scalar_head': client.model.config.num_labels,
104
+ 'interpretation': 'Interface checks and illustrative decisions, not a clinical benchmark or general accuracy estimate.',
105
+ 'checks': checks}
106
+
107
+
108
+ def main():
109
+ p = argparse.ArgumentParser()
110
+ p.add_argument('--audit-only', action='store_true')
111
+ p.add_argument('--interface-only', action='store_true')
112
+ p.add_argument('--model')
113
+ p.add_argument('--revision')
114
+ p.add_argument('--device', default='cpu')
115
+ p.add_argument('--output-dir', type=Path, default=ROOT / 'supplemental-evaluation')
116
+ args = p.parse_args()
117
+ out = args.output_dir; out.mkdir(parents=True, exist_ok=True)
118
+ test = load_test(); baseline = audit(test)
119
+ write_json(out/'per-task-baselines.json', baseline)
120
+ print(json.dumps(baseline), flush=True)
121
+ if args.audit_only: return
122
+ if not args.model: p.error('--model is required for checkpoint evaluation')
123
+ local = Path(args.model).is_dir()
124
+ if not local and not args.revision: p.error('Pin --revision for a Hub model')
125
+ if local:
126
+ weights = sorted(Path(args.model).glob('*.safetensors'))
127
+ assert weights, 'Local checkpoint must have saved safetensors'
128
+ h = hashlib.sha256()
129
+ for path in weights:
130
+ with path.open('rb') as f:
131
+ for chunk in iter(lambda:f.read(8*1024*1024), b''): h.update(chunk)
132
+ identity = h.hexdigest()
133
+ else: identity = args.revision
134
+ manifest = {'model': args.model, 'checkpoint_identity': identity, 'dataset_revision': DATA_REV}
135
+ if (out/'manifest.json').exists(): assert json.loads((out/'manifest.json').read_text()) == manifest
136
+ write_json(out/'manifest.json', manifest)
137
+ spec = importlib.util.spec_from_file_location('inference', ROOT/'predict_open_labels.py')
138
+ helper = importlib.util.module_from_spec(spec); spec.loader.exec_module(helper)
139
+ import torch
140
+ torch.set_num_threads(4)
141
+ client = helper.DecisionModel(args.model, device=args.device, revision=args.revision)
142
+ assert client.model.config.num_labels == 1
143
+ write_json(out/'open-label-checks.json', open_labels(client))
144
+ if args.interface_only: return
145
+ rows = [r for r in test if r['task_family'] == 'when_to_call_tool']
146
+ path = out/'when2call-predictions.jsonl'; previous = {}
147
+ if path.exists():
148
+ for line in path.read_text().splitlines():
149
+ if line: item = json.loads(line); previous[item['row_id']] = item
150
+ assert set(previous).issubset({r['row_id'] for r in rows})
151
+ with path.open('a') as f:
152
+ for row in rows:
153
+ if row['row_id'] in previous: continue
154
+ pred = client.decide(state=json.loads(row['state_json']), question=row['question_text'],
155
+ criteria=json.loads(row['criteria_json']), candidate_batch_size=4)
156
+ rec = {'row_id': row['row_id'], 'gold': row['gold_label'], **pred}
157
+ f.write(json.dumps(rec)+'\n'); f.flush(); previous[row['row_id']] = rec
158
+ if len(previous) % 100 == 0: print(f'When2Call {len(previous)}/{len(rows)}', flush=True)
159
+ correct = sum(previous[r['row_id']]['predicted_label'] == r['gold_label'] for r in rows)
160
+ n = len(rows); reference = baseline['families']['when_to_call_tool']['majority_accuracy']
161
+ result = {'task': 'when_to_call_tool', 'correct': correct, 'n': n, 'accuracy': correct/n,
162
+ 'majority_reference': reference, 'lift_percentage_points': 100*(correct/n-reference),
163
+ 'unseen_task_family': True, 'checkpoint_identity': identity,
164
+ 'truncated_rows': sum(x['truncated'] for x in previous.values())}
165
+ write_json(out/'when2call-results.json', result)
166
+ print(json.dumps(result), flush=True)
167
+
168
+
169
+ if __name__ == '__main__': main()
evaluation/open-label-checks.json ADDED
@@ -0,0 +1,1098 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "scalar_head": 1,
3
+ "interpretation": "Interface checks and illustrative decisions, not a clinical benchmark or general accuracy estimate.",
4
+ "checks": [
5
+ {
6
+ "answer_count": 2,
7
+ "interface_pass": true,
8
+ "reorder_same_label": true,
9
+ "rename_same_description": true,
10
+ "selected_expected_description": true,
11
+ "outputs": [
12
+ {
13
+ "predicted_label": "fresh_choice_2_0",
14
+ "allowed_choices": [
15
+ "fresh_choice_2_0",
16
+ "fresh_choice_2_1"
17
+ ],
18
+ "candidates": [
19
+ {
20
+ "label": "fresh_choice_2_0",
21
+ "score": 0.8051604628562927,
22
+ "raw_score": -24.558311462402344,
23
+ "rank": 1
24
+ },
25
+ {
26
+ "label": "fresh_choice_2_1",
27
+ "score": 0.19483955204486847,
28
+ "raw_score": -25.977176666259766,
29
+ "rank": 2
30
+ }
31
+ ],
32
+ "truncated": false,
33
+ "max_sequence_length": 4096,
34
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
35
+ },
36
+ {
37
+ "predicted_label": "fresh_choice_2_0",
38
+ "allowed_choices": [
39
+ "fresh_choice_2_1",
40
+ "fresh_choice_2_0"
41
+ ],
42
+ "candidates": [
43
+ {
44
+ "label": "fresh_choice_2_0",
45
+ "score": 0.8051604628562927,
46
+ "raw_score": -24.558311462402344,
47
+ "rank": 1
48
+ },
49
+ {
50
+ "label": "fresh_choice_2_1",
51
+ "score": 0.19483955204486847,
52
+ "raw_score": -25.977176666259766,
53
+ "rank": 2
54
+ }
55
+ ],
56
+ "truncated": false,
57
+ "max_sequence_length": 4096,
58
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
59
+ },
60
+ {
61
+ "predicted_label": "opaque_2_0",
62
+ "allowed_choices": [
63
+ "opaque_2_0",
64
+ "opaque_2_1"
65
+ ],
66
+ "candidates": [
67
+ {
68
+ "label": "opaque_2_0",
69
+ "score": 0.7122752070426941,
70
+ "raw_score": -22.207761764526367,
71
+ "rank": 1
72
+ },
73
+ {
74
+ "label": "opaque_2_1",
75
+ "score": 0.2877248227596283,
76
+ "raw_score": -23.114221572875977,
77
+ "rank": 2
78
+ }
79
+ ],
80
+ "truncated": false,
81
+ "max_sequence_length": 4096,
82
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
83
+ },
84
+ {
85
+ "predicted_label": "Escalate this damaged parcel to a human support agent.",
86
+ "allowed_choices": [
87
+ "Escalate this damaged parcel to a human support agent.",
88
+ "Search the product catalog for a new item."
89
+ ],
90
+ "candidates": [
91
+ {
92
+ "label": "Escalate this damaged parcel to a human support agent.",
93
+ "score": 0.7943422794342041,
94
+ "raw_score": -19.58828353881836,
95
+ "rank": 1
96
+ },
97
+ {
98
+ "label": "Search the product catalog for a new item.",
99
+ "score": 0.2056577205657959,
100
+ "raw_score": -20.939584732055664,
101
+ "rank": 2
102
+ }
103
+ ],
104
+ "truncated": false,
105
+ "max_sequence_length": 4096,
106
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
107
+ }
108
+ ]
109
+ },
110
+ {
111
+ "answer_count": 3,
112
+ "interface_pass": true,
113
+ "reorder_same_label": true,
114
+ "rename_same_description": false,
115
+ "selected_expected_description": true,
116
+ "outputs": [
117
+ {
118
+ "predicted_label": "fresh_choice_3_0",
119
+ "allowed_choices": [
120
+ "fresh_choice_3_0",
121
+ "fresh_choice_3_1",
122
+ "fresh_choice_3_2"
123
+ ],
124
+ "candidates": [
125
+ {
126
+ "label": "fresh_choice_3_0",
127
+ "score": 0.5637261867523193,
128
+ "raw_score": -24.52130699157715,
129
+ "rank": 1
130
+ },
131
+ {
132
+ "label": "fresh_choice_3_2",
133
+ "score": 0.28383156657218933,
134
+ "raw_score": -25.207494735717773,
135
+ "rank": 2
136
+ },
137
+ {
138
+ "label": "fresh_choice_3_1",
139
+ "score": 0.15244220197200775,
140
+ "raw_score": -25.829090118408203,
141
+ "rank": 3
142
+ }
143
+ ],
144
+ "truncated": false,
145
+ "max_sequence_length": 4096,
146
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
147
+ },
148
+ {
149
+ "predicted_label": "fresh_choice_3_0",
150
+ "allowed_choices": [
151
+ "fresh_choice_3_2",
152
+ "fresh_choice_3_1",
153
+ "fresh_choice_3_0"
154
+ ],
155
+ "candidates": [
156
+ {
157
+ "label": "fresh_choice_3_0",
158
+ "score": 0.563726544380188,
159
+ "raw_score": -24.52130699157715,
160
+ "rank": 1
161
+ },
162
+ {
163
+ "label": "fresh_choice_3_2",
164
+ "score": 0.2838311791419983,
165
+ "raw_score": -25.207496643066406,
166
+ "rank": 2
167
+ },
168
+ {
169
+ "label": "fresh_choice_3_1",
170
+ "score": 0.1524423062801361,
171
+ "raw_score": -25.829090118408203,
172
+ "rank": 3
173
+ }
174
+ ],
175
+ "truncated": false,
176
+ "max_sequence_length": 4096,
177
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
178
+ },
179
+ {
180
+ "predicted_label": "opaque_3_2",
181
+ "allowed_choices": [
182
+ "opaque_3_0",
183
+ "opaque_3_1",
184
+ "opaque_3_2"
185
+ ],
186
+ "candidates": [
187
+ {
188
+ "label": "opaque_3_2",
189
+ "score": 0.5318363308906555,
190
+ "raw_score": -21.825939178466797,
191
+ "rank": 1
192
+ },
193
+ {
194
+ "label": "opaque_3_0",
195
+ "score": 0.34079551696777344,
196
+ "raw_score": -22.270992279052734,
197
+ "rank": 2
198
+ },
199
+ {
200
+ "label": "opaque_3_1",
201
+ "score": 0.12736809253692627,
202
+ "raw_score": -23.25519371032715,
203
+ "rank": 3
204
+ }
205
+ ],
206
+ "truncated": false,
207
+ "max_sequence_length": 4096,
208
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
209
+ },
210
+ {
211
+ "predicted_label": "Route to unrelated department number 0 for a different request.",
212
+ "allowed_choices": [
213
+ "Escalate this damaged parcel to a human support agent.",
214
+ "Search the product catalog for a new item.",
215
+ "Route to unrelated department number 0 for a different request."
216
+ ],
217
+ "candidates": [
218
+ {
219
+ "label": "Route to unrelated department number 0 for a different request.",
220
+ "score": 0.6520880460739136,
221
+ "raw_score": -18.729812622070312,
222
+ "rank": 1
223
+ },
224
+ {
225
+ "label": "Escalate this damaged parcel to a human support agent.",
226
+ "score": 0.2763611972332001,
227
+ "raw_score": -19.58828353881836,
228
+ "rank": 2
229
+ },
230
+ {
231
+ "label": "Search the product catalog for a new item.",
232
+ "score": 0.07155078649520874,
233
+ "raw_score": -20.939584732055664,
234
+ "rank": 3
235
+ }
236
+ ],
237
+ "truncated": false,
238
+ "max_sequence_length": 4096,
239
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
240
+ }
241
+ ]
242
+ },
243
+ {
244
+ "answer_count": 7,
245
+ "interface_pass": true,
246
+ "reorder_same_label": true,
247
+ "rename_same_description": false,
248
+ "selected_expected_description": true,
249
+ "outputs": [
250
+ {
251
+ "predicted_label": "fresh_choice_7_0",
252
+ "allowed_choices": [
253
+ "fresh_choice_7_0",
254
+ "fresh_choice_7_1",
255
+ "fresh_choice_7_2",
256
+ "fresh_choice_7_3",
257
+ "fresh_choice_7_4",
258
+ "fresh_choice_7_5",
259
+ "fresh_choice_7_6"
260
+ ],
261
+ "candidates": [
262
+ {
263
+ "label": "fresh_choice_7_0",
264
+ "score": 0.33595216274261475,
265
+ "raw_score": -23.554637908935547,
266
+ "rank": 1
267
+ },
268
+ {
269
+ "label": "fresh_choice_7_6",
270
+ "score": 0.2605591416358948,
271
+ "raw_score": -23.80877685546875,
272
+ "rank": 2
273
+ },
274
+ {
275
+ "label": "fresh_choice_7_4",
276
+ "score": 0.10382113605737686,
277
+ "raw_score": -24.72893714904785,
278
+ "rank": 3
279
+ },
280
+ {
281
+ "label": "fresh_choice_7_3",
282
+ "score": 0.08751508593559265,
283
+ "raw_score": -24.899795532226562,
284
+ "rank": 4
285
+ },
286
+ {
287
+ "label": "fresh_choice_7_5",
288
+ "score": 0.07923448085784912,
289
+ "raw_score": -24.999195098876953,
290
+ "rank": 5
291
+ },
292
+ {
293
+ "label": "fresh_choice_7_2",
294
+ "score": 0.07844322174787521,
295
+ "raw_score": -25.009231567382812,
296
+ "rank": 6
297
+ },
298
+ {
299
+ "label": "fresh_choice_7_1",
300
+ "score": 0.05447477847337723,
301
+ "raw_score": -25.373868942260742,
302
+ "rank": 7
303
+ }
304
+ ],
305
+ "truncated": false,
306
+ "max_sequence_length": 4096,
307
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
308
+ },
309
+ {
310
+ "predicted_label": "fresh_choice_7_0",
311
+ "allowed_choices": [
312
+ "fresh_choice_7_6",
313
+ "fresh_choice_7_5",
314
+ "fresh_choice_7_4",
315
+ "fresh_choice_7_3",
316
+ "fresh_choice_7_2",
317
+ "fresh_choice_7_1",
318
+ "fresh_choice_7_0"
319
+ ],
320
+ "candidates": [
321
+ {
322
+ "label": "fresh_choice_7_0",
323
+ "score": 0.3359522819519043,
324
+ "raw_score": -23.554637908935547,
325
+ "rank": 1
326
+ },
327
+ {
328
+ "label": "fresh_choice_7_6",
329
+ "score": 0.26055923104286194,
330
+ "raw_score": -23.80877685546875,
331
+ "rank": 2
332
+ },
333
+ {
334
+ "label": "fresh_choice_7_4",
335
+ "score": 0.10382117331027985,
336
+ "raw_score": -24.72893714904785,
337
+ "rank": 3
338
+ },
339
+ {
340
+ "label": "fresh_choice_7_3",
341
+ "score": 0.08751478046178818,
342
+ "raw_score": -24.899799346923828,
343
+ "rank": 4
344
+ },
345
+ {
346
+ "label": "fresh_choice_7_5",
347
+ "score": 0.07923451066017151,
348
+ "raw_score": -24.999195098876953,
349
+ "rank": 5
350
+ },
351
+ {
352
+ "label": "fresh_choice_7_2",
353
+ "score": 0.0784432515501976,
354
+ "raw_score": -25.009231567382812,
355
+ "rank": 6
356
+ },
357
+ {
358
+ "label": "fresh_choice_7_1",
359
+ "score": 0.05447479709982872,
360
+ "raw_score": -25.373868942260742,
361
+ "rank": 7
362
+ }
363
+ ],
364
+ "truncated": false,
365
+ "max_sequence_length": 4096,
366
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
367
+ },
368
+ {
369
+ "predicted_label": "opaque_7_6",
370
+ "allowed_choices": [
371
+ "opaque_7_0",
372
+ "opaque_7_1",
373
+ "opaque_7_2",
374
+ "opaque_7_3",
375
+ "opaque_7_4",
376
+ "opaque_7_5",
377
+ "opaque_7_6"
378
+ ],
379
+ "candidates": [
380
+ {
381
+ "label": "opaque_7_6",
382
+ "score": 0.3010425269603729,
383
+ "raw_score": -21.268993377685547,
384
+ "rank": 1
385
+ },
386
+ {
387
+ "label": "opaque_7_0",
388
+ "score": 0.22324608266353607,
389
+ "raw_score": -21.567970275878906,
390
+ "rank": 2
391
+ },
392
+ {
393
+ "label": "opaque_7_4",
394
+ "score": 0.114542156457901,
395
+ "raw_score": -22.235301971435547,
396
+ "rank": 3
397
+ },
398
+ {
399
+ "label": "opaque_7_2",
400
+ "score": 0.11444498598575592,
401
+ "raw_score": -22.23615074157715,
402
+ "rank": 4
403
+ },
404
+ {
405
+ "label": "opaque_7_3",
406
+ "score": 0.11412934213876724,
407
+ "raw_score": -22.23891258239746,
408
+ "rank": 5
409
+ },
410
+ {
411
+ "label": "opaque_7_5",
412
+ "score": 0.07986637204885483,
413
+ "raw_score": -22.595890045166016,
414
+ "rank": 6
415
+ },
416
+ {
417
+ "label": "opaque_7_1",
418
+ "score": 0.05272847041487694,
419
+ "raw_score": -23.011089324951172,
420
+ "rank": 7
421
+ }
422
+ ],
423
+ "truncated": false,
424
+ "max_sequence_length": 4096,
425
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
426
+ },
427
+ {
428
+ "predicted_label": "Route to unrelated department number 0 for a different request.",
429
+ "allowed_choices": [
430
+ "Escalate this damaged parcel to a human support agent.",
431
+ "Search the product catalog for a new item.",
432
+ "Route to unrelated department number 0 for a different request.",
433
+ "Route to unrelated department number 1 for a different request.",
434
+ "Route to unrelated department number 2 for a different request.",
435
+ "Route to unrelated department number 3 for a different request.",
436
+ "Route to unrelated department number 4 for a different request."
437
+ ],
438
+ "candidates": [
439
+ {
440
+ "label": "Route to unrelated department number 0 for a different request.",
441
+ "score": 0.2675093412399292,
442
+ "raw_score": -18.729812622070312,
443
+ "rank": 1
444
+ },
445
+ {
446
+ "label": "Route to unrelated department number 4 for a different request.",
447
+ "score": 0.2286183089017868,
448
+ "raw_score": -18.886913299560547,
449
+ "rank": 2
450
+ },
451
+ {
452
+ "label": "Route to unrelated department number 1 for a different request.",
453
+ "score": 0.1698872447013855,
454
+ "raw_score": -19.1838321685791,
455
+ "rank": 3
456
+ },
457
+ {
458
+ "label": "Escalate this damaged parcel to a human support agent.",
459
+ "score": 0.1133730337023735,
460
+ "raw_score": -19.58828353881836,
461
+ "rank": 4
462
+ },
463
+ {
464
+ "label": "Route to unrelated department number 3 for a different request.",
465
+ "score": 0.10056313127279282,
466
+ "raw_score": -19.708181381225586,
467
+ "rank": 5
468
+ },
469
+ {
470
+ "label": "Route to unrelated department number 2 for a different request.",
471
+ "score": 0.09069626778364182,
472
+ "raw_score": -19.811450958251953,
473
+ "rank": 6
474
+ },
475
+ {
476
+ "label": "Search the product catalog for a new item.",
477
+ "score": 0.029352637007832527,
478
+ "raw_score": -20.939584732055664,
479
+ "rank": 7
480
+ }
481
+ ],
482
+ "truncated": false,
483
+ "max_sequence_length": 4096,
484
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
485
+ }
486
+ ]
487
+ },
488
+ {
489
+ "answer_count": 20,
490
+ "interface_pass": true,
491
+ "reorder_same_label": true,
492
+ "rename_same_description": true,
493
+ "selected_expected_description": false,
494
+ "outputs": [
495
+ {
496
+ "predicted_label": "fresh_choice_20_11",
497
+ "allowed_choices": [
498
+ "fresh_choice_20_0",
499
+ "fresh_choice_20_1",
500
+ "fresh_choice_20_2",
501
+ "fresh_choice_20_3",
502
+ "fresh_choice_20_4",
503
+ "fresh_choice_20_5",
504
+ "fresh_choice_20_6",
505
+ "fresh_choice_20_7",
506
+ "fresh_choice_20_8",
507
+ "fresh_choice_20_9",
508
+ "fresh_choice_20_10",
509
+ "fresh_choice_20_11",
510
+ "fresh_choice_20_12",
511
+ "fresh_choice_20_13",
512
+ "fresh_choice_20_14",
513
+ "fresh_choice_20_15",
514
+ "fresh_choice_20_16",
515
+ "fresh_choice_20_17",
516
+ "fresh_choice_20_18",
517
+ "fresh_choice_20_19"
518
+ ],
519
+ "candidates": [
520
+ {
521
+ "label": "fresh_choice_20_11",
522
+ "score": 0.15572811663150787,
523
+ "raw_score": -24.59111785888672,
524
+ "rank": 1
525
+ },
526
+ {
527
+ "label": "fresh_choice_20_10",
528
+ "score": 0.1482214629650116,
529
+ "raw_score": -24.640522003173828,
530
+ "rank": 2
531
+ },
532
+ {
533
+ "label": "fresh_choice_20_0",
534
+ "score": 0.10967609286308289,
535
+ "raw_score": -24.94169807434082,
536
+ "rank": 3
537
+ },
538
+ {
539
+ "label": "fresh_choice_20_9",
540
+ "score": 0.09642379730939865,
541
+ "raw_score": -25.070476531982422,
542
+ "rank": 4
543
+ },
544
+ {
545
+ "label": "fresh_choice_20_6",
546
+ "score": 0.062188707292079926,
547
+ "raw_score": -25.509056091308594,
548
+ "rank": 5
549
+ },
550
+ {
551
+ "label": "fresh_choice_20_8",
552
+ "score": 0.057998985052108765,
553
+ "raw_score": -25.57880401611328,
554
+ "rank": 6
555
+ },
556
+ {
557
+ "label": "fresh_choice_20_7",
558
+ "score": 0.05253210663795471,
559
+ "raw_score": -25.677804946899414,
560
+ "rank": 7
561
+ },
562
+ {
563
+ "label": "fresh_choice_20_13",
564
+ "score": 0.038757286965847015,
565
+ "raw_score": -25.981910705566406,
566
+ "rank": 8
567
+ },
568
+ {
569
+ "label": "fresh_choice_20_14",
570
+ "score": 0.03664031997323036,
571
+ "raw_score": -26.0380802154541,
572
+ "rank": 9
573
+ },
574
+ {
575
+ "label": "fresh_choice_20_12",
576
+ "score": 0.03644376993179321,
577
+ "raw_score": -26.043458938598633,
578
+ "rank": 10
579
+ },
580
+ {
581
+ "label": "fresh_choice_20_2",
582
+ "score": 0.026935890316963196,
583
+ "raw_score": -26.34576988220215,
584
+ "rank": 11
585
+ },
586
+ {
587
+ "label": "fresh_choice_20_4",
588
+ "score": 0.024683654308319092,
589
+ "raw_score": -26.433088302612305,
590
+ "rank": 12
591
+ },
592
+ {
593
+ "label": "fresh_choice_20_5",
594
+ "score": 0.024211952462792397,
595
+ "raw_score": -26.452383041381836,
596
+ "rank": 13
597
+ },
598
+ {
599
+ "label": "fresh_choice_20_19",
600
+ "score": 0.023727986961603165,
601
+ "raw_score": -26.47257423400879,
602
+ "rank": 14
603
+ },
604
+ {
605
+ "label": "fresh_choice_20_3",
606
+ "score": 0.022604424506425858,
607
+ "raw_score": -26.52108383178711,
608
+ "rank": 15
609
+ },
610
+ {
611
+ "label": "fresh_choice_20_1",
612
+ "score": 0.021840117871761322,
613
+ "raw_score": -26.55548095703125,
614
+ "rank": 16
615
+ },
616
+ {
617
+ "label": "fresh_choice_20_15",
618
+ "score": 0.02026970498263836,
619
+ "raw_score": -26.630102157592773,
620
+ "rank": 17
621
+ },
622
+ {
623
+ "label": "fresh_choice_20_16",
624
+ "score": 0.019944682717323303,
625
+ "raw_score": -26.64626693725586,
626
+ "rank": 18
627
+ },
628
+ {
629
+ "label": "fresh_choice_20_18",
630
+ "score": 0.011263599619269371,
631
+ "raw_score": -27.217653274536133,
632
+ "rank": 19
633
+ },
634
+ {
635
+ "label": "fresh_choice_20_17",
636
+ "score": 0.00990742165595293,
637
+ "raw_score": -27.345945358276367,
638
+ "rank": 20
639
+ }
640
+ ],
641
+ "truncated": false,
642
+ "max_sequence_length": 4096,
643
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
644
+ },
645
+ {
646
+ "predicted_label": "fresh_choice_20_11",
647
+ "allowed_choices": [
648
+ "fresh_choice_20_19",
649
+ "fresh_choice_20_18",
650
+ "fresh_choice_20_17",
651
+ "fresh_choice_20_16",
652
+ "fresh_choice_20_15",
653
+ "fresh_choice_20_14",
654
+ "fresh_choice_20_13",
655
+ "fresh_choice_20_12",
656
+ "fresh_choice_20_11",
657
+ "fresh_choice_20_10",
658
+ "fresh_choice_20_9",
659
+ "fresh_choice_20_8",
660
+ "fresh_choice_20_7",
661
+ "fresh_choice_20_6",
662
+ "fresh_choice_20_5",
663
+ "fresh_choice_20_4",
664
+ "fresh_choice_20_3",
665
+ "fresh_choice_20_2",
666
+ "fresh_choice_20_1",
667
+ "fresh_choice_20_0"
668
+ ],
669
+ "candidates": [
670
+ {
671
+ "label": "fresh_choice_20_11",
672
+ "score": 0.15572810173034668,
673
+ "raw_score": -24.59111785888672,
674
+ "rank": 1
675
+ },
676
+ {
677
+ "label": "fresh_choice_20_10",
678
+ "score": 0.1482214480638504,
679
+ "raw_score": -24.640522003173828,
680
+ "rank": 2
681
+ },
682
+ {
683
+ "label": "fresh_choice_20_0",
684
+ "score": 0.10967607796192169,
685
+ "raw_score": -24.94169807434082,
686
+ "rank": 3
687
+ },
688
+ {
689
+ "label": "fresh_choice_20_9",
690
+ "score": 0.09642378240823746,
691
+ "raw_score": -25.070476531982422,
692
+ "rank": 4
693
+ },
694
+ {
695
+ "label": "fresh_choice_20_6",
696
+ "score": 0.06218870356678963,
697
+ "raw_score": -25.509056091308594,
698
+ "rank": 5
699
+ },
700
+ {
701
+ "label": "fresh_choice_20_8",
702
+ "score": 0.057998981326818466,
703
+ "raw_score": -25.57880401611328,
704
+ "rank": 6
705
+ },
706
+ {
707
+ "label": "fresh_choice_20_7",
708
+ "score": 0.05253210291266441,
709
+ "raw_score": -25.677804946899414,
710
+ "rank": 7
711
+ },
712
+ {
713
+ "label": "fresh_choice_20_13",
714
+ "score": 0.03875728324055672,
715
+ "raw_score": -25.981910705566406,
716
+ "rank": 8
717
+ },
718
+ {
719
+ "label": "fresh_choice_20_14",
720
+ "score": 0.036640316247940063,
721
+ "raw_score": -26.0380802154541,
722
+ "rank": 9
723
+ },
724
+ {
725
+ "label": "fresh_choice_20_12",
726
+ "score": 0.036443766206502914,
727
+ "raw_score": -26.043458938598633,
728
+ "rank": 10
729
+ },
730
+ {
731
+ "label": "fresh_choice_20_2",
732
+ "score": 0.026935888454318047,
733
+ "raw_score": -26.34576988220215,
734
+ "rank": 11
735
+ },
736
+ {
737
+ "label": "fresh_choice_20_4",
738
+ "score": 0.024683650583028793,
739
+ "raw_score": -26.433088302612305,
740
+ "rank": 12
741
+ },
742
+ {
743
+ "label": "fresh_choice_20_5",
744
+ "score": 0.024211950600147247,
745
+ "raw_score": -26.452383041381836,
746
+ "rank": 13
747
+ },
748
+ {
749
+ "label": "fresh_choice_20_19",
750
+ "score": 0.023727985098958015,
751
+ "raw_score": -26.47257423400879,
752
+ "rank": 14
753
+ },
754
+ {
755
+ "label": "fresh_choice_20_3",
756
+ "score": 0.02260442264378071,
757
+ "raw_score": -26.52108383178711,
758
+ "rank": 15
759
+ },
760
+ {
761
+ "label": "fresh_choice_20_1",
762
+ "score": 0.021840116009116173,
763
+ "raw_score": -26.55548095703125,
764
+ "rank": 16
765
+ },
766
+ {
767
+ "label": "fresh_choice_20_15",
768
+ "score": 0.02026970311999321,
769
+ "raw_score": -26.630102157592773,
770
+ "rank": 17
771
+ },
772
+ {
773
+ "label": "fresh_choice_20_16",
774
+ "score": 0.019944680854678154,
775
+ "raw_score": -26.64626693725586,
776
+ "rank": 18
777
+ },
778
+ {
779
+ "label": "fresh_choice_20_18",
780
+ "score": 0.011263598687946796,
781
+ "raw_score": -27.217653274536133,
782
+ "rank": 19
783
+ },
784
+ {
785
+ "label": "fresh_choice_20_17",
786
+ "score": 0.009907420724630356,
787
+ "raw_score": -27.345945358276367,
788
+ "rank": 20
789
+ }
790
+ ],
791
+ "truncated": false,
792
+ "max_sequence_length": 4096,
793
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
794
+ },
795
+ {
796
+ "predicted_label": "opaque_20_11",
797
+ "allowed_choices": [
798
+ "opaque_20_0",
799
+ "opaque_20_1",
800
+ "opaque_20_2",
801
+ "opaque_20_3",
802
+ "opaque_20_4",
803
+ "opaque_20_5",
804
+ "opaque_20_6",
805
+ "opaque_20_7",
806
+ "opaque_20_8",
807
+ "opaque_20_9",
808
+ "opaque_20_10",
809
+ "opaque_20_11",
810
+ "opaque_20_12",
811
+ "opaque_20_13",
812
+ "opaque_20_14",
813
+ "opaque_20_15",
814
+ "opaque_20_16",
815
+ "opaque_20_17",
816
+ "opaque_20_18",
817
+ "opaque_20_19"
818
+ ],
819
+ "candidates": [
820
+ {
821
+ "label": "opaque_20_11",
822
+ "score": 0.17951472103595734,
823
+ "raw_score": -21.678247451782227,
824
+ "rank": 1
825
+ },
826
+ {
827
+ "label": "opaque_20_10",
828
+ "score": 0.16386361420154572,
829
+ "raw_score": -21.76947021484375,
830
+ "rank": 2
831
+ },
832
+ {
833
+ "label": "opaque_20_9",
834
+ "score": 0.1559082716703415,
835
+ "raw_score": -21.819236755371094,
836
+ "rank": 3
837
+ },
838
+ {
839
+ "label": "opaque_20_8",
840
+ "score": 0.06722499430179596,
841
+ "raw_score": -22.660459518432617,
842
+ "rank": 4
843
+ },
844
+ {
845
+ "label": "opaque_20_6",
846
+ "score": 0.061009395867586136,
847
+ "raw_score": -22.757476806640625,
848
+ "rank": 5
849
+ },
850
+ {
851
+ "label": "opaque_20_0",
852
+ "score": 0.05371956154704094,
853
+ "raw_score": -22.884727478027344,
854
+ "rank": 6
855
+ },
856
+ {
857
+ "label": "opaque_20_7",
858
+ "score": 0.046380579471588135,
859
+ "raw_score": -23.03162384033203,
860
+ "rank": 7
861
+ },
862
+ {
863
+ "label": "opaque_20_2",
864
+ "score": 0.031455062329769135,
865
+ "raw_score": -23.419944763183594,
866
+ "rank": 8
867
+ },
868
+ {
869
+ "label": "opaque_20_12",
870
+ "score": 0.0313202403485775,
871
+ "raw_score": -23.424240112304688,
872
+ "rank": 9
873
+ },
874
+ {
875
+ "label": "opaque_20_14",
876
+ "score": 0.029301172122359276,
877
+ "raw_score": -23.490877151489258,
878
+ "rank": 10
879
+ },
880
+ {
881
+ "label": "opaque_20_13",
882
+ "score": 0.026141302660107613,
883
+ "raw_score": -23.60498809814453,
884
+ "rank": 11
885
+ },
886
+ {
887
+ "label": "opaque_20_19",
888
+ "score": 0.02265744097530842,
889
+ "raw_score": -23.748016357421875,
890
+ "rank": 12
891
+ },
892
+ {
893
+ "label": "opaque_20_4",
894
+ "score": 0.022092893719673157,
895
+ "raw_score": -23.77324867248535,
896
+ "rank": 13
897
+ },
898
+ {
899
+ "label": "opaque_20_3",
900
+ "score": 0.02106008678674698,
901
+ "raw_score": -23.821125030517578,
902
+ "rank": 14
903
+ },
904
+ {
905
+ "label": "opaque_20_5",
906
+ "score": 0.019397132098674774,
907
+ "raw_score": -23.903379440307617,
908
+ "rank": 15
909
+ },
910
+ {
911
+ "label": "opaque_20_1",
912
+ "score": 0.018865037709474564,
913
+ "raw_score": -23.931194305419922,
914
+ "rank": 16
915
+ },
916
+ {
917
+ "label": "opaque_20_15",
918
+ "score": 0.016952816396951675,
919
+ "raw_score": -24.038070678710938,
920
+ "rank": 17
921
+ },
922
+ {
923
+ "label": "opaque_20_16",
924
+ "score": 0.013765859417617321,
925
+ "raw_score": -24.246313095092773,
926
+ "rank": 18
927
+ },
928
+ {
929
+ "label": "opaque_20_18",
930
+ "score": 0.012271336279809475,
931
+ "raw_score": -24.361238479614258,
932
+ "rank": 19
933
+ },
934
+ {
935
+ "label": "opaque_20_17",
936
+ "score": 0.007098542992025614,
937
+ "raw_score": -24.908615112304688,
938
+ "rank": 20
939
+ }
940
+ ],
941
+ "truncated": false,
942
+ "max_sequence_length": 4096,
943
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
944
+ },
945
+ {
946
+ "predicted_label": "Route to unrelated department number 9 for a different request.",
947
+ "allowed_choices": [
948
+ "Escalate this damaged parcel to a human support agent.",
949
+ "Search the product catalog for a new item.",
950
+ "Route to unrelated department number 0 for a different request.",
951
+ "Route to unrelated department number 1 for a different request.",
952
+ "Route to unrelated department number 2 for a different request.",
953
+ "Route to unrelated department number 3 for a different request.",
954
+ "Route to unrelated department number 4 for a different request.",
955
+ "Route to unrelated department number 5 for a different request.",
956
+ "Route to unrelated department number 6 for a different request.",
957
+ "Route to unrelated department number 7 for a different request.",
958
+ "Route to unrelated department number 8 for a different request.",
959
+ "Route to unrelated department number 9 for a different request.",
960
+ "Route to unrelated department number 10 for a different request.",
961
+ "Route to unrelated department number 11 for a different request.",
962
+ "Route to unrelated department number 12 for a different request.",
963
+ "Route to unrelated department number 13 for a different request.",
964
+ "Route to unrelated department number 14 for a different request.",
965
+ "Route to unrelated department number 15 for a different request.",
966
+ "Route to unrelated department number 16 for a different request.",
967
+ "Route to unrelated department number 17 for a different request."
968
+ ],
969
+ "candidates": [
970
+ {
971
+ "label": "Route to unrelated department number 9 for a different request.",
972
+ "score": 0.11579183489084244,
973
+ "raw_score": -18.278705596923828,
974
+ "rank": 1
975
+ },
976
+ {
977
+ "label": "Route to unrelated department number 8 for a different request.",
978
+ "score": 0.11046068370342255,
979
+ "raw_score": -18.32583999633789,
980
+ "rank": 2
981
+ },
982
+ {
983
+ "label": "Route to unrelated department number 11 for a different request.",
984
+ "score": 0.09375981986522675,
985
+ "raw_score": -18.489763259887695,
986
+ "rank": 3
987
+ },
988
+ {
989
+ "label": "Route to unrelated department number 7 for a different request.",
990
+ "score": 0.07637739926576614,
991
+ "raw_score": -18.694812774658203,
992
+ "rank": 4
993
+ },
994
+ {
995
+ "label": "Route to unrelated department number 0 for a different request.",
996
+ "score": 0.07375044375658035,
997
+ "raw_score": -18.729812622070312,
998
+ "rank": 5
999
+ },
1000
+ {
1001
+ "label": "Route to unrelated department number 10 for a different request.",
1002
+ "score": 0.06551573425531387,
1003
+ "raw_score": -18.848209381103516,
1004
+ "rank": 6
1005
+ },
1006
+ {
1007
+ "label": "Route to unrelated department number 4 for a different request.",
1008
+ "score": 0.06302846223115921,
1009
+ "raw_score": -18.886913299560547,
1010
+ "rank": 7
1011
+ },
1012
+ {
1013
+ "label": "Route to unrelated department number 6 for a different request.",
1014
+ "score": 0.051630470901727676,
1015
+ "raw_score": -19.086387634277344,
1016
+ "rank": 8
1017
+ },
1018
+ {
1019
+ "label": "Route to unrelated department number 12 for a different request.",
1020
+ "score": 0.04935718700289726,
1021
+ "raw_score": -19.13141632080078,
1022
+ "rank": 9
1023
+ },
1024
+ {
1025
+ "label": "Route to unrelated department number 1 for a different request.",
1026
+ "score": 0.046836718916893005,
1027
+ "raw_score": -19.1838321685791,
1028
+ "rank": 10
1029
+ },
1030
+ {
1031
+ "label": "Route to unrelated department number 16 for a different request.",
1032
+ "score": 0.03468938171863556,
1033
+ "raw_score": -19.484066009521484,
1034
+ "rank": 11
1035
+ },
1036
+ {
1037
+ "label": "Route to unrelated department number 13 for a different request.",
1038
+ "score": 0.03264395520091057,
1039
+ "raw_score": -19.54483985900879,
1040
+ "rank": 12
1041
+ },
1042
+ {
1043
+ "label": "Escalate this damaged parcel to a human support agent.",
1044
+ "score": 0.03125615045428276,
1045
+ "raw_score": -19.58828353881836,
1046
+ "rank": 13
1047
+ },
1048
+ {
1049
+ "label": "Route to unrelated department number 14 for a different request.",
1050
+ "score": 0.03116500936448574,
1051
+ "raw_score": -19.591203689575195,
1052
+ "rank": 14
1053
+ },
1054
+ {
1055
+ "label": "Route to unrelated department number 5 for a different request.",
1056
+ "score": 0.02856743521988392,
1057
+ "raw_score": -19.678232192993164,
1058
+ "rank": 15
1059
+ },
1060
+ {
1061
+ "label": "Route to unrelated department number 3 for a different request.",
1062
+ "score": 0.027724549174308777,
1063
+ "raw_score": -19.708181381225586,
1064
+ "rank": 16
1065
+ },
1066
+ {
1067
+ "label": "Route to unrelated department number 2 for a different request.",
1068
+ "score": 0.025004321709275246,
1069
+ "raw_score": -19.811450958251953,
1070
+ "rank": 17
1071
+ },
1072
+ {
1073
+ "label": "Route to unrelated department number 15 for a different request.",
1074
+ "score": 0.02304059825837612,
1075
+ "raw_score": -19.89324188232422,
1076
+ "rank": 18
1077
+ },
1078
+ {
1079
+ "label": "Route to unrelated department number 17 for a different request.",
1080
+ "score": 0.011307581327855587,
1081
+ "raw_score": -20.605026245117188,
1082
+ "rank": 19
1083
+ },
1084
+ {
1085
+ "label": "Search the product catalog for a new item.",
1086
+ "score": 0.008092314936220646,
1087
+ "raw_score": -20.939584732055664,
1088
+ "rank": 20
1089
+ }
1090
+ ],
1091
+ "truncated": false,
1092
+ "max_sequence_length": 4096,
1093
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
1094
+ }
1095
+ ]
1096
+ }
1097
+ ]
1098
+ }
evaluation/per-task-baselines.json ADDED
@@ -0,0 +1,67 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "dataset_revision": "f2fb14e4ec977c420f376c08785664cd38763d7e",
3
+ "families": {
4
+ "agent_next_action_type": {
5
+ "n": 1158,
6
+ "majority_correct": 616,
7
+ "majority_accuracy": 0.531951640759931,
8
+ "majority_labels": [
9
+ "text_response"
10
+ ],
11
+ "candidate_count": {
12
+ "median": 3.0,
13
+ "min": 3,
14
+ "max": 3
15
+ },
16
+ "uniform_expected_accuracy": 0.3333333333333333,
17
+ "reference_definition": "Descriptive test-set constant-label majority; no model fitting",
18
+ "unseen_task_family": false,
19
+ "training_fixed_majority": {
20
+ "label": "text_response",
21
+ "correct": 616,
22
+ "n": 1158,
23
+ "majority_label_absent_from_choices": 0
24
+ }
25
+ },
26
+ "tool_selection": {
27
+ "n": 542,
28
+ "majority_correct": 79,
29
+ "majority_accuracy": 0.14575645756457564,
30
+ "majority_labels": [
31
+ "transfer_to_human_agent"
32
+ ],
33
+ "candidate_count": {
34
+ "median": 20.0,
35
+ "min": 11,
36
+ "max": 32
37
+ },
38
+ "uniform_expected_accuracy": 0.051859254651728436,
39
+ "reference_definition": "Descriptive test-set constant-label majority; no model fitting",
40
+ "unseen_task_family": false,
41
+ "training_fixed_majority": {
42
+ "label": "transfer_to_human_agent",
43
+ "correct": 79,
44
+ "n": 542,
45
+ "majority_label_absent_from_choices": 235
46
+ }
47
+ },
48
+ "when_to_call_tool": {
49
+ "n": 3652,
50
+ "majority_correct": 1295,
51
+ "majority_accuracy": 0.3546002190580504,
52
+ "majority_labels": [
53
+ "cannot_answer",
54
+ "tool_call"
55
+ ],
56
+ "candidate_count": {
57
+ "median": 4.0,
58
+ "min": 4,
59
+ "max": 4
60
+ },
61
+ "uniform_expected_accuracy": 0.25,
62
+ "reference_definition": "Descriptive test-set constant-label majority; no model fitting",
63
+ "unseen_task_family": true
64
+ }
65
+ },
66
+ "when2call_training_rows": 0
67
+ }
evaluation/when2call-predictions.jsonl ADDED
The diff for this file is too large to render. See raw diff
 
evaluation/when2call-results.json ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "task": "when_to_call_tool",
3
+ "correct": 1256,
4
+ "n": 3652,
5
+ "accuracy": 0.343921139101862,
6
+ "majority_reference": 0.3546002190580504,
7
+ "lift_percentage_points": -1.0679079956188386,
8
+ "unseen_task_family": true,
9
+ "checkpoint_identity": "cd9b0bfaea164d5ad186bca3a4b4aaf3cb4896c18fae14c488be597afa82f0cc",
10
+ "truncated_rows": 0
11
+ }
example-output.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "predicted_label": "lookup_order",
3
+ "allowed_choices": [
4
+ "lookup_order",
5
+ "search_catalog"
6
+ ],
7
+ "candidates": [
8
+ {
9
+ "label": "lookup_order",
10
+ "score": 0.7431679368019104,
11
+ "raw_score": -13.1875,
12
+ "rank": 1
13
+ },
14
+ {
15
+ "label": "search_catalog",
16
+ "score": 0.25683197379112244,
17
+ "raw_score": -14.25,
18
+ "rank": 2
19
+ }
20
+ ],
21
+ "truncated": false,
22
+ "max_sequence_length": 4096,
23
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
24
+ }
metrics.jsonl ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {"event": "env", "mode": "prototype", "device": "cuda", "gpu": "NVIDIA A100-SXM4-80GB", "torch": "2.12.0+cu126", "transformers": "5.17.0", "attention": "kernels-community/flash-attn2@f50dc99ed079b35990bc895d43fd353ea0cb376d", "ts": 1790774736.9}
2
+ {"event": "training_api_validated", "mode": "prototype", "ts": 1790774737.1}
3
+ {"event": "data_prep", "n_rows": 60000, "n_a": 40809, "n_t": 19191, "family_counts": {"agent_next_action_type": 40809, "tool_selection": 19191}, "gold_class_counts": {"agent_next_action_type": {"tool_call": 19173, "text_response": 21636}, "tool_selection": {"distinct": 2840}}, "selected_row_ids_sha256": "c5b306b0471ba104161051a7241b3bce4b69e1d959ff3e34fb06ef8eb4b077d9", "seconds": 25.0, "ts": 1790774762.8}
4
+ {"event": "lazy_pairs_ready", "tag": "train", "n_rows": 60000, "n_pairs": 199191, "max_len": 4096, "mapped_pairs": 0, "ts": 1790774763.4}
5
+ {"event": "prototype_start", "pool": "sampled4", "n_rows_selected": 60000, "n_a": 40809, "n_t": 19191, "ts": 1790774763.9}
6
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 599, "n_pairs": 4767, "max_len": 4096, "mapped_pairs": 0, "ts": 1790774765.7}
7
+ {"event": "coverage_progress", "step": 1, "rows_seen": 17, "target": 60000, "elapsed_seconds": 44.8, "ts": 1790774781.7}
8
+ {"event": "coverage_progress", "step": 2, "rows_seen": 33, "target": 60000, "elapsed_seconds": 46.2, "ts": 1790774783.0}
9
+ {"event": "coverage_progress", "step": 3, "rows_seen": 49, "target": 60000, "elapsed_seconds": 47.5, "ts": 1790774784.4}
10
+ {"event": "coverage_progress", "step": 4, "rows_seen": 65, "target": 60000, "elapsed_seconds": 48.9, "ts": 1790774785.7}
11
+ {"event": "coverage_progress", "step": 5, "rows_seen": 81, "target": 60000, "elapsed_seconds": 50.4, "ts": 1790774787.2}
12
+ {"event": "coverage_progress", "step": 100, "rows_seen": 1618, "target": 60000, "elapsed_seconds": 226.6, "ts": 1790774963.4}
13
+ {"event": "coverage_progress", "step": 200, "rows_seen": 3228, "target": 60000, "elapsed_seconds": 414.1, "ts": 1790775150.9}
14
+ {"event": "coverage_progress", "step": 300, "rows_seen": 4853, "target": 60000, "elapsed_seconds": 602.0, "ts": 1790775338.8}
15
+ {"event": "coverage_progress", "step": 400, "rows_seen": 6455, "target": 60000, "elapsed_seconds": 787.4, "ts": 1790775524.2}
16
+ {"event": "val_acc", "step": 400, "n_rows": 599, "acc": 0.5008347245409015, "n": 599, "macro_accuracy": 0.44902728815772297, "per_family": {"agent_next_action_type": {"acc": 0.5845410628019324, "n": 414}, "tool_selection": {"acc": 0.31351351351351353, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 86.2, "ts": 1790775610.4}
17
+ {"event": "coverage_progress", "step": 500, "rows_seen": 8078, "target": 60000, "elapsed_seconds": 1059.2, "ts": 1790775796.0}
18
+ {"event": "coverage_progress", "step": 600, "rows_seen": 9706, "target": 60000, "elapsed_seconds": 1250.8, "ts": 1790775987.7}
19
+ {"event": "coverage_progress", "step": 700, "rows_seen": 11313, "target": 60000, "elapsed_seconds": 1436.6, "ts": 1790776173.5}
20
+ {"event": "coverage_progress", "step": 800, "rows_seen": 12925, "target": 60000, "elapsed_seconds": 1623.0, "ts": 1790776359.8}
21
+ {"event": "val_acc", "step": 800, "n_rows": 599, "acc": 0.5709515859766278, "n": 599, "macro_accuracy": 0.5147016581799191, "per_family": {"agent_next_action_type": {"acc": 0.6618357487922706, "n": 414}, "tool_selection": {"acc": 0.3675675675675676, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 86.8, "ts": 1790776446.6}
22
+ {"event": "coverage_progress", "step": 900, "rows_seen": 14539, "target": 60000, "elapsed_seconds": 1896.8, "ts": 1790776633.6}
23
+ {"event": "coverage_progress", "step": 1000, "rows_seen": 16165, "target": 60000, "elapsed_seconds": 2084.6, "ts": 1790776821.4}
24
+ {"event": "coverage_progress", "step": 1100, "rows_seen": 17789, "target": 60000, "elapsed_seconds": 2274.0, "ts": 1790777010.8}
25
+ {"event": "coverage_progress", "step": 1200, "rows_seen": 19410, "target": 60000, "elapsed_seconds": 2463.7, "ts": 1790777200.5}
26
+ {"event": "val_acc", "step": 1200, "n_rows": 599, "acc": 0.6010016694490818, "n": 599, "macro_accuracy": 0.557370413892153, "per_family": {"agent_next_action_type": {"acc": 0.6714975845410628, "n": 414}, "tool_selection": {"acc": 0.44324324324324327, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 86.5, "ts": 1790777287.0}
27
+ {"event": "coverage_progress", "step": 1300, "rows_seen": 21023, "target": 60000, "elapsed_seconds": 2741.0, "ts": 1790777477.9}
28
+ {"event": "coverage_progress", "step": 1400, "rows_seen": 22641, "target": 60000, "elapsed_seconds": 2927.7, "ts": 1790777664.6}
29
+ {"event": "coverage_progress", "step": 1500, "rows_seen": 24251, "target": 60000, "elapsed_seconds": 3116.2, "ts": 1790777853.0}
30
+ {"event": "coverage_progress", "step": 1600, "rows_seen": 25854, "target": 60000, "elapsed_seconds": 3308.0, "ts": 1790778044.9}
31
+ {"event": "val_acc", "step": 1600, "n_rows": 599, "acc": 0.6060100166944908, "n": 599, "macro_accuracy": 0.5654785220002612, "per_family": {"agent_next_action_type": {"acc": 0.6714975845410628, "n": 414}, "tool_selection": {"acc": 0.4594594594594595, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 87.7, "ts": 1790778132.6}
32
+ {"event": "coverage_progress", "step": 1700, "rows_seen": 27460, "target": 60000, "elapsed_seconds": 3581.6, "ts": 1790778318.4}
33
+ {"event": "coverage_progress", "step": 1800, "rows_seen": 29067, "target": 60000, "elapsed_seconds": 3769.7, "ts": 1790778506.6}
34
+ {"event": "coverage_progress", "step": 1900, "rows_seen": 30690, "target": 60000, "elapsed_seconds": 3959.4, "ts": 1790778696.2}
35
+ {"event": "coverage_progress", "step": 2000, "rows_seen": 32306, "target": 60000, "elapsed_seconds": 4151.0, "ts": 1790778887.8}
36
+ {"event": "val_acc", "step": 2000, "n_rows": 599, "acc": 0.6227045075125208, "n": 599, "macro_accuracy": 0.5880206293249772, "per_family": {"agent_next_action_type": {"acc": 0.678743961352657, "n": 414}, "tool_selection": {"acc": 0.4972972972972973, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 87.0, "ts": 1790778975.7}
37
+ {"event": "coverage_progress", "step": 2100, "rows_seen": 33913, "target": 60000, "elapsed_seconds": 4425.4, "ts": 1790779162.2}
38
+ {"event": "coverage_progress", "step": 2200, "rows_seen": 35510, "target": 60000, "elapsed_seconds": 4610.9, "ts": 1790779347.7}
39
+ {"event": "coverage_progress", "step": 2300, "rows_seen": 37119, "target": 60000, "elapsed_seconds": 4800.3, "ts": 1790779537.2}
40
+ {"event": "coverage_progress", "step": 2400, "rows_seen": 38723, "target": 60000, "elapsed_seconds": 4986.8, "ts": 1790779723.7}
41
+ {"event": "val_acc", "step": 2400, "n_rows": 599, "acc": 0.659432387312187, "n": 599, "macro_accuracy": 0.6340253296775036, "per_family": {"agent_next_action_type": {"acc": 0.7004830917874396, "n": 414}, "tool_selection": {"acc": 0.5675675675675675, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 86.8, "ts": 1790779810.5}
42
+ {"event": "coverage_progress", "step": 2500, "rows_seen": 40337, "target": 60000, "elapsed_seconds": 5263.2, "ts": 1790780000.0}
43
+ {"event": "coverage_progress", "step": 2600, "rows_seen": 41950, "target": 60000, "elapsed_seconds": 5454.1, "ts": 1790780191.0}
44
+ {"event": "coverage_progress", "step": 2700, "rows_seen": 43568, "target": 60000, "elapsed_seconds": 5643.4, "ts": 1790780380.3}
45
+ {"event": "coverage_progress", "step": 2800, "rows_seen": 45191, "target": 60000, "elapsed_seconds": 5831.1, "ts": 1790780567.9}
46
+ {"event": "val_acc", "step": 2800, "n_rows": 599, "acc": 0.657762938230384, "n": 599, "macro_accuracy": 0.6253427340383862, "per_family": {"agent_next_action_type": {"acc": 0.7101449275362319, "n": 414}, "tool_selection": {"acc": 0.5405405405405406, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 88.0, "ts": 1790780655.9}
47
+ {"event": "coverage_progress", "step": 2900, "rows_seen": 46798, "target": 60000, "elapsed_seconds": 6106.0, "ts": 1790780842.9}
48
+ {"event": "coverage_progress", "step": 3000, "rows_seen": 48405, "target": 60000, "elapsed_seconds": 6295.0, "ts": 1790781031.8}
49
+ {"event": "coverage_progress", "step": 3100, "rows_seen": 50021, "target": 60000, "elapsed_seconds": 6482.9, "ts": 1790781219.8}
50
+ {"event": "coverage_progress", "step": 3200, "rows_seen": 51627, "target": 60000, "elapsed_seconds": 6671.1, "ts": 1790781407.9}
51
+ {"event": "val_acc", "step": 3200, "n_rows": 599, "acc": 0.6944908180300501, "n": 599, "macro_accuracy": 0.6638725682203943, "per_family": {"agent_next_action_type": {"acc": 0.7439613526570048, "n": 414}, "tool_selection": {"acc": 0.5837837837837838, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 86.7, "ts": 1790781494.7}
52
+ {"event": "coverage_progress", "step": 3300, "rows_seen": 53230, "target": 60000, "elapsed_seconds": 6949.0, "ts": 1790781685.9}
53
+ {"event": "coverage_progress", "step": 3400, "rows_seen": 54841, "target": 60000, "elapsed_seconds": 7135.8, "ts": 1790781872.6}
54
+ {"event": "coverage_progress", "step": 3500, "rows_seen": 56445, "target": 60000, "elapsed_seconds": 7323.8, "ts": 1790782060.7}
55
+ {"event": "coverage_progress", "step": 3600, "rows_seen": 58065, "target": 60000, "elapsed_seconds": 7513.4, "ts": 1790782250.3}
56
+ {"event": "val_acc", "step": 3600, "n_rows": 599, "acc": 0.671118530884808, "n": 599, "macro_accuracy": 0.6394894894894895, "per_family": {"agent_next_action_type": {"acc": 0.7222222222222222, "n": 414}, "tool_selection": {"acc": 0.5567567567567567, "n": 185}}, "n_pairs": 4767, "n_truncated_pairs": 158, "seconds": 87.9, "ts": 1790782338.2}
57
+ {"event": "coverage_progress", "step": 3700, "rows_seen": 59683, "target": 60000, "elapsed_seconds": 7790.1, "ts": 1790782527.0}
58
+ {"event": "train_done", "pool": "sampled4", "seconds": 7786.5, "steps": 3720, "rows_covered": 60000, "n_rows_selected": 60000, "n_pairs_processed": 199191, "n_truncated_pairs": 9249, "note": "rows_covered counts unique decisions actually iterated; no full-epoch guarantee", "ts": 1790782565.0}
59
+ {"event": "checkpoint_saved_before_evaluation", "rows_covered": 60000, "ts": 1790782579.5}
60
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 1798, "n_pairs": 14449, "max_len": 4096, "mapped_pairs": 0, "ts": 1790782579.5}
61
+ {"event": "val_full", "acc": 0.6840934371523916, "n": 1798, "macro_accuracy": 0.6529034180751321, "per_family": {"agent_next_action_type": {"acc": 0.7348912167606769, "n": 1241}, "tool_selection": {"acc": 0.5709156193895871, "n": 557}}, "n_pairs": 14449, "n_truncated_pairs": 443, "seconds": 266.6, "ts": 1790782846.2}
62
+ {"event": "test_prep", "n_rows": 1700, "skipped": 0, "ts": 1790782848.2}
63
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 1700, "n_pairs": 14251, "max_len": 4096, "mapped_pairs": 0, "ts": 1790782848.2}
64
+ {"event": "final_test", "acc": 0.6982352941176471, "n": 1700, "macro_accuracy": 0.6778976986661058, "per_family": {"agent_next_action_type": {"acc": 0.7340241796200345, "n": 1158}, "tool_selection": {"acc": 0.6217712177121771, "n": 542}}, "n_pairs": 14251, "n_truncated_pairs": 414, "seconds": 274.1, "ts": 1790783122.4}
65
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 299, "n_pairs": 2523, "max_len": 4096, "mapped_pairs": 0, "ts": 1790783122.4}
66
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 299, "n_pairs": 2523, "max_len": 4096, "mapped_pairs": 0, "ts": 1790783172.5}
67
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 299, "n_pairs": 2523, "max_len": 4096, "mapped_pairs": 0, "ts": 1790783222.7}
68
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 299, "n_pairs": 2523, "max_len": 4096, "mapped_pairs": 0, "ts": 1790783272.9}
69
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 299, "n_pairs": 2523, "max_len": 4096, "mapped_pairs": 0, "ts": 1790783323.2}
70
+ {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 299, "n_pairs": 2523, "max_len": 4096, "mapped_pairs": 0, "ts": 1790783373.7}
71
+ {"event": "shuffled_invariance", "invariance_rate": 1.0, "perms": 5, "n_rows": 299, "ts": 1790783423.8}
72
+ {"event": "baselines", "val": {"uniform_expected": {"acc": 0.24665407679797663, "n": 1798, "per_family": {"agent_next_action_type": {"acc": 0.33333333333332843, "n": 1241}, "tool_selection": {"acc": 0.05353207076499354, "n": 557}}, "macro_accuracy": 0.193432702049161}, "train_frequency": {"acc": 0.47052280311457173, "n": 1798, "per_family": {"agent_next_action_type": {"acc": 0.5511684125705076, "n": 1241}, "tool_selection": {"acc": 0.29084380610412924, "n": 557}}, "macro_accuracy": 0.4210061093373184}}, "test": {"uniform_expected": {"acc": 0.24359277413013672, "n": 1700, "per_family": {"agent_next_action_type": {"acc": 0.33333333333332943, "n": 1158}, "tool_selection": {"acc": 0.051859254651728665, "n": 542}}, "macro_accuracy": 0.19259629399252903}, "train_frequency": {"acc": 0.43529411764705883, "n": 1700, "per_family": {"agent_next_action_type": {"acc": 0.531951640759931, "n": 1158}, "tool_selection": {"acc": 0.22878228782287824, "n": 542}}, "macro_accuracy": 0.3803669642914046}}, "ts": 1790783423.9}
73
+ {"event": "baseline_untrained_head", "val": {"acc": 0.2925472747497219, "n": 1798, "macro_accuracy": 0.21439969214610907, "per_family": {"agent_next_action_type": {"acc": 0.41982272360999195, "n": 1241}, "tool_selection": {"acc": 0.008976660682226212, "n": 557}}, "n_pairs": 14449, "n_truncated_pairs": 443, "seconds": 267.2}, "test": {"acc": 0.29411764705882354, "n": 1700, "macro_accuracy": 0.2178523857777438, "per_family": {"agent_next_action_type": {"acc": 0.4283246977547496, "n": 1158}, "tool_selection": {"acc": 0.007380073800738007, "n": 542}}, "n_pairs": 14251, "n_truncated_pairs": 414, "seconds": 274.4}, "ts": 1790783965.7}
74
+ {"event": "latency_gpu", "measurement": "one candidate forward pass, excludes tokenization", "warmup_pairs": 5, "pair_ms_p50": 23.6, "pair_ms_p95": 24.2, "ts": 1790783969.6}
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cd9b0bfaea164d5ad186bca3a4b4aaf3cb4896c18fae14c488be597afa82f0cc
3
+ size 598436708
per-task-results.json ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "checkpoint_revision": "5623e421346eeaad6a1c7abe959edc4e1f5875e3",
3
+ "weights_sha256": "cd9b0bfaea164d5ad186bca3a4b4aaf3cb4896c18fae14c488be597afa82f0cc",
4
+ "dataset_revision": "f2fb14e4ec977c420f376c08785664cd38763d7e",
5
+ "training_rows": 60000,
6
+ "optimizer_steps": 3720,
7
+ "tasks": {
8
+ "agent_next_action_type": {
9
+ "correct": 850,
10
+ "n": 1158,
11
+ "accuracy": 0.7340241796200345,
12
+ "constant_majority_reference": 0.531951640759931,
13
+ "majority_correct": 616,
14
+ "lift_percentage_points": 20.207253886010356,
15
+ "candidate_count": {
16
+ "median": 3.0,
17
+ "min": 3,
18
+ "max": 3
19
+ },
20
+ "unseen_task_family": false,
21
+ "untrained_modernbert_scalar_head_accuracy": 0.4283246977547496,
22
+ "allowed_choice_training_frequency_accuracy": 0.531951640759931,
23
+ "uniform_expected_accuracy": 0.3333333333333333
24
+ },
25
+ "tool_selection": {
26
+ "correct": 337,
27
+ "n": 542,
28
+ "accuracy": 0.6217712177121771,
29
+ "constant_majority_reference": 0.14575645756457564,
30
+ "majority_correct": 79,
31
+ "lift_percentage_points": 47.601476014760145,
32
+ "candidate_count": {
33
+ "median": 20.0,
34
+ "min": 11,
35
+ "max": 32
36
+ },
37
+ "unseen_task_family": false,
38
+ "untrained_modernbert_scalar_head_accuracy": 0.007380073800738007,
39
+ "allowed_choice_training_frequency_accuracy": 0.22878228782287824,
40
+ "uniform_expected_accuracy": 0.051859254651728436
41
+ },
42
+ "when_to_call_tool": {
43
+ "correct": 1256,
44
+ "n": 3652,
45
+ "accuracy": 0.343921139101862,
46
+ "constant_majority_reference": 0.3546002190580504,
47
+ "majority_correct": 1295,
48
+ "lift_percentage_points": -1.0679079956188386,
49
+ "candidate_count": {
50
+ "median": 4.0,
51
+ "min": 4,
52
+ "max": 4
53
+ },
54
+ "unseen_task_family": true,
55
+ "uniform_expected_accuracy": 0.25
56
+ }
57
+ },
58
+ "open_label_checks": [
59
+ {
60
+ "answer_count": 2,
61
+ "interface_pass": true,
62
+ "reorder_same_label": true,
63
+ "rename_same_description": true,
64
+ "selected_expected_description": true
65
+ },
66
+ {
67
+ "answer_count": 3,
68
+ "interface_pass": true,
69
+ "reorder_same_label": true,
70
+ "rename_same_description": false,
71
+ "selected_expected_description": true
72
+ },
73
+ {
74
+ "answer_count": 7,
75
+ "interface_pass": true,
76
+ "reorder_same_label": true,
77
+ "rename_same_description": false,
78
+ "selected_expected_description": true
79
+ },
80
+ {
81
+ "answer_count": 20,
82
+ "interface_pass": true,
83
+ "reorder_same_label": true,
84
+ "rename_same_description": true,
85
+ "selected_expected_description": false
86
+ }
87
+ ],
88
+ "candidate_order_check": {
89
+ "invariance_rate": 1.0,
90
+ "perms": 5,
91
+ "n_rows": 299
92
+ },
93
+ "cost_estimate_including_failed_jobs_usd": 8.7919,
94
+ "invoice_verified": false
95
+ }
pilot/checkpoint-reload-smoke.json ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "pilot_only": true,
3
+ "coverage_not_prototype": true,
4
+ "checkpoint_parameter_equality": true,
5
+ "same_precision_logit_equality": true,
6
+ "float32_score": 0.1755414605140686,
7
+ "bf16_score": 0.1669921875,
8
+ "absolute_precision_difference": 0.008549273014068604,
9
+ "inference_helper_returned_valid_choice": true,
10
+ "illustrative_prediction": {
11
+ "predicted_label": "search_catalog",
12
+ "allowed_choices": [
13
+ "lookup_order",
14
+ "search_catalog"
15
+ ],
16
+ "candidates": [
17
+ {
18
+ "label": "search_catalog",
19
+ "score": 0.5236632227897644,
20
+ "raw_score": -0.6127176284790039,
21
+ "rank": 1
22
+ },
23
+ {
24
+ "label": "lookup_order",
25
+ "score": 0.476336807012558,
26
+ "raw_score": -0.7074412107467651,
27
+ "rank": 2
28
+ }
29
+ ],
30
+ "truncated": false,
31
+ "max_sequence_length": 4096,
32
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."
33
+ }
34
+ }
pilot/feasibility.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "remaining_seconds": 4407,
3
+ "training_allowance_seconds": 2607,
4
+ "measured_rows_per_s": 18.72,
5
+ "estimated_training_seconds": 3205.1282051282055,
6
+ "full_60000_feasible": false
7
+ }
pilot/pilot_metrics.jsonl ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ {"event": "env", "mode": "pilot", "device": "cuda", "gpu": "NVIDIA H200", "torch": "2.12.0+cu126", "transformers": "5.17.0", "attention": "kernels-community/flash-attn2@f50dc99ed079b35990bc895d43fd353ea0cb376d", "ts": 1790772005.8}
2
+ {"event": "data_prep", "n_rows": 500, "n_a": 340, "n_t": 160, "family_counts": {"agent_next_action_type": 340, "tool_selection": 160}, "gold_class_counts": {"agent_next_action_type": {"text_response": 192, "tool_call": 148}, "tool_selection": {"distinct": 97}}, "selected_row_ids_sha256": "fc9bcbbbfdbfee3d6df4bf7d56a9b658829e70ff7db3039d64fb1bf8add77587", "seconds": 12.8, "ts": 1790772020.4}
3
+ {"event": "tokenize", "tag": "train", "n_pairs": 4205, "n_at_max": 236, "n_truncated": 236, "len_p50": 2548, "len_p95": 4096, "len_max": 4096, "seconds": 6.9, "ts": 1790772027.7}
4
+ {"event": "pilot_phase", "pool": "sampled4", "steps": 40, "seconds": 28.85, "steps_per_s": 1.386, "rows_seen": 500, "rows_per_s": 18.72, "pairs_seen": 1789, "pairs_per_s": 62.0, "tokens_per_s": 165721, "max_mem_gb": 80.82, "ts": 1790772065.7}
5
+ {"event": "latency_gpu", "measurement": "one candidate forward pass, excludes tokenization", "warmup_pairs": 5, "pair_ms_p50": 14.8, "pair_ms_p95": 22.0, "ts": 1790772068.2}
pilot/pilot_results.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "phases": {
3
+ "sampled4": {
4
+ "steps": 40,
5
+ "seconds": 28.85,
6
+ "steps_per_s": 1.386,
7
+ "rows_seen": 500,
8
+ "rows_per_s": 18.72,
9
+ "pairs_seen": 1789,
10
+ "pairs_per_s": 62.0,
11
+ "tokens_per_s": 165721,
12
+ "max_mem_gb": 80.82
13
+ }
14
+ },
15
+ "latency": {
16
+ "measurement": "one candidate forward pass, excludes tokenization",
17
+ "warmup_pairs": 5,
18
+ "pair_ms_p50": 14.8,
19
+ "pair_ms_p95": 22.0
20
+ },
21
+ "env": {
22
+ "gpu": "NVIDIA H200",
23
+ "torch": "2.12.0+cu126"
24
+ },
25
+ "checkpoint_parameter_equality": true,
26
+ "checkpoint_reload_verified": true
27
+ }
predict.py ADDED
@@ -0,0 +1,91 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Typed-choice inference for ModernJEV-Decide-Preview.
2
+ The encoder scores each declared (state, candidate) pair; it does not generate text.
3
+ """
4
+ import json
5
+ import os
6
+ import torch
7
+ from transformers import AutoModelForSequenceClassification, AutoTokenizer
8
+
9
+ MAX_LEN = 4096
10
+ MODEL_ID = "OpenMed/ModernJEV-Decide-Preview"
11
+
12
+ def serialize_state(state):
13
+ if not isinstance(state, dict):
14
+ raise TypeError("state must be a dict containing conversation and available_tools")
15
+ conv = state.get("conversation") or []
16
+ policy = state.get("policy")
17
+ first = conv[0] if conv else None
18
+ duplicate = (policy is not None and isinstance(first, dict)
19
+ and first.get("role") == "system" and first.get("content") == policy)
20
+ compact = {"available_tools": state.get("available_tools") or [], "conversation": conv}
21
+ if policy is not None and not duplicate:
22
+ compact["policy"] = policy
23
+ return json.dumps(compact, ensure_ascii=False)
24
+
25
+ def normalize_criteria(criteria):
26
+ if isinstance(criteria, list):
27
+ if not criteria or any(not isinstance(v, str) or not v.strip() for v in criteria):
28
+ raise ValueError("Answer list must contain nonempty strings")
29
+ if len(set(criteria)) != len(criteria):
30
+ raise ValueError("Answer list must contain unique strings")
31
+ criteria = {v: v for v in criteria}
32
+ if not isinstance(criteria, dict) or not criteria:
33
+ raise ValueError("criteria must be a nonempty mapping or list of unique answer strings")
34
+ if any(not isinstance(k, str) or not k.strip() or not isinstance(v, str)
35
+ for k, v in criteria.items()):
36
+ raise TypeError("Choice labels must be nonempty strings and descriptions must be strings")
37
+ return criteria
38
+
39
+ class DecisionModel:
40
+ def __init__(self, model_path=MODEL_ID, device=None, revision=None):
41
+ self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
42
+ self.tokenizer = AutoTokenizer.from_pretrained(model_path, revision=revision)
43
+ self.model = AutoModelForSequenceClassification.from_pretrained(
44
+ model_path, revision=revision, attn_implementation="sdpa").to(self.device).eval()
45
+ if self.model.config.num_labels != 1:
46
+ raise ValueError("Expected a trained scalar candidate-scoring head")
47
+
48
+ @torch.inference_mode()
49
+ def decide(self, *, state, question, criteria, candidate_batch_size=8):
50
+ if not isinstance(question, str) or not question.strip():
51
+ raise ValueError("question must be nonempty text")
52
+ criteria = normalize_criteria(criteria)
53
+ if not isinstance(candidate_batch_size, int) or isinstance(candidate_batch_size, bool) or candidate_batch_size < 1:
54
+ raise ValueError("candidate_batch_size must be a positive integer")
55
+ keys = list(criteria)
56
+ text_a = question + "\n\nSTATE:\n" + serialize_state(state)
57
+ text_bs = [f"{k}: {criteria[k]}" for k in keys]
58
+ raw_a_length = len(self.tokenizer(text_a, add_special_tokens=False, verbose=False)["input_ids"])
59
+ special = self.tokenizer.num_special_tokens_to_add(pair=True)
60
+ raw_lengths = [raw_a_length + len(self.tokenizer(t, add_special_tokens=False)["input_ids"]) + special for t in text_bs]
61
+ scores = []
62
+ for begin in range(0, len(keys), candidate_batch_size):
63
+ ts = text_bs[begin:begin + candidate_batch_size]
64
+ encoded = self.tokenizer([text_a] * len(ts), ts, truncation="only_first",
65
+ max_length=MAX_LEN, padding=True, return_tensors="pt", verbose=False).to(self.device)
66
+ with torch.autocast(device_type=self.device.split(":")[0],
67
+ dtype=torch.bfloat16, enabled=self.device.startswith("cuda")):
68
+ logits = self.model(input_ids=encoded["input_ids"],
69
+ attention_mask=encoded["attention_mask"]).logits.squeeze(-1)
70
+ scores.extend(logits.float().cpu().tolist())
71
+ probabilities = torch.softmax(torch.tensor(scores), dim=0).tolist()
72
+ order = sorted(range(len(keys)), key=lambda i: (-scores[i], keys[i]))
73
+ return {"predicted_label": keys[order[0]], "allowed_choices": keys,
74
+ "candidates": [{"label": keys[i], "score": probabilities[i], "raw_score": scores[i],
75
+ "rank": rank + 1} for rank,i in enumerate(order)],
76
+ "truncated": any(n > MAX_LEN for n in raw_lengths),
77
+ "max_sequence_length": MAX_LEN,
78
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."}
79
+
80
+ _default = None
81
+ def predict_typed(question_text, state_json, criteria_json):
82
+ global _default
83
+ if _default is None:
84
+ _default = DecisionModel(os.environ.get("MODEL_PATH", MODEL_ID))
85
+ state = json.loads(state_json) if isinstance(state_json, str) else state_json
86
+ criteria = json.loads(criteria_json) if isinstance(criteria_json, str) else criteria_json
87
+ return _default.decide(state=state, question=question_text, criteria=criteria)
88
+
89
+ def predict_batch(rows):
90
+ return [predict_typed(r["question_text"], r["state_json"], r["criteria_json"]) for r in rows]
91
+
recipe/finalize.py ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse,json,os,time,importlib.util,importlib.metadata
2
+ from pathlib import Path
3
+ from huggingface_hub import HfApi
4
+ p=argparse.ArgumentParser();p.add_argument('--run_dir',required=True);p.add_argument('--repo',default='OpenMed/ModernJEV-Decide-Preview');args=p.parse_args()
5
+ folder=Path(args.run_dir);results=json.loads((folder/'results.json').read_text())
6
+ assert results['test']['n']==1700
7
+ spec=importlib.util.spec_from_file_location('predict_helper',folder/'predict.py');helper=importlib.util.module_from_spec(spec);spec.loader.exec_module(helper)
8
+ client=helper.DecisionModel(str(folder/'model'))
9
+ state={'policy':'Use the order lookup tool when a customer asks about an order.','conversation':[{'role':'user','content':'Where is order A123?'}],'available_tools':[{'name':'lookup_order','description':"Retrieve an order's delivery status."},{'name':'search_catalog','description':'Find products in the catalog.'}]}
10
+ criteria={'lookup_order':"Retrieve the delivery status of the customer's order.",'search_catalog':'Search for products that match a shopping request.'}
11
+ example=client.decide(state=state,question='Which tool should the assistant call next?',criteria=criteria)
12
+ assert example['predicted_label'] in criteria and all(c['label'] in criteria for c in example['candidates'])
13
+ (folder/'example-output.json').write_text(json.dumps(example,indent=2)+'\n')
14
+ versions={name:importlib.metadata.version(name) for name in ['torch','transformers','datasets','accelerate','huggingface-hub','kernels']}
15
+ (folder/'runtime-versions.json').write_text(json.dumps(versions,indent=2)+'\n')
16
+ card="---\nlicense: apache-2.0\nlibrary_name: transformers\npipeline_tag: text-classification\nbase_model: answerdotai/ModernBERT-base\nbase_model_relation: finetune\ndatasets:\n - MaziyarPanahi/AgentToolDecisions-180K\ntags:\n - modernbert\n - encoder\n - decision-model\n - tool-routing\n - agentic\n - preview\nlanguage:\n - en\n---\n\n<div align=\"center\">\n\n# ModernJEV-Decide-Preview\n\n### Give your agent a next move.\n\n**A small encoder for choosing an action or a tool from the options you provide.**\n\n149.6M parameters · 4,096-token inputs · Typed choices · Built on ModernBERT\n\n</div>\n\n> **Training preview.** Training and evaluation are pending. This draft contains no measured accuracy or speed claims. The final card will report actual training coverage and held-out results.\n\nAn agent doesn't need to write a paragraph every time it makes a decision. Sometimes the useful output is simply **answer**, **call a tool**, or **choose this tool**.\n\nModernJEV-Decide-Preview reads a conversation, its policy and available tools, then scores the choices supplied with your question. Your application receives one of those labels, along with the scores for the alternatives. The model ranks choices; your application executes the selected action.\n\nThe architecture is a **ModernBERT cross-encoder with one scalar scoring head**. Choice labels and descriptions are input text, so the output is not limited to a fixed vocabulary of tool names. This is a Jev-style choice model built from public agent decisions, with its own training and evaluation; it is not a reproduction of Jev.\n\n## Where it fits\n\n| Use case | Supply | Receive |\n|---|---|---|\n| **Next-action routing** | Current conversation, policy, tool list | A declared action such as `text_response` or `tool_call` |\n| **Tool selection** | The task and each available tool's purpose | The label of the selected tool |\n| **Workflow branching** | A text state and clearly described alternatives | One branch label; evaluate on your own workflow before adoption |\n| **Decision-model experiments** | Your choices and held-out tasks | Candidate rankings you can inspect and compare |\n\nUseful places to start: support agents choosing between lookup and escalation, assistants selecting an API, and workflows choosing between a direct answer and a tool-backed step. The model does not generate tool arguments or execute tools.\n\n## Quick start\n\nInstall PyTorch and Transformers, then use the included inference helper so the input formatting matches training:\n\n```bash\npip install torch transformers huggingface-hub\n```\n\nDownload the supplied helper, then run it alongside your application. During the private preview, your Hugging Face account must have access to the repository.\n\n```bash\nhf download OpenMed/ModernJEV-Decide-Preview predict.py --local-dir modernjev\ncd modernjev\n```\n\nThe helper's input contract has been reviewed. The examples' predictions will be recorded after the trained checkpoint passes evaluation.\n\n```python\nfrom predict import DecisionModel\n\nmodel = DecisionModel(\"OpenMed/ModernJEV-Decide-Preview\")\ndecision = model.decide(\n state={\n \"policy\": \"Use the order lookup tool when a customer asks about an order.\",\n \"conversation\": [\n {\"role\": \"user\", \"content\": \"Where is order A123?\"}\n ],\n \"available_tools\": [\n {\"name\": \"lookup_order\", \"description\": \"Retrieve an order's delivery status.\"},\n {\"name\": \"search_catalog\", \"description\": \"Find products in the catalog.\"},\n ],\n },\n question=\"Which tool should the assistant call next?\",\n criteria={\n \"lookup_order\": \"Retrieve the delivery status of the customer's order.\",\n \"search_catalog\": \"Search for products that match a shopping request.\",\n },\n)\nprint(decision[\"predicted_label\"])\nprint(decision[\"candidates\"])\n```\n\nThe result contains the selected label, all candidate scores, and whether input truncation occurred. Scores are normalized **within the supplied choice set**; they are not calibrated confidence estimates and cannot be compared directly across unrelated requests.\n\n## Training recipe\n\n| Setting | Planned configuration |\n|---|---|\n| Base | `answerdotai/ModernBERT-base` |\n| Base revision | `8949b909ec900327062f0ebf497f51aef5e6f0c8` |\n| Parameters with scalar head | 149,605,633 |\n| Input limit | 4,096 tokens including both sequences and special tokens |\n| Dataset | `MaziyarPanahi/AgentToolDecisions-180K` |\n| Dataset revision | `f2fb14e4ec977c420f376c08785664cd38763d7e` |\n| Selected training target | Exactly 60,000 train-split decisions, stratified by family, seed 42 |\n| Task families | `agent_next_action_type`, `tool_selection` |\n| Objective | Per-decision softmax cross-entropy over candidate scores |\n| Candidate grouping | `row_id`, preserving `group_id` only for episode/split boundaries |\n| Candidate pool | Pilot will choose full or sampled declared alternatives; final recipe records which ran |\n| Evaluation | Monitor official validation; evaluate the final fixed-epoch checkpoint on 1,700 in-scope official test decisions |\n\nThe dataset has 180,000 total rows. This prototype targets 60,000 of its 112,973 in-scope training decisions; **180,000 is not the number used for training**. Exact selected row IDs, coverage, steps, dependency versions and measured cost will accompany the checkpoint.\n\n## Evaluation\n\n**Pending.** The final comparison will include:\n\n| System | Next action | Tool selection | Macro accuracy |\n|---|---|---|---|\n| ModernJEV-Decide-Preview | Pending | Pending | Pending |\n| ModernBERT + untrained scalar head, seed 42 | Pending | Pending | Pending |\n| Training-derived allowed-choice frequency baseline | Pending | Pending | Pending |\n| Uniform over supplied choices, expected accuracy | Pending | Pending | Pending |\n\nUntrained ModernBERT is a pretrained encoder, not a pretrained decision classifier. Its randomly initialized scalar head is a starting-point comparison, not a claim about ModernBERT's quality on other tasks. A frozen-backbone trained-head comparison will be reported separately if completed within the run budget.\n\nThe report will also include candidate-order checks, input truncation, per-family sample counts, measured inference latency, and actual training coverage. The published split audit found no cross-split conflicts for episode IDs or full upstream source identities.\n\n## Input design matters\n\n- Give every choice a concrete, distinct description. Include the same available tool information the application actually has.\n- Use the supplied serialization helper. Long state may need truncation at the 4,096-token limit; the helper reports it.\n- The training data is English agent conversations and tool decisions. New domains and arbitrary workflow labels need their own evaluation.\n- `tool_and_response` is a declared next-action choice but never a gold label in this dataset. That action is outside demonstrated positive training coverage.\n- This preview implements **choice ranking**. It does not implement Jev's other primitives, vision inputs, generated explanations, or autonomous tool execution.\n\n## Make your own\n\nThe package includes a reusable **HuggingChat ML Intern prompt**, the pinned training recipe and prediction examples. Change the namespace to your own account or organization, review the compute proposal, grant your budget, and keep your own validation/test data separate from training.\n\n## Attribution\n\nBase model: [Answer.AI ModernBERT](https://huggingface.co/answerdotai/ModernBERT-base), Apache 2.0.\n\nDecision dataset: [MaziyarPanahi/AgentToolDecisions-180K](https://huggingface.co/datasets/MaziyarPanahi/AgentToolDecisions-180K), transformed from the upstream agentic sources documented in its card. Keep the pinned revisions and source attribution with derivative work, and follow the upstream dataset licenses when redistributing data.\n\n**Workflow:** HuggingChat ML Intern prepared the proposal, training code, helper and model-card draft. Codex audited the data and corrected candidate targets, padding budgets, coverage accounting and checkpoint verification. Hugging Face Jobs was launched with the operator's approved local HF credential after the connector's write permissions returned 403. This is an assisted training workflow; HuggingChat did not autonomously execute the training job.\n"
17
+ status=f"> **Experimental preview.** The final checkpoint processed **{results['rows_covered']:,}/{results['n_rows_selected']:,} selected training decisions**. Held-out results below cover all **1,700** in-scope official test decisions. This is a small choice-ranking prototype, not a Jev reproduction."
18
+ card=card.replace('> **Training preview.** Training and evaluation are pending. This draft contains no measured accuracy or speed claims. The final card will report actual training coverage and held-out results.',status)
19
+ if results['rows_covered'] != results['n_rows_selected']: card=card.replace(status,status+'\n\n> The time limit stopped training before a complete epoch. Selected rows are not the same as trained rows.')
20
+ card=card.replace('The helper\'s input contract has been reviewed. The examples\' predictions will be recorded after the trained checkpoint passes evaluation.','The included helper was executed against this checkpoint. See [the observed example output](example-output.json); the snippet below prints the actual prediction rather than promising a fixed answer.')
21
+ card=card.replace('| Setting | Planned configuration |','| Setting | Recorded configuration |')
22
+ card=card.replace('| Selected training target | Exactly 60,000 train-split decisions, stratified by family, seed 42 |',f"| Selected training target | Exactly 60,000 train-split decisions, stratified by family, seed 42 |\n| Actual training coverage | {results['rows_covered']:,} decisions; {results['steps']:,} optimizer steps |\n| Attention backend | {results['attention']} |\n| Training elapsed | {results['train_seconds']/60:.1f} minutes |\n| GPU | {results['gpu']} |")
23
+ card=card.replace('| Candidate pool | Pilot will choose full or sampled declared alternatives; final recipe records which ran |','| Candidate pool | Gold plus up to three declared negatives per training decision; evaluation ranks every declared choice |')
24
+ systems=[('ModernJEV-Decide-Preview',results['test']),('ModernBERT + untrained scalar head, seed 42',results['baseline_untrained_head']['test']),('Allowed-choice training-frequency baseline',results['baselines']['test']['train_frequency']),('Uniform over allowed choices, expected accuracy',results['baselines']['test']['uniform_expected'])]
25
+ evaluation='## Evaluation\n\nOfficial held-out test decisions. Next action: **1,158** cases. Tool selection: **542** cases. Macro accuracy gives both families equal weight.\n\n| System | Next action | Tool selection | Macro accuracy | Overall accuracy |\n|---|---|---|---|---|\n'
26
+ for name,r in systems:
27
+ evaluation+=f"| {name} | {r['per_family']['agent_next_action_type']['acc']*100:.2f}% | {r['per_family']['tool_selection']['acc']*100:.2f}% | {r['macro_accuracy']*100:.2f}% | {r['acc']*100:.2f}% |\n"
28
+ inv=results['shuffled_invariance'];lat=results['latency_gpu']
29
+ evaluation+=f"\nCandidate-order label agreement: **{inv['invariance_rate']*100:.2f}%** across {inv['perms']} permutations of {inv['n_rows']} held-out decisions. Ties use lexical label order.\n\nGPU latency for **one candidate forward pass**, excluding tokenization: median {lat['pair_ms_p50']:.1f} ms, p95 {lat['pair_ms_p95']:.1f} ms. A full decision may contain several candidates; these are not end-to-end decision timings.\n\n"
30
+ evaluation+='Untrained ModernBERT is a pretrained encoder with a randomly initialized scoring head, not a pretrained decision classifier. No frozen-backbone trained-head comparison or Jev API comparison was completed in this bounded run.\n\n'
31
+ evaluation+='See [results.json](results.json), [training coverage](training_coverage.json), [metrics](metrics.jsonl), [runtime versions](runtime-versions.json) and [the pinned recipe](recipe/train.py). Input truncation is measured in the tokenization events in the metrics file. The published split audit found no cross-split episode or full upstream source conflicts.\n\n'
32
+ start=card.index('## Evaluation');end=card.index('## Input design matters',start);card=card[:start]+evaluation+card[end:]
33
+ card=card.replace('Exact selected row IDs, coverage, steps, dependency versions and measured cost will accompany the checkpoint.','Exact selected row IDs, coverage, steps and dependency versions accompany this checkpoint. This repository contains model artifacts; it does not expose an inference endpoint.')
34
+ card=card.replace('## Where it fits','![How decisions are scored](architecture.svg)\n\n## Where it fits')
35
+ (folder/'README.md').write_text(card)
36
+ svg='<svg xmlns="http://www.w3.org/2000/svg" width="1080" height="220" viewBox="0 0 1080 220"><rect width="1080" height="220" rx="22" fill="#101820"/><g font-family="sans-serif" text-anchor="middle"><text x="540" y="40" font-size="18" fill="#a8d9bd">READ THE STATE. SCORE THE CHOICES. RETURN A LABEL.</text><g fill="#20332d" stroke="#4e9f79"><rect x="30" y="80" width="245" height="100" rx="12"/><rect x="350" y="80" width="350" height="100" rx="12"/><rect x="775" y="80" width="275" height="100" rx="12"/></g><g font-size="20" fill="#f1f4f2"><text x="152" y="117">STATE + QUESTION</text><text x="152" y="150">DECLARED CHOICES</text><text x="525" y="117">SHARED MODERNBERT</text><text x="525" y="150">CANDIDATE SCORER</text><text x="912" y="117">SELECTED LABEL</text><text x="912" y="150">+ CHOICE RANKINGS</text></g><g fill="#8bd9af" font-size="32"><text x="311" y="140">→</text><text x="738" y="140">→</text></g></g></svg>'
37
+ (folder/'architecture.svg').write_text(svg)
38
+ api=HfApi();assert api.model_info(args.repo).private
39
+ for filename in ['README.md','architecture.svg','runtime-versions.json','example-output.json']:
40
+ api.upload_file(path_or_fileobj=str(folder/filename),path_in_repo=filename,repo_id=args.repo,repo_type='model',commit_message='Measured prototype results and verified quick start')
41
+ print('MODEL_CARD_FINALIZED '+json.dumps({'rows_covered':results['rows_covered'],'test_accuracy':results['test']['acc'],'private':True}),flush=True)
recipe/selected-row-ids.txt ADDED
The diff for this file is too large to render. See raw diff
 
recipe/subset-60000.json ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "target_training_decisions": 60000,
3
+ "seed": 42,
4
+ "dataset_revision": "f2fb14e4ec977c420f376c08785664cd38763d7e",
5
+ "counts": {
6
+ "agent_next_action_type": 40809,
7
+ "tool_selection": 19191
8
+ },
9
+ "selected_row_ids_sha256": "c5b306b0471ba104161051a7241b3bce4b69e1d959ff3e34fb06ef8eb4b077d9",
10
+ "full_candidate_pairs": 504012,
11
+ "sampled4_candidate_pairs": 199191,
12
+ "next_action_gold_labels": {
13
+ "text_response": 21636,
14
+ "tool_call": 19173
15
+ },
16
+ "training_started": false
17
+ }
recipe/train.py ADDED
@@ -0,0 +1,859 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ModernJEV-Decide-Preview — training, evaluation, persistence.
2
+
3
+ Dataset : MaziyarPanahi/AgentToolDecisions-180K @ f2fb14e4ec977c420f376c08785664cd38763d7e
4
+ Base : answerdotai/ModernBERT-base @ 8949b909ec900327062f0ebf497f51aef5e6f0c8
5
+ Scope : task_family in {agent_next_action_type, tool_selection} ONLY (choice primitive).
6
+
7
+ Objective: shared-encoder candidate scalar scorer; per-row_id softmax cross-entropy
8
+ over the row's DECLARED candidates. group_id is the EPISODE, not the decision —
9
+ grouping is ALWAYS by row_id (asserted). Candidates enter the softmax as text
10
+ (label + criterion), so variable tool names need no fixed head and no candidate
11
+ index is exposed to the model. Inputs contain NO gold_label / gold_json /
12
+ gold_score / label_source / source metadata.
13
+
14
+ Training pool (--pool):
15
+ full — every declared candidate of the row joins the softmax.
16
+ sampled4 — gold + up to 3 declared negatives, deterministic per (SEED, epoch,
17
+ row_id). This is ordinary sampled-choice softmax within the same loss;
18
+ EVALUATION ALWAYS RANKS ALL DECLARED CANDIDATES regardless of --pool.
19
+ The pilot benchmarks both objectives and the launch report states which one ran.
20
+
21
+ Modes:
22
+ pilot — benchmarks full vs sampled4 throughput/memory, GPU latency, save check.
23
+ prototype — time-guarded training on the exact stratified subset, interval monitoring
24
+ on a stratified subset of the OFFICIAL validation split (full-val eval
25
+ recorded for the final fixed-epoch checkpoint), ONE final test evaluation on all in-scope
26
+ test rows, baselines (uniform / train-frequency / untrained ModernBERT
27
+ head), shuffled-order label invariance, optional frozen-backbone probe,
28
+ and persistence to --save_dir (default /output/modernjev). No Hub push
29
+ unless --push is passed explicitly.
30
+
31
+ No Space is created anywhere; metrics are metrics.jsonl + stdout only.
32
+ """
33
+ import argparse
34
+ import gzip
35
+ import hashlib
36
+ import json
37
+ import os
38
+ import random
39
+ import time
40
+
41
+ import numpy as np
42
+ import torch
43
+ import torch.nn.functional as F
44
+ from datasets import Dataset, load_dataset
45
+ from torch.utils.data import DataLoader, Dataset as TorchDataset, Sampler
46
+ from transformers import (AutoModelForSequenceClassification, AutoTokenizer,
47
+ Trainer, TrainerCallback, TrainingArguments)
48
+
49
+ DS_ID = "MaziyarPanahi/AgentToolDecisions-180K"
50
+ DS_REV = "f2fb14e4ec977c420f376c08785664cd38763d7e"
51
+ BASE_ID = "answerdotai/ModernBERT-base"
52
+ BASE_REV = "8949b909ec900327062f0ebf497f51aef5e6f0c8"
53
+ FOCUS = ("agent_next_action_type", "tool_selection")
54
+ SEED = 42
55
+ MAX_LEN = 4096
56
+ ATTN_IMPL = os.environ.get("MODERNJEV_ATTN", "sdpa")
57
+ EXPECTED = {"train": 171056, "validation": 2713, "test": 6231}
58
+ EXPECTED_FOCUS_TRAIN = 112973
59
+
60
+ METRICS = []
61
+
62
+
63
+ def log_metric(d):
64
+ d["ts"] = round(time.time(), 1)
65
+ METRICS.append(d)
66
+ print("METRIC " + json.dumps(d, default=str), flush=True)
67
+
68
+
69
+ def save_metrics(path):
70
+ with open(path, "w") as f:
71
+ for d in METRICS:
72
+ f.write(json.dumps(d, default=str) + "\n")
73
+
74
+
75
+ def serialize_state(row):
76
+ """Compact the state. Drops the policy key ONLY when it exactly equals the
77
+ first system message (audited dedupe rule)."""
78
+ state = json.loads(row["state_json"])
79
+ conv = state.get("conversation") or []
80
+ policy = state.get("policy")
81
+ first = conv[0] if conv else None
82
+ dup = (policy is not None and isinstance(first, dict)
83
+ and first.get("role") == "system" and first.get("content") == policy)
84
+ compact = {"available_tools": state.get("available_tools") or [],
85
+ "conversation": conv}
86
+ if policy is not None and not dup:
87
+ compact["policy"] = policy
88
+ return json.dumps(compact, ensure_ascii=False)
89
+
90
+
91
+ def parse_row(r):
92
+ criteria = json.loads(r["criteria_json"])
93
+ keys = list(criteria.keys())
94
+ gold = r["gold_label"]
95
+ return {"row_id": r["row_id"], "group_id": r["group_id"],
96
+ "family": r["task_family"],
97
+ "text_a": r["question_text"] + "\n\nSTATE:\n" + serialize_state(r),
98
+ "cand_keys": keys,
99
+ "cand_texts": [f"{k}: {criteria[k]}" for k in keys],
100
+ "gold_idx": keys.index(gold) if gold in keys else None}
101
+
102
+
103
+ def load_and_prepare(n_rows, seed=SEED):
104
+ t0 = time.time()
105
+ ds = load_dataset(DS_ID, revision=DS_REV)
106
+ for split, n in EXPECTED.items():
107
+ assert len(ds[split]) == n, f"{split}: {len(ds[split])} != {n}"
108
+ train = ds["train"].filter(lambda r: r["task_family"] in FOCUS, num_proc=8)
109
+ assert len(train) == EXPECTED_FOCUS_TRAIN, f"{len(train)} != {EXPECTED_FOCUS_TRAIN}"
110
+
111
+ metadata = train.select_columns(["row_id", "task_family"])[:]
112
+ fam_idx = {f: [] for f in FOCUS}
113
+ for i, fam in enumerate(metadata["task_family"]):
114
+ fam_idx[fam].append(i)
115
+ n_a = n_rows * len(fam_idx["agent_next_action_type"]) // len(train)
116
+ n_t = n_rows - n_a
117
+ rng = random.Random(seed)
118
+ picked = []
119
+ for fam, k in (("agent_next_action_type", n_a), ("tool_selection", n_t)):
120
+ all_ids = metadata["row_id"]
121
+ idxs = sorted(fam_idx[fam], key=lambda i: all_ids[i])
122
+ picked += rng.sample(idxs, k)
123
+ assert len(picked) == n_rows
124
+ picked.sort()
125
+ fields = ["row_id", "group_id", "task_family", "question_text", "state_json", "criteria_json", "gold_label"]
126
+ selected = train.select(picked).select_columns(fields)[:]
127
+ parsed = [parse_row(dict(zip(fields, values))) for values in zip(*(selected[f] for f in fields))]
128
+ for r in parsed:
129
+ assert r["gold_idx"] is not None, f"gold missing in {r['row_id']}"
130
+
131
+ fam_counts = {f: sum(1 for r in parsed if r["family"] == f) for f in FOCUS}
132
+ gold_classes = {f: {} for f in FOCUS}
133
+ for r in parsed:
134
+ g = r["cand_keys"][r["gold_idx"]]
135
+ gold_classes[r["family"]][g] = gold_classes[r["family"]].get(g, 0) + 1
136
+ ids_hash = hashlib.sha256(
137
+ "\n".join(sorted(r["row_id"] for r in parsed)).encode()).hexdigest()
138
+ log_metric({"event": "data_prep", "n_rows": len(parsed), "n_a": n_a, "n_t": n_t,
139
+ "family_counts": fam_counts, "gold_class_counts": {f: gold_classes[f] if f == "agent_next_action_type" else {"distinct": len(gold_classes[f])} for f in FOCUS},
140
+ "selected_row_ids_sha256": ids_hash,
141
+ "seconds": round(time.time() - t0, 1)})
142
+ return parsed, ds, ids_hash, (n_a, n_t)
143
+
144
+
145
+ class LazyPairDataset:
146
+ """Tokenize only requested pairs. Metadata stays small; no up-front Dataset.map."""
147
+ def __init__(self, parsed, tokenizer, tag, max_len, training_pool):
148
+ self.parsed, self.tokenizer, self.tag, self.max_len = parsed, tokenizer, tag, max_len
149
+ self.refs = []
150
+ for row in range(len(parsed)):
151
+ candidates, _ = pool_for(parsed, row, 1, training_pool)
152
+ self.refs.extend((row, candidate) for candidate in candidates)
153
+ self.metadata = Dataset.from_dict({
154
+ "row_idx": [r for r, _ in self.refs],
155
+ "cand_idx": [c for _, c in self.refs],
156
+ # Conservative packing bound; actual batch padding uses actual token lengths.
157
+ "length": [max_len] * len(self.refs),
158
+ })
159
+ self.row_sortlen = np.array([
160
+ min(max_len, max(1, len(r["text_a"]) // 4)) for r in parsed], dtype=np.int64)
161
+ log_metric({"event": "lazy_pairs_ready", "tag": tag,
162
+ "n_rows": len(parsed), "n_pairs": len(self.refs),
163
+ "max_len": max_len, "mapped_pairs": 0})
164
+
165
+ def __len__(self):
166
+ return len(self.refs)
167
+
168
+ def select_columns(self, columns):
169
+ return self.metadata.select_columns(columns)
170
+
171
+ def __getitem__(self, indices):
172
+ scalar = isinstance(indices, (int, np.integer))
173
+ if scalar:
174
+ indices = [int(indices)]
175
+ elif isinstance(indices, slice):
176
+ indices = list(range(*indices.indices(len(self))))
177
+ else:
178
+ indices = list(indices)
179
+ refs = [self.refs[int(i)] for i in indices]
180
+ encoded = self.tokenizer(
181
+ [self.parsed[r]["text_a"] for r, c in refs],
182
+ [self.parsed[r]["cand_texts"][c] for r, c in refs],
183
+ truncation="only_first", max_length=self.max_len,
184
+ padding=False, verbose=False)
185
+ result = {"row_idx": [r for r, c in refs],
186
+ "cand_idx": [c for r, c in refs],
187
+ "input_ids": encoded["input_ids"],
188
+ "attention_mask": encoded["attention_mask"],
189
+ "length": [len(ids) for ids in encoded["input_ids"]],
190
+ "truncated": [bool(e.overflowing) for e in encoded.encodings]}
191
+ return {k: v[0] for k, v in result.items()} if scalar else result
192
+
193
+
194
+ def tokenize_pairs(parsed, tokenizer, tag, max_len=MAX_LEN, num_proc=8, training_pool="full"):
195
+ return LazyPairDataset(parsed, tokenizer, tag, max_len, training_pool)
196
+
197
+
198
+ class RowIndexer:
199
+ def __init__(self, parsed, pair_ds):
200
+ meta = pair_ds.select_columns(["row_idx", "cand_idx", "length"]).with_format("numpy")[:]
201
+ assert np.all(np.diff(meta["row_idx"]) >= 0), "pairs must be row-major"
202
+ self.lookup = {}
203
+ self.row_maxlen = np.zeros(len(parsed), dtype=np.int64)
204
+ for flat, (row, cand, length) in enumerate(zip(meta["row_idx"], meta["cand_idx"], meta["length"])):
205
+ row, cand = int(row), int(cand)
206
+ assert (row, cand) not in self.lookup
207
+ self.lookup[(row, cand)] = flat
208
+ self.row_maxlen[row] = max(self.row_maxlen[row], int(length))
209
+ assert np.all(self.row_maxlen > 0)
210
+ self.row_sortlen = getattr(pair_ds, "row_sortlen", self.row_maxlen)
211
+ def flat(self, row, cand):
212
+ return self.lookup[(row, cand)]
213
+
214
+
215
+ def pool_for(parsed, row, epoch, mode):
216
+ """Training pool for one row: (pool_cand_indices, gold_pos_in_pool).
217
+ full: all declared candidates. sampled4: gold + up to 3 declared negatives,
218
+ deterministic per (SEED, epoch, row_id)."""
219
+ r = parsed[row]
220
+ k = len(r["cand_keys"])
221
+ gold = r["gold_idx"]
222
+ if mode == "full" or k <= 4:
223
+ return list(range(k)), gold
224
+ rng = random.Random(f"{SEED}:{epoch}:{r['row_id']}")
225
+ others = sorted(set(range(k)) - {gold})
226
+ negs = rng.sample(others, min(3, len(others)))
227
+ pool = [gold] + sorted(negs)
228
+ rng.shuffle(pool)
229
+ return pool, pool.index(gold)
230
+
231
+
232
+ class TrainRefs(TorchDataset):
233
+ """Position i -> (row, cand_actual) for the current epoch; gold_pos is
234
+ recomputed in collate from the pool definition (deterministic)."""
235
+
236
+ def __init__(self, sampler):
237
+ self.sampler = sampler
238
+
239
+ def __len__(self):
240
+ return len(self.sampler.refs)
241
+
242
+ def __getitem__(self, i):
243
+ return self.sampler.refs[i]
244
+
245
+
246
+ class WholeRowBatchSampler(Sampler):
247
+ """Yields batches of positions into self.refs. Every row's pooled pairs stay
248
+ in one batch (padding-token budget checked at row boundaries)."""
249
+
250
+ def __init__(self, parsed, indexer, mode, token_budget, max_rows=64,
251
+ seed=SEED, shuffle=True, length_buckets=True):
252
+ self.parsed, self.indexer, self.mode = parsed, indexer, mode
253
+ self.token_budget, self.max_rows = token_budget, max_rows
254
+ self.rng = random.Random(seed)
255
+ self.shuffle = shuffle
256
+ self.length_buckets = length_buckets
257
+ self.epoch = 0
258
+ self.refs = []
259
+ self.rows_cum = 0
260
+ self.pairs_cum = 0
261
+ self.tokens_cum = 0
262
+ self.batches = self._build_epoch()
263
+
264
+ def _build_epoch(self):
265
+ self.epoch += 1
266
+ order = list(range(len(self.parsed)))
267
+ if self.shuffle:
268
+ self.rng.shuffle(order)
269
+ # Sort locally within randomized buckets to reduce padding, preserving all rows.
270
+ if self.length_buckets:
271
+ order = [row for begin in range(0, len(order), 256)
272
+ for row in sorted(order[begin:begin + 256], key=lambda r: int(self.indexer.row_sortlen[r]))]
273
+ refs, batches, cur = [], [], []
274
+ batch_maxlen, batch_rows = 0, 0
275
+ for r in order:
276
+ pool, _ = pool_for(self.parsed, r, self.epoch, self.mode)
277
+ row_len = int(self.indexer.row_maxlen[r])
278
+ new_maxlen = max(batch_maxlen, row_len)
279
+ if cur and (new_maxlen * (len(cur) + len(pool)) > self.token_budget
280
+ or batch_rows >= self.max_rows):
281
+ batches.append(cur)
282
+ cur, batch_maxlen, batch_rows = [], 0, 0
283
+ start = len(refs)
284
+ refs.extend((r, cand) for cand in pool)
285
+ cur.extend(range(start, start + len(pool)))
286
+ batch_maxlen = max(batch_maxlen, row_len)
287
+ batch_rows += 1
288
+ if cur:
289
+ batches.append(cur)
290
+ self.refs = refs
291
+ return batches
292
+
293
+ def __iter__(self):
294
+ return iter(self.batches)
295
+
296
+ def __len__(self):
297
+ return len(self.batches)
298
+
299
+
300
+ def make_collate(pair_ds, indexer, pad_id, parsed, mode):
301
+ def collate(items):
302
+ rows = torch.tensor([it[0] for it in items], dtype=torch.long)
303
+ by_row = {}
304
+ for row_index, candidate_index in items:
305
+ by_row.setdefault(row_index, []).append(candidate_index)
306
+ gold_positions = {}
307
+ for row_index, candidate_indices in by_row.items():
308
+ gold_index = parsed[row_index]["gold_idx"]
309
+ assert candidate_indices.count(gold_index) == 1, "Each row needs exactly one gold candidate"
310
+ gold_positions[row_index] = candidate_indices.index(gold_index)
311
+ golds = torch.tensor([gold_positions[it[0]] for it in items], dtype=torch.long)
312
+ flats = [indexer.flat(it[0], it[1]) for it in items]
313
+ rec = pair_ds[flats]
314
+ maxlen = max(rec["length"])
315
+ B = len(items)
316
+ ids = torch.full((B, maxlen), pad_id, dtype=torch.long)
317
+ att = torch.zeros((B, maxlen), dtype=torch.long)
318
+ for i, (ii, aa) in enumerate(zip(rec["input_ids"], rec["attention_mask"])):
319
+ n = len(ii)
320
+ ids[i, :n] = torch.tensor(ii, dtype=torch.long)
321
+ att[i, :n] = torch.tensor(aa, dtype=torch.long)
322
+ return {"input_ids": ids, "attention_mask": att, "row": rows, "gold": golds,
323
+ "_n_truncated": sum(rec["truncated"]), "_n_at_max": sum(n >= MAX_LEN for n in rec["length"])}
324
+ collate.epoch = 1
325
+ return collate
326
+
327
+
328
+ def grouped_ce(logits, rows, gold):
329
+ """Per-row softmax CE. rows must be contiguous per row (asserted)."""
330
+ uniq_c = torch.unique_consecutive(rows)
331
+ uniq_all = torch.unique(rows)
332
+ assert len(uniq_c) == len(uniq_all), "row pairs not contiguous — grouping unsafe"
333
+ counts = torch.unique_consecutive(rows, return_counts=True)[1].tolist()
334
+ losses = []
335
+ ofs = 0
336
+ for c in counts:
337
+ seg = logits[ofs:ofs + c]
338
+ g = int(gold[ofs].item())
339
+ assert g < c, f"gold {g} out of pool size {c}"
340
+ losses.append(F.cross_entropy(seg.unsqueeze(0),
341
+ torch.tensor([g], device=seg.device)))
342
+ ofs += c
343
+ return torch.stack(losses).mean(), len(counts)
344
+
345
+
346
+ class GroupTrainer(Trainer):
347
+ def __init__(self, *args, **kwargs):
348
+ super().__init__(*args, **kwargs)
349
+ self.model_accepts_loss_kwargs = False
350
+
351
+ def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
352
+ gold = inputs.pop("gold")
353
+ rows = inputs.pop("row")
354
+ if not hasattr(self, "seen_rows"):
355
+ self.seen_rows = set()
356
+ self.rows_processed = 0
357
+ self.pairs_processed = 0
358
+ self.tokens_processed = 0
359
+ self.truncated_pairs = 0
360
+ self.pairs_at_max = 0
361
+ self.truncated_pairs += int(inputs.pop("_n_truncated", 0))
362
+ self.pairs_at_max += int(inputs.pop("_n_at_max", 0))
363
+ batch_rows = torch.unique_consecutive(rows).detach().cpu().tolist()
364
+ self.seen_rows.update(batch_rows)
365
+ self.rows_processed += len(batch_rows)
366
+ self.pairs_processed += len(rows)
367
+ self.tokens_processed += int(inputs["attention_mask"].sum().item())
368
+ out = model(input_ids=inputs["input_ids"],
369
+ attention_mask=inputs["attention_mask"])
370
+ logits = out.logits.squeeze(-1).float()
371
+ loss, _ = grouped_ce(logits, rows, gold)
372
+ return (loss, out) if return_outputs else loss
373
+
374
+ def get_train_dataloader(self):
375
+ workers = self.args.dataloader_num_workers
376
+ return DataLoader(self._refs_ds, batch_sampler=self._sampler,
377
+ collate_fn=self._collate, num_workers=workers,
378
+ persistent_workers=workers > 0,
379
+ prefetch_factor=2 if workers else None,
380
+ pin_memory=self.args.device.type == "cuda")
381
+
382
+
383
+ class EvalPairs(TorchDataset):
384
+ def __init__(self, pair_ds):
385
+ self.pair_ds = pair_ds
386
+
387
+ def __len__(self):
388
+ return len(self.pair_ds)
389
+
390
+ def __getitem__(self, i):
391
+ return i
392
+
393
+
394
+ def make_eval_collate(pair_ds, pad_id):
395
+ def collate(idx_list):
396
+ rec = pair_ds[idx_list]
397
+ maxlen = max(rec["length"])
398
+ B = len(idx_list)
399
+ ids = torch.full((B, maxlen), pad_id, dtype=torch.long)
400
+ att = torch.zeros((B, maxlen), dtype=torch.long)
401
+ for i, (ii, aa) in enumerate(zip(rec["input_ids"], rec["attention_mask"])):
402
+ n = len(ii)
403
+ ids[i, :n] = torch.tensor(ii, dtype=torch.long)
404
+ att[i, :n] = torch.tensor(aa, dtype=torch.long)
405
+ return {"input_ids": ids, "attention_mask": att,
406
+ "row": torch.tensor(rec["row_idx"], dtype=torch.long),
407
+ "_n_truncated": sum(rec["truncated"])}
408
+ return collate
409
+
410
+
411
+ @torch.no_grad()
412
+ def evaluate_rows(model, pair_ds, parsed, device, batch_size=8):
413
+ """Row-grouped accuracy over ALL declared candidates."""
414
+ model.eval()
415
+ scores = [[] for _ in range(len(parsed))]
416
+ n_truncated_pairs = 0
417
+ dl = DataLoader(EvalPairs(pair_ds), batch_size=batch_size,
418
+ collate_fn=make_eval_collate(pair_ds, model.config.pad_token_id), num_workers=0)
419
+ t0 = time.time()
420
+ for batch in dl:
421
+ n_truncated_pairs += int(batch["_n_truncated"])
422
+ with torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
423
+ out = model(input_ids=batch["input_ids"].to(device),
424
+ attention_mask=batch["attention_mask"].to(device))
425
+ lg = out.logits.squeeze(-1).float().cpu()
426
+ for j, r in enumerate(batch["row"].tolist()):
427
+ scores[r].append(lg[j].item())
428
+ correct, n = 0, 0
429
+ per_fam = {f: [0, 0] for f in FOCUS}
430
+ pred_labels = []
431
+ for r, row in enumerate(parsed):
432
+ if not scores[r]:
433
+ continue
434
+ p = min(range(len(scores[r])), key=lambda c: (-scores[r][c], row["cand_keys"][c]))
435
+ ok = int(p == row["gold_idx"])
436
+ correct += ok
437
+ n += 1
438
+ per_fam[row["family"]][0] += ok
439
+ per_fam[row["family"]][1] += 1
440
+ pred_labels.append(row["cand_keys"][p])
441
+ model.train()
442
+ return {"acc": correct / max(n, 1), "n": n,
443
+ "macro_accuracy": sum(per_fam[f][0] / max(per_fam[f][1], 1) for f in FOCUS) / len(FOCUS),
444
+ "per_family": {f: {"acc": per_fam[f][0] / max(per_fam[f][1], 1),
445
+ "n": per_fam[f][1]} for f in FOCUS},
446
+ "pred_labels": pred_labels,
447
+ "n_pairs": len(pair_ds), "n_truncated_pairs": n_truncated_pairs,
448
+ "seconds": round(time.time() - t0, 1)}
449
+
450
+
451
+ def stratified_subset(parsed, n, seed=SEED):
452
+ fam_rows = {f: sorted([r for r in parsed if r["family"] == f],
453
+ key=lambda r: r["row_id"]) for f in FOCUS}
454
+ total = sum(len(v) for v in fam_rows.values())
455
+ out = []
456
+ for f in FOCUS:
457
+ k = n * len(fam_rows[f]) // total
458
+ out += random.Random(seed).sample(fam_rows[f], k)
459
+ return out[:n]
460
+
461
+
462
+ def parse_eval_rows(split_ds):
463
+ parsed, skipped = [], 0
464
+ for r in split_ds:
465
+ row = parse_row(r)
466
+ if row["gold_idx"] is None:
467
+ skipped += 1
468
+ continue
469
+ parsed.append(row)
470
+ return parsed, skipped
471
+
472
+
473
+ def tokenize_eval(parsed, tokenizer):
474
+ pd = tokenize_pairs(parsed, tokenizer, tag="eval", num_proc=4)
475
+ idx = RowIndexer(parsed, pd)
476
+ return pd, idx
477
+
478
+
479
+ def build_baseline_results(parsed_tr, parsed_val, parsed_te):
480
+ """Analytic uniform expected accuracy and allowed-choice training frequency."""
481
+ freq = {f: {} for f in FOCUS}
482
+ for row in parsed_tr:
483
+ g = row["cand_keys"][row["gold_idx"]]
484
+ fam_freq = freq[row["family"]]
485
+ fam_freq[g] = fam_freq.get(g, 0) + 1
486
+ res = {}
487
+ for name, parsed in (("val", parsed_val), ("test", parsed_te)):
488
+ methods = {"uniform_expected": {f: [0., 0] for f in FOCUS},
489
+ "train_frequency": {f: [0., 0] for f in FOCUS}}
490
+ for row in parsed:
491
+ fam = row["family"]
492
+ expected = 1. / len(row["cand_keys"])
493
+ best = min(row["cand_keys"], key=lambda key: (-freq[fam].get(key, 0), key))
494
+ matched = int(best == row["cand_keys"][row["gold_idx"]])
495
+ for method, value in (("uniform_expected", expected), ("train_frequency", matched)):
496
+ methods[method][fam][0] += value
497
+ methods[method][fam][1] += 1
498
+ res[name] = {}
499
+ for method, counts in methods.items():
500
+ per = {f: {"acc": c / n, "n": n} for f, (c, n) in counts.items()}
501
+ res[name][method] = {"acc": sum(c for c, _ in counts.values()) / len(parsed),
502
+ "n": len(parsed), "per_family": per,
503
+ "macro_accuracy": sum(v["acc"] for v in per.values()) / len(FOCUS)}
504
+ return res
505
+
506
+
507
+ @torch.no_grad()
508
+ def gpu_latency(model, pair_ds, device, n=100):
509
+ model.eval()
510
+ lat = []
511
+ for i in range(min(n + 5, len(pair_ds))):
512
+ rec = pair_ds[[i]]
513
+ ids = torch.tensor(rec["input_ids"][0], dtype=torch.long,
514
+ device=device).unsqueeze(0)
515
+ att = torch.tensor(rec["attention_mask"][0], dtype=torch.long,
516
+ device=device).unsqueeze(0)
517
+ if device == "cuda":
518
+ torch.cuda.synchronize()
519
+ t = time.time()
520
+ with torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
521
+ model(input_ids=ids, attention_mask=att)
522
+ if device == "cuda":
523
+ torch.cuda.synchronize()
524
+ if i >= 5:
525
+ lat.append(time.time() - t)
526
+ lat.sort()
527
+ return {"measurement": "one candidate forward pass, excludes tokenization", "warmup_pairs": 5,
528
+ "pair_ms_p50": round(1000 * lat[len(lat) // 2], 1),
529
+ "pair_ms_p95": round(1000 * lat[int(0.95 * len(lat))], 1)}
530
+
531
+
532
+ def shuffled_invariance(model, tokenizer, parsed_subset, device, k=5, batch_size=8):
533
+ """Permute candidate order k times; the predicted LABEL must be identical."""
534
+ rng = random.Random(SEED + 1)
535
+ pd0, _ = tokenize_eval(parsed_subset, tokenizer)
536
+ base_labels = evaluate_rows(model, pd0, parsed_subset, device)["pred_labels"]
537
+ agree, total = 0, 0
538
+ for rep in range(k):
539
+ perm = []
540
+ for r in parsed_subset:
541
+ order = list(range(len(r["cand_texts"])))
542
+ rng.shuffle(order)
543
+ perm.append({**r, "cand_keys": [r["cand_keys"][j] for j in order],
544
+ "cand_texts": [r["cand_texts"][j] for j in order],
545
+ "gold_idx": order.index(r["gold_idx"])})
546
+ pd, _ = tokenize_eval(perm, tokenizer)
547
+ acc = evaluate_rows(model, pd, perm, device, batch_size)
548
+ agree += sum(1 for a, b in zip(base_labels, acc["pred_labels"]) if a == b)
549
+ total += len(base_labels)
550
+ return {"invariance_rate": agree / max(total, 1), "perms": k,
551
+ "n_rows": len(parsed_subset)}
552
+
553
+
554
+ def parse_args():
555
+ p = argparse.ArgumentParser()
556
+ p.add_argument("--mode", choices=["pilot", "prototype"], required=True)
557
+ p.add_argument("--pool", choices=["full", "sampled4"], default="sampled4",
558
+ help="training candidate pool (eval always uses all declared)")
559
+ p.add_argument("--n_rows", type=int, default=60000)
560
+ p.add_argument("--max_steps", type=int, default=25)
561
+ p.add_argument("--pilot_sampled_only", action="store_true")
562
+ p.add_argument("--pilot_length_buckets", action="store_true")
563
+ p.add_argument("--token_budget", type=int, default=32768)
564
+ p.add_argument("--grad_accum", type=int, default=2)
565
+ p.add_argument("--lr", type=float, default=2e-5)
566
+ p.add_argument("--eval_steps", type=int, default=400)
567
+ p.add_argument("--deadline_seconds", type=int, default=6600)
568
+ p.add_argument("--eval_reserve_seconds", type=int, default=1200)
569
+ p.add_argument("--save_dir", default="/output/modernjev")
570
+ p.add_argument("--push", action="store_true", default=False)
571
+ p.add_argument("--hub_model_id", default="OpenMed/ModernJEV-Decide-Preview")
572
+ return p.parse_args()
573
+
574
+
575
+ def prototype_training_args(args, save_dir, device):
576
+ return TrainingArguments(
577
+ output_dir=os.path.join(save_dir, "ckpt"), per_device_train_batch_size=1,
578
+ gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr,
579
+ bf16=device == "cuda", use_cpu=device == "cpu",
580
+ num_train_epochs=1.0, logging_steps=25, logging_first_step=True,
581
+ save_strategy="no", eval_strategy="no", report_to="none",
582
+ seed=SEED, remove_unused_columns=False,
583
+ lr_scheduler_type="linear", warmup_steps=0.03,
584
+ dataloader_num_workers=2 if device == "cuda" else 0)
585
+
586
+
587
+ def load_model():
588
+ return AutoModelForSequenceClassification.from_pretrained(
589
+ BASE_ID, revision=BASE_REV, num_labels=1, attn_implementation=ATTN_IMPL)
590
+
591
+
592
+ def main():
593
+ args = parse_args()
594
+ t_start = time.time()
595
+ torch.manual_seed(SEED)
596
+ random.seed(SEED)
597
+ device = "cuda" if torch.cuda.is_available() else "cpu"
598
+ gpu_name = torch.cuda.get_device_name(0) if device == "cuda" else "cpu"
599
+ import transformers
600
+ log_metric({"event": "env", "mode": args.mode, "device": device, "gpu": gpu_name,
601
+ "torch": torch.__version__, "transformers": transformers.__version__, "attention": ATTN_IMPL})
602
+ assert device == "cuda", "GPU required"
603
+ # Validate the complete prototype argument branch before any dataset work.
604
+ validated_training_args = prototype_training_args(args, args.save_dir, device) if args.mode == "prototype" else None
605
+ log_metric({"event": "training_api_validated", "mode": args.mode})
606
+
607
+ tokenizer = AutoTokenizer.from_pretrained(BASE_ID, revision=BASE_REV)
608
+ pad_id = tokenizer.pad_token_id
609
+ assert pad_id is not None, "tokenizer has no pad token"
610
+
611
+ n_train_rows = 500 if args.mode == "pilot" else args.n_rows
612
+ parsed, ds, ids_hash, (n_a, n_t) = load_and_prepare(n_train_rows)
613
+ if args.mode == "prototype":
614
+ assert ids_hash == "c5b306b0471ba104161051a7241b3bce4b69e1d959ff3e34fb06ef8eb4b077d9", "Selected subset differs from authorized frozen manifest"
615
+ os.makedirs(args.save_dir, exist_ok=True)
616
+ pair_ds = tokenize_pairs(parsed, tokenizer, tag="train",
617
+ training_pool=args.pool if args.mode == "prototype" else "full")
618
+ indexer = RowIndexer(parsed, pair_ds)
619
+ collate = make_collate(pair_ds, indexer, pad_id, parsed, args.pool)
620
+
621
+ if args.mode == "pilot":
622
+ results = {"phases": {}}
623
+ pilot_modes = ("sampled4",) if args.pilot_sampled_only else ("full", "sampled4")
624
+ for phase_i, mode in enumerate(pilot_modes):
625
+ torch.manual_seed(SEED)
626
+ phase_steps = args.max_steps if args.pilot_sampled_only else args.max_steps // 2 + (args.max_steps % 2 if phase_i == 1 else 0)
627
+ model = load_model().to(device)
628
+ sampler = WholeRowBatchSampler(parsed, indexer, mode, args.token_budget, length_buckets=args.pilot_length_buckets)
629
+ targs = TrainingArguments(
630
+ output_dir="/tmp/pilot_" + mode, per_device_train_batch_size=1,
631
+ gradient_accumulation_steps=1, learning_rate=args.lr, bf16=True,
632
+ max_steps=phase_steps, logging_steps=3, save_strategy="no",
633
+ eval_strategy="no", report_to="none", seed=SEED,
634
+ remove_unused_columns=False)
635
+ trainer = GroupTrainer(model=model, args=targs,
636
+ train_dataset=TrainRefs(sampler))
637
+ trainer._sampler = sampler
638
+ trainer._collate = make_collate(pair_ds, indexer, pad_id, parsed, mode)
639
+ trainer._refs_ds = TrainRefs(sampler)
640
+ torch.cuda.reset_peak_memory_stats()
641
+ t0 = time.time()
642
+ trainer.train()
643
+ dt = time.time() - t0
644
+ results["phases"][mode] = {
645
+ "steps": trainer.state.global_step, "seconds": round(dt, 2),
646
+ "steps_per_s": round(trainer.state.global_step / dt, 3),
647
+ "rows_seen": len(trainer.seen_rows),
648
+ "rows_per_s": round(trainer.rows_processed / dt, 2),
649
+ "pairs_seen": trainer.pairs_processed,
650
+ "pairs_per_s": round(trainer.pairs_processed / dt, 1),
651
+ "tokens_per_s": round(trainer.tokens_processed / dt),
652
+ "max_mem_gb": round(torch.cuda.max_memory_allocated() / 1e9, 2)}
653
+ log_metric({"event": "pilot_phase", "pool": mode,
654
+ **results["phases"][mode]})
655
+ checkpoint = os.path.join(args.save_dir, "checkpoint_" + mode)
656
+ model.save_pretrained(checkpoint)
657
+ tokenizer.save_pretrained(checkpoint)
658
+ del trainer
659
+ if mode == "full":
660
+ del model
661
+ torch.cuda.empty_cache()
662
+ results["latency"] = gpu_latency(model, pair_ds, device)
663
+ log_metric({"event": "latency_gpu", **results["latency"]})
664
+ results["env"] = {"gpu": gpu_name, "torch": torch.__version__}
665
+ os.makedirs(args.save_dir, exist_ok=True)
666
+ with open(os.path.join(args.save_dir, "pilot_results.json"), "w") as f:
667
+ json.dump(results, f, indent=2, default=str)
668
+ save_metrics(os.path.join(args.save_dir, "pilot_metrics.jsonl"))
669
+ reload_model = AutoModelForSequenceClassification.from_pretrained(os.path.join(args.save_dir, "checkpoint_sampled4"), attn_implementation=ATTN_IMPL).to(device)
670
+ rec = pair_ds[[0]]
671
+ ids = torch.tensor(rec["input_ids"], device=device)
672
+ att = torch.tensor(rec["attention_mask"], device=device)
673
+ model.eval(); reload_model.eval()
674
+ saved_weights = reload_model.state_dict()
675
+ for name, value in model.state_dict().items():
676
+ assert torch.equal(value, saved_weights[name]), f"Checkpoint changed parameter: {name}"
677
+ with torch.no_grad(), torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
678
+ expected_logits = model(input_ids=ids, attention_mask=att).logits.float()
679
+ actual_logits = reload_model(input_ids=ids, attention_mask=att).logits.float()
680
+ assert torch.isfinite(actual_logits).all()
681
+ assert torch.allclose(expected_logits, actual_logits, atol=1e-4, rtol=1e-4), "Same-precision reload mismatch"
682
+ results["checkpoint_parameter_equality"] = True
683
+ results["checkpoint_reload_verified"] = True
684
+ with open(os.path.join(args.save_dir, "pilot_results.json"), "w") as f:
685
+ json.dump(results, f, indent=2)
686
+ if args.push:
687
+ from huggingface_hub import HfApi
688
+ api = HfApi()
689
+ for filename in ("pilot_results.json", "pilot_metrics.jsonl"):
690
+ api.upload_file(path_or_fileobj=os.path.join(args.save_dir, filename), path_in_repo="pilot/" + filename, repo_id=args.hub_model_id, repo_type="model")
691
+ print("PILOT_DONE", flush=True)
692
+ return
693
+
694
+ # ---------------- PROTOTYPE ----------------
695
+ save_dir = args.save_dir
696
+ os.makedirs(save_dir, exist_ok=True)
697
+ log_metric({"event": "prototype_start", "pool": args.pool,
698
+ "n_rows_selected": len(parsed), "n_a": n_a, "n_t": n_t})
699
+
700
+ val_full, _ = parse_eval_rows(ds["validation"].filter(
701
+ lambda r: r["task_family"] in FOCUS, num_proc=8))
702
+ val_sub = stratified_subset(val_full, 600, seed=SEED)
703
+ val_pd, _ = tokenize_eval(val_sub, tokenizer)
704
+
705
+ model = load_model().to(device)
706
+ sampler = WholeRowBatchSampler(parsed, indexer, args.pool, args.token_budget)
707
+ targs = validated_training_args
708
+ trainer = GroupTrainer(model=model, args=targs, train_dataset=TrainRefs(sampler))
709
+ trainer._sampler = sampler
710
+ trainer._collate = collate
711
+ trainer._refs_ds = TrainRefs(sampler)
712
+
713
+ best = {"acc": -1.0, "step": -1}
714
+
715
+ def run_val(step):
716
+ acc = evaluate_rows(model, val_pd, val_sub, device)
717
+ acc.pop("pred_labels")
718
+ log_metric({"event": "val_acc", "step": step,
719
+ "n_rows": len(val_sub), **acc})
720
+ if acc["acc"] > best["acc"]:
721
+ best.update(acc=acc["acc"], step=step)
722
+
723
+
724
+ class Guards(TrainerCallback):
725
+ def on_step_end(self, targs2, state, control, **kw):
726
+ if state.global_step <= 5 or state.global_step % 100 == 0:
727
+ elapsed = time.time() - t_start
728
+ log_metric({"event": "coverage_progress", "step": state.global_step,
729
+ "rows_seen": len(trainer.seen_rows), "target": len(parsed),
730
+ "elapsed_seconds": round(elapsed, 1)})
731
+ if state.global_step % 1000 == 0:
732
+ latest = os.path.join(save_dir, "latest-checkpoint")
733
+ model.save_pretrained(latest)
734
+ tokenizer.save_pretrained(latest)
735
+ with open(os.path.join(latest, "coverage.json"), "w") as f:
736
+ json.dump({"rows_seen": len(trainer.seen_rows), "step": state.global_step}, f)
737
+ if state.global_step % args.eval_steps == 0 and state.global_step > 0:
738
+ run_val(state.global_step)
739
+ if time.time() - t_start > args.deadline_seconds - args.eval_reserve_seconds:
740
+ control.should_training_stop = True
741
+ log_metric({"event": "time_guard_stop", "step": state.global_step,
742
+ "rows_cum": len(trainer.seen_rows)})
743
+
744
+ trainer.add_callback(Guards())
745
+ t0 = time.time()
746
+ trainer.train()
747
+ train_seconds = time.time() - t0
748
+ rows_trained = len(trainer.seen_rows)
749
+ steps_done = trainer.state.global_step
750
+ log_metric({"event": "train_done", "pool": args.pool,
751
+ "seconds": round(train_seconds, 1), "steps": steps_done,
752
+ "rows_covered": rows_trained, "n_rows_selected": len(parsed),
753
+ "n_pairs_processed": trainer.pairs_processed, "n_truncated_pairs": trainer.truncated_pairs,
754
+ "note": "rows_covered counts unique decisions actually iterated; "
755
+ "no full-epoch guarantee"})
756
+
757
+ # One fixed epoch: final weights are the selected checkpoint.
758
+ # Validation is monitored without rewinding to a partially trained checkpoint.
759
+ model_dir = os.path.join(save_dir, "model")
760
+ model.save_pretrained(model_dir)
761
+ tokenizer.save_pretrained(model_dir)
762
+ with open(os.path.join(save_dir, "training_coverage.json"), "w") as f:
763
+ json.dump({"target": len(parsed), "rows_seen": rows_trained,
764
+ "complete": rows_trained == len(parsed),
765
+ "steps": steps_done, "selection": "final fixed-epoch checkpoint",
766
+ "seen_row_ids": sorted(parsed[i]["row_id"] for i in trainer.seen_rows)}, f)
767
+ if args.push:
768
+ from huggingface_hub import HfApi
769
+ api = HfApi()
770
+ api.upload_folder(folder_path=model_dir, repo_id=args.hub_model_id, repo_type="model",
771
+ commit_message="Persist final prototype before evaluation")
772
+ api.upload_file(path_or_fileobj=os.path.join(save_dir, "training_coverage.json"),
773
+ path_in_repo="training_coverage.json", repo_id=args.hub_model_id, repo_type="model")
774
+ log_metric({"event": "checkpoint_saved_before_evaluation", "rows_covered": rows_trained})
775
+
776
+ results = {"model": "ModernJEV-Decide-Preview",
777
+ "dataset": {"id": DS_ID, "revision": DS_REV},
778
+ "base": {"id": BASE_ID, "revision": BASE_REV},
779
+ "train_pool": args.pool, "max_len": MAX_LEN, "seed": SEED,
780
+ "n_rows_selected": len(parsed), "n_a": n_a, "n_t": n_t,
781
+ "selected_row_ids_sha256": ids_hash,
782
+ "rows_covered": rows_trained, "steps": steps_done,
783
+ "train_seconds": round(train_seconds, 1),
784
+ "lr": args.lr, "token_budget": args.token_budget,
785
+ "grad_accum": args.grad_accum, "validation_monitor": best,
786
+ "input_preparation": "lazy per batch, no upfront map",
787
+ "training_input_stats": {"pairs": trainer.pairs_processed, "truncated": trainer.truncated_pairs, "at_max": trainer.pairs_at_max},
788
+ "checkpoint_selection": "final fixed-epoch checkpoint", "complete_training_coverage": rows_trained == len(parsed),
789
+ "gpu": gpu_name, "torch": torch.__version__, "attention": ATTN_IMPL,
790
+ "transformers": transformers.__version__}
791
+
792
+ val_pd_full, _ = tokenize_eval(val_full, tokenizer)
793
+ results["val_full"] = {k: v for k, v in
794
+ evaluate_rows(model, val_pd_full, val_full, device).items()
795
+ if k != "pred_labels"}
796
+ log_metric({"event": "val_full", **results["val_full"]})
797
+
798
+ test_focus, skipped_te = parse_eval_rows(ds["test"].filter(
799
+ lambda r: r["task_family"] in FOCUS, num_proc=8))
800
+ log_metric({"event": "test_prep", "n_rows": len(test_focus),
801
+ "skipped": skipped_te})
802
+ te_pd, _ = tokenize_eval(test_focus, tokenizer)
803
+ te = evaluate_rows(model, te_pd, test_focus, device)
804
+ results["test"] = {k: v for k, v in te.items() if k != "pred_labels"}
805
+ log_metric({"event": "final_test", **results["test"]})
806
+
807
+ inv_rows = stratified_subset(test_focus, 300, seed=SEED)
808
+ results["shuffled_invariance"] = shuffled_invariance(
809
+ model, tokenizer, inv_rows, device)
810
+ log_metric({"event": "shuffled_invariance", **results["shuffled_invariance"]})
811
+
812
+ results["baselines"] = build_baseline_results(parsed, val_full, test_focus)
813
+ log_metric({"event": "baselines", **results["baselines"]})
814
+
815
+ torch.manual_seed(SEED)
816
+ base_model = load_model().to(device)
817
+ results["baseline_untrained_head"] = {
818
+ "val": {k: v for k, v in evaluate_rows(
819
+ base_model, val_pd_full, val_full, device).items() if k != "pred_labels"},
820
+ "test": {k: v for k, v in evaluate_rows(
821
+ base_model, te_pd, test_focus, device).items() if k != "pred_labels"}}
822
+ log_metric({"event": "baseline_untrained_head",
823
+ **results["baseline_untrained_head"]})
824
+ del base_model
825
+ torch.cuda.empty_cache()
826
+
827
+ results["latency_gpu"] = gpu_latency(model, te_pd, device)
828
+ log_metric({"event": "latency_gpu", **results["latency_gpu"]})
829
+
830
+ results["probe"] = {"omitted": True, "reason": "Budget reserved for full prototype coverage and held-out evaluation"}
831
+
832
+ model_dir = os.path.join(save_dir, "model")
833
+ with open(os.path.join(save_dir, "results.json"), "w") as f:
834
+ json.dump(results, f, indent=2, default=str)
835
+ save_metrics(os.path.join(save_dir, "metrics.jsonl"))
836
+ with gzip.open(os.path.join(save_dir, "selected_row_ids.json.gz"), "wt") as f:
837
+ json.dump({"sha256": ids_hash, "n": len(parsed), "n_a": n_a, "n_t": n_t,
838
+ "pool": args.pool, "rows_covered": rows_trained,
839
+ "row_ids": [r["row_id"] for r in parsed]}, f)
840
+ here = os.path.dirname(os.path.abspath(__file__))
841
+ if os.path.exists(os.path.join(here, "predict.py")):
842
+ import shutil
843
+ shutil.copy(os.path.join(here, "predict.py"),
844
+ os.path.join(save_dir, "predict.py"))
845
+ if args.push:
846
+ from huggingface_hub import HfApi
847
+ api = HfApi()
848
+ assert api.model_info(args.hub_model_id).private, "Private model required"
849
+ for name in ["results.json", "metrics.jsonl", "selected_row_ids.json.gz",
850
+ "predict.py"]:
851
+ api.upload_file(path_or_fileobj=os.path.join(save_dir, name),
852
+ repo_id=args.hub_model_id, repo_type="model",
853
+ path_in_repo=name)
854
+ log_metric({"event": "persisted", "save_dir": save_dir, "pushed": args.push})
855
+ print("PROTOTYPE_DONE", flush=True)
856
+
857
+
858
+ if __name__ == "__main__":
859
+ main()
results.json ADDED
@@ -0,0 +1,195 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model": "ModernJEV-Decide-Preview",
3
+ "dataset": {
4
+ "id": "MaziyarPanahi/AgentToolDecisions-180K",
5
+ "revision": "f2fb14e4ec977c420f376c08785664cd38763d7e"
6
+ },
7
+ "base": {
8
+ "id": "answerdotai/ModernBERT-base",
9
+ "revision": "8949b909ec900327062f0ebf497f51aef5e6f0c8"
10
+ },
11
+ "train_pool": "sampled4",
12
+ "max_len": 4096,
13
+ "seed": 42,
14
+ "n_rows_selected": 60000,
15
+ "n_a": 40809,
16
+ "n_t": 19191,
17
+ "selected_row_ids_sha256": "c5b306b0471ba104161051a7241b3bce4b69e1d959ff3e34fb06ef8eb4b077d9",
18
+ "rows_covered": 60000,
19
+ "steps": 3720,
20
+ "train_seconds": 7786.5,
21
+ "lr": 2e-05,
22
+ "token_budget": 114688,
23
+ "grad_accum": 2,
24
+ "validation_monitor": {
25
+ "acc": 0.6944908180300501,
26
+ "step": 3200
27
+ },
28
+ "input_preparation": "lazy per batch, no upfront map",
29
+ "training_input_stats": {
30
+ "pairs": 199191,
31
+ "truncated": 9249,
32
+ "at_max": 9265
33
+ },
34
+ "checkpoint_selection": "final fixed-epoch checkpoint",
35
+ "complete_training_coverage": true,
36
+ "gpu": "NVIDIA A100-SXM4-80GB",
37
+ "torch": "2.12.0+cu126",
38
+ "attention": "kernels-community/flash-attn2@f50dc99ed079b35990bc895d43fd353ea0cb376d",
39
+ "transformers": "5.17.0",
40
+ "val_full": {
41
+ "acc": 0.6840934371523916,
42
+ "n": 1798,
43
+ "macro_accuracy": 0.6529034180751321,
44
+ "per_family": {
45
+ "agent_next_action_type": {
46
+ "acc": 0.7348912167606769,
47
+ "n": 1241
48
+ },
49
+ "tool_selection": {
50
+ "acc": 0.5709156193895871,
51
+ "n": 557
52
+ }
53
+ },
54
+ "n_pairs": 14449,
55
+ "n_truncated_pairs": 443,
56
+ "seconds": 266.6
57
+ },
58
+ "test": {
59
+ "acc": 0.6982352941176471,
60
+ "n": 1700,
61
+ "macro_accuracy": 0.6778976986661058,
62
+ "per_family": {
63
+ "agent_next_action_type": {
64
+ "acc": 0.7340241796200345,
65
+ "n": 1158
66
+ },
67
+ "tool_selection": {
68
+ "acc": 0.6217712177121771,
69
+ "n": 542
70
+ }
71
+ },
72
+ "n_pairs": 14251,
73
+ "n_truncated_pairs": 414,
74
+ "seconds": 274.1
75
+ },
76
+ "shuffled_invariance": {
77
+ "invariance_rate": 1.0,
78
+ "perms": 5,
79
+ "n_rows": 299
80
+ },
81
+ "baselines": {
82
+ "val": {
83
+ "uniform_expected": {
84
+ "acc": 0.24665407679797663,
85
+ "n": 1798,
86
+ "per_family": {
87
+ "agent_next_action_type": {
88
+ "acc": 0.33333333333332843,
89
+ "n": 1241
90
+ },
91
+ "tool_selection": {
92
+ "acc": 0.05353207076499354,
93
+ "n": 557
94
+ }
95
+ },
96
+ "macro_accuracy": 0.193432702049161
97
+ },
98
+ "train_frequency": {
99
+ "acc": 0.47052280311457173,
100
+ "n": 1798,
101
+ "per_family": {
102
+ "agent_next_action_type": {
103
+ "acc": 0.5511684125705076,
104
+ "n": 1241
105
+ },
106
+ "tool_selection": {
107
+ "acc": 0.29084380610412924,
108
+ "n": 557
109
+ }
110
+ },
111
+ "macro_accuracy": 0.4210061093373184
112
+ }
113
+ },
114
+ "test": {
115
+ "uniform_expected": {
116
+ "acc": 0.24359277413013672,
117
+ "n": 1700,
118
+ "per_family": {
119
+ "agent_next_action_type": {
120
+ "acc": 0.33333333333332943,
121
+ "n": 1158
122
+ },
123
+ "tool_selection": {
124
+ "acc": 0.051859254651728665,
125
+ "n": 542
126
+ }
127
+ },
128
+ "macro_accuracy": 0.19259629399252903
129
+ },
130
+ "train_frequency": {
131
+ "acc": 0.43529411764705883,
132
+ "n": 1700,
133
+ "per_family": {
134
+ "agent_next_action_type": {
135
+ "acc": 0.531951640759931,
136
+ "n": 1158
137
+ },
138
+ "tool_selection": {
139
+ "acc": 0.22878228782287824,
140
+ "n": 542
141
+ }
142
+ },
143
+ "macro_accuracy": 0.3803669642914046
144
+ }
145
+ }
146
+ },
147
+ "baseline_untrained_head": {
148
+ "val": {
149
+ "acc": 0.2925472747497219,
150
+ "n": 1798,
151
+ "macro_accuracy": 0.21439969214610907,
152
+ "per_family": {
153
+ "agent_next_action_type": {
154
+ "acc": 0.41982272360999195,
155
+ "n": 1241
156
+ },
157
+ "tool_selection": {
158
+ "acc": 0.008976660682226212,
159
+ "n": 557
160
+ }
161
+ },
162
+ "n_pairs": 14449,
163
+ "n_truncated_pairs": 443,
164
+ "seconds": 267.2
165
+ },
166
+ "test": {
167
+ "acc": 0.29411764705882354,
168
+ "n": 1700,
169
+ "macro_accuracy": 0.2178523857777438,
170
+ "per_family": {
171
+ "agent_next_action_type": {
172
+ "acc": 0.4283246977547496,
173
+ "n": 1158
174
+ },
175
+ "tool_selection": {
176
+ "acc": 0.007380073800738007,
177
+ "n": 542
178
+ }
179
+ },
180
+ "n_pairs": 14251,
181
+ "n_truncated_pairs": 414,
182
+ "seconds": 274.4
183
+ }
184
+ },
185
+ "latency_gpu": {
186
+ "measurement": "one candidate forward pass, excludes tokenization",
187
+ "warmup_pairs": 5,
188
+ "pair_ms_p50": 23.6,
189
+ "pair_ms_p95": 24.2
190
+ },
191
+ "probe": {
192
+ "omitted": true,
193
+ "reason": "Budget reserved for full prototype coverage and held-out evaluation"
194
+ }
195
+ }
runtime-versions.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "torch": "2.12.0+cu126",
3
+ "transformers": "5.17.0",
4
+ "datasets": "5.0.1",
5
+ "accelerate": "1.15.0",
6
+ "huggingface-hub": "1.33.0",
7
+ "kernels": "0.16.0"
8
+ }
selected_row_ids.json.gz ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cd7a24a0cd78a6908610758f4ed1cc7d0cac3ca89442bdaefc3ef92c114d0f72
3
+ size 1186946
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "backend": "tokenizers",
3
+ "clean_up_tokenization_spaces": true,
4
+ "cls_token": "[CLS]",
5
+ "is_local": false,
6
+ "local_files_only": false,
7
+ "mask_token": "[MASK]",
8
+ "model_input_names": [
9
+ "input_ids",
10
+ "attention_mask"
11
+ ],
12
+ "model_max_length": 8192,
13
+ "pad_token": "[PAD]",
14
+ "sep_token": "[SEP]",
15
+ "tokenizer_class": "TokenizersBackend",
16
+ "unk_token": "[UNK]"
17
+ }
training-status.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "status": "training_preparation",
3
+ "model": "ModernJEV-Decide-Preview",
4
+ "training_decisions_target": 60000,
5
+ "max_sequence_length": 4096,
6
+ "compute_budget_usd": 10,
7
+ "weights_available": false,
8
+ "evaluation_available": false
9
+ }
training_coverage.json ADDED
The diff for this file is too large to render. See raw diff
 
verification/final-coverage.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "target": 60000,
3
+ "rows_seen": 60000,
4
+ "complete": true,
5
+ "steps": 3720,
6
+ "selection": "final fixed-epoch checkpoint",
7
+ "exact_selected_row_set_verified": true
8
+ }
verification/full-training-branch-smoke.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "passed": true,
3
+ "test": "complete one-epoch production trainer and callback path on 24 synthetic decisions",
4
+ "optimizer_steps": 6,
5
+ "rows_seen": 24,
6
+ "validation_calls": 3,
7
+ "checkpoint_parameter_equality": true,
8
+ "lazy_tokens_equal_to_eager": true,
9
+ "long_input_truncation_verified": true,
10
+ "warmup_steps_for_100": 3,
11
+ "transformers": "5.17.0",
12
+ "not_a_quality_evaluation": true
13
+ }
verification/full-training-worker-smoke.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "passed": true,
3
+ "test": "complete production training path with two forked prefetch workers on 24 synthetic decisions",
4
+ "optimizer_steps": 6,
5
+ "rows_seen": 24,
6
+ "validation_calls": 3,
7
+ "checkpoint_parameter_equality": true,
8
+ "lazy_tokens_equal_to_eager": true,
9
+ "long_input_truncation_verified": true,
10
+ "warmup_steps_for_100": 3,
11
+ "transformers": "5.17.0",
12
+ "not_a_quality_evaluation": true
13
+ }
verification/training-start.log ADDED
@@ -0,0 +1,170 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0
  0%| | 0/3720 [00:00<?, ?it/s]
 
 
1
  0%| | 1/3720 [00:02<2:48:33, 2.72s/it]
 
2
 
 
 
3
  0%| | 1/3720 [00:02<2:48:33, 2.72s/it]
 
 
4
  0%| | 2/3720 [00:04<1:57:23, 1.89s/it]
 
 
5
  0%| | 3/3720 [00:05<1:41:43, 1.64s/it]
 
 
6
  0%| | 4/3720 [00:06<1:35:31, 1.54s/it]
 
 
1
+ WARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager, possibly rendering your system unusable. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv. Use the --root-user-action option if you know what you are doing and want to suppress this warning.
2
+ WARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager, possibly rendering your system unusable. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv. Use the --root-user-action option if you know what you are doing and want to suppress this warning.
3
+
4
+ [notice] A new release of pip is available: 25.0.1 -> 26.2.1
5
+ [notice] To update, run: pip install --upgrade pip
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
+ CUDA_KERNEL_VERIFIED {"gpu": "NVIDIA A100-SXM4-80GB", "torch": "2.12.0+cu126"}
70
+
71
+
72
+
73
+
74
+
75
+
76
+ ACTUAL_TRAINING_LAUNCH {"n_rows": 60000, "max_len": 4096, "input_preparation": "lazy per batch; no upfront map", "remaining_seconds": 9571, "warmup_steps": 0.03}
77
+ METRIC {"event": "env", "mode": "prototype", "device": "cuda", "gpu": "NVIDIA A100-SXM4-80GB", "torch": "2.12.0+cu126", "transformers": "5.17.0", "attention": "kernels-community/flash-attn2@f50dc99ed079b35990bc895d43fd353ea0cb376d", "ts": 1790774736.9}
78
+ METRIC {"event": "training_api_validated", "mode": "prototype", "ts": 1790774737.1}
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
+ METRIC {"event": "data_prep", "n_rows": 60000, "n_a": 40809, "n_t": 19191, "family_counts": {"agent_next_action_type": 40809, "tool_selection": 19191}, "gold_class_counts": {"agent_next_action_type": {"tool_call": 19173, "text_response": 21636}, "tool_selection": {"distinct": 2840}}, "selected_row_ids_sha256": "c5b306b0471ba104161051a7241b3bce4b69e1d959ff3e34fb06ef8eb4b077d9", "seconds": 25.0, "ts": 1790774762.8}
123
+ METRIC {"event": "lazy_pairs_ready", "tag": "train", "n_rows": 60000, "n_pairs": 199191, "max_len": 4096, "mapped_pairs": 0, "ts": 1790774763.4}
124
+ METRIC {"event": "prototype_start", "pool": "sampled4", "n_rows_selected": 60000, "n_a": 40809, "n_t": 19191, "ts": 1790774763.9}
125
+
126
+
127
+ METRIC {"event": "lazy_pairs_ready", "tag": "eval", "n_rows": 599, "n_pairs": 4767, "max_len": 4096, "mapped_pairs": 0, "ts": 1790774765.7}
128
+
129
+
130
+
131
+
132
+
133
+
134
+
135
+
136
+
137
+
138
+
139
+
140
+
141
+
142
+
143
+
144
+
145
+
146
+
147
+ [transformers] ModernBertForSequenceClassification LOAD REPORT from: answerdotai/ModernBERT-base
148
+ Key | Status |
149
+ ------------------+------------+-
150
+ decoder.bias | UNEXPECTED |
151
+ classifier.bias | MISSING |
152
+ classifier.weight | MISSING |
153
+
154
+ Notes:
155
+ - UNEXPECTED: can be ignored when loading from different task/architecture; not ok if you expect identical arch.
156
+ - MISSING: those params were newly initialized because missing from the checkpoint. Consider training on your downstream task.
157
+
158
+
159
  0%| | 0/3720 [00:00<?, ?it/s]
160
+ METRIC {"event": "coverage_progress", "step": 1, "rows_seen": 17, "target": 60000, "elapsed_seconds": 44.8, "ts": 1790774781.7}
161
+
162
  0%| | 1/3720 [00:02<2:48:33, 2.72s/it]
163
+
164
 
165
+ {'loss': '1.181', 'grad_norm': '2.406', 'learning_rate': '0', 'epoch': '0.0002688'}
166
+
167
  0%| | 1/3720 [00:02<2:48:33, 2.72s/it]
168
+ METRIC {"event": "coverage_progress", "step": 2, "rows_seen": 33, "target": 60000, "elapsed_seconds": 46.2, "ts": 1790774783.0}
169
+
170
  0%| | 2/3720 [00:04<1:57:23, 1.89s/it]
171
+ METRIC {"event": "coverage_progress", "step": 3, "rows_seen": 49, "target": 60000, "elapsed_seconds": 47.5, "ts": 1790774784.4}
172
+
173
  0%| | 3/3720 [00:05<1:41:43, 1.64s/it]
174
+ METRIC {"event": "coverage_progress", "step": 4, "rows_seen": 65, "target": 60000, "elapsed_seconds": 48.9, "ts": 1790774785.7}
175
+
176
  0%| | 4/3720 [00:06<1:35:31, 1.54s/it]
177
+ METRIC {"event": "coverage_progress", "step": 5, "rows_seen": 81, "target": 60000, "elapsed_seconds": 50.4, "ts": 1790774787.2}
workflow/ONE-PROMPT-REPLAY.md ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ML Intern one-prompt recipe replay: provenance
2
+
3
+ ## What this package contains
4
+
5
+ The model weights in the parent repository are the original 60,000-decision checkpoint. This replay is a separate 6,000-decision checkpoint. None of its metrics replace the parent model's measurements.
6
+
7
+ ## Attempt 1 — blocked before compute
8
+
9
+ Conversation: https://huggingface.co/chat/conversation/6abe15a0d2dc32fb39119ec9
10
+ The source recipe could not be read via the HuggingChat file tools, including at its canonical private namespace. A local HF read-only check could access it. The tool failure is documented, but its platform/permission root cause is not conclusively diagnosed. ML Intern created a private provenance repository and stopped before GPU spend.
11
+
12
+ ## Attempt 2 — training and upload verified
13
+
14
+ Conversation: https://huggingface.co/chat/conversation/6abe292848531444d725312b
15
+ Job: https://huggingface.co/jobs/OpenMed/6abe2b6ffbc85ba6823612cc
16
+ Destination: OpenMed/ModernJEV-Decide-OnePrompt-6K-20261001-R2 (private)
17
+ One execution message contained the attached prompt and tested source bundle. OpenMed was selected in HuggingChat Billing; the conversation budget control was set to $10 after submission. Source preparation and those UI setup actions are disclosed. Codex performed only read-only observation/verification after submission: zero follow-up execution messages, zero code patches and zero job launches.
18
+
19
+ ML Intern adapted the source, corrected its own rehearsal code, froze 6,000 rows (4,080 next action and 1,920 tool selection), performed a production-path rehearsal, reset to base for training, and launched the A100 job using its own tools. The runtime reported Torch 2.12.0+cu130, differing from the supplied cu126 reference; the pinned attention kernel ran successfully.
20
+
21
+ Verified training coverage: 6,000/6,000 unique IDs; exact set equality to the frozen selection; 1,500 optimizer steps. Frozen selection hash: b1ce539158acc3f52319e843a17826021c1da30394f2632da750a8f97b6900eb.
22
+ Checkpoint upload revision: 72cc524bfeefb2c025fc8d99dbfeafccb64259ac.
23
+ Training elapsed (reported runtime metric): 996.9 seconds. Job setup, chat preparation and evaluation are additional.
24
+ Rehearsal checkpoint equality passed. The first GPU run logged exact final-checkpoint reload and changes in all 136 backbone tensors. Its helper-copy finalizer then failed on undefined `__file__`, and the When2Call baseline report reversed numerator/denominator. ML Intern detected both and launched its one authorized corrective retry without follow-up instructions. The retry trained the same frozen 6,000 decisions in 1,500 optimizer steps (1006.3 seconds) and pushed checkpoint dc4067f0de30b8d46a0330d0bbaf70984577ce82. The owner subsequently narrowed the objective to one-prompt launch and real training. At 10:59 UTC Codex stopped generation and canceled the retry A100 and preparation CPU sandbox; the monitor was paused. Full replay evaluation/packaging was not completed. This supports the narrower training workflow claim; it does not satisfy the original full packaging contract.
25
+
26
+ The job has a 90-minute timeout at $2.50/hour: $3.75 maximum GPU reservation for this job, not the actual cost. CPU sandbox costs, any failures and inference accounting must be recorded separately. The $6.23 remaining displayed by HuggingChat includes outstanding holds; it is not a final spend receipt. No public publication or Space change was performed.
27
+
28
+ ## Supported claim
29
+
30
+ With a tested recipe attached and billing/budget configured, one execution message in a fresh HuggingChat ML Intern conversation led to a real 6,000-decision training run and a private checkpoint upload, without additional implementation instructions from the operator or Codex. A prior attempt failed at private-source access and is disclosed above. This is one observed workflow execution, not proof of universal reliability.
31
+
32
+ ## Repository packaging
33
+
34
+ The owner requested a clean main-branch history for the original 60k model under MaziyarPanahi. Historical source revision identifiers above describe what was used during the experiment; after the authorized history squash, retrieve the supplied recipe-bundle.txt and verify recipe-manifest.json instead of relying on those old revisions remaining fetchable. The source contents and actual training history are preserved in the documentation.
workflow/PROMPT-R2.txt ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Run this experiment end to end in HuggingChat ML Intern. This is my single execution prompt, including authorization to launch bounded compute now. We are testing whether ML Intern can reproduce a working training recipe without further human/Codex instructions. Do not stop after a proposal. Use your own tools to prepare, train, evaluate, save, and monitor to completion. If genuinely blocked, report the exact blocker; do not pretend completion.
2
+
3
+ Billing namespace: OpenMed for inference, sandboxes and EVERY Job. The HuggingChat Billing selector is already OpenMed; verify your job namespace too. New destination: OpenMed/ModernJEV-Decide-OnePrompt-6K-20261001-R2, PRIVATE model repository. Never overwrite the existing OpenMed/ModernJEV-Decide-Preview. Never create, update, deploy, or change visibility of any Space. No dashboard Spaces. Do not publish anything.
4
+
5
+ Budget: maximum $10 total new Jobs/sandbox compute including preflight and failed attempts, no extensions. One A100 80GB (a100-large) at the verified current price, no H200 and no concurrent GPU jobs. Prefer one job containing rehearsal then training, timeout at most 90 minutes if price is $2.50/hour; otherwise shorten to stay under the cap. Allow at most one autonomous corrective retry only if remaining budget covers its maximum timeout. Stop any idle sandbox when done. Chat inference must also bill OpenMed; report it separately from compute if measurable.
6
+
7
+ Use the previously successful private recipe as input, not its trained weights or metrics:
8
+ source repo MaziyarPanahi/ModernJEV-Decide-Preview
9
+ source revision 5b31f4c2778f186cda66499ec6bf1d922e00149d
10
+ recipe/train.py, predict.py, runtime-versions.json, and available evaluation helpers/reports as format references only.
11
+ The exact source files are supplied in the attached recipe-bundle.txt because the prior attempt demonstrated that your private-source read route returned missing despite the local HF login verifying the source exists. Use the supplied bundle as the recipe source; no private-source download is required. This source-material attachment is explicit preparation, not work you generated. You must verify your own write access to the new PRIVATE model repo before GPU spend. The bundle also includes the previous evaluate_open_labels.py runner; parameterize its old local paths to your new job and pair it with predict.py. Never copy old model metrics as new results. Do not ask for a token pasted into chat. Verify write access with a small provenance file in the new repo. Do not use external Codex or the operator's local login to run the experiment.
12
+
13
+ Train a fresh ModernBERT scalar candidate ranker from answerdotai/ModernBERT-base revision 8949b909ec900327062f0ebf497f51aef5e6f0c8. Dataset: MaziyarPanahi/AgentToolDecisions-180K revision f2fb14e4ec977c420f376c08785664cd38763d7e. Exactly 6000 training decisions, task-stratified seed 42, from agent_next_action_type and tool_selection only, one epoch, 4096-token maximum. Preserve published train/validation/test splits. Freeze and hash selected row IDs before training. The original script's prototype hash assertion is for 60000 rows: parameterize it for the independently frozen 6000-row manifest and assert exact coverage, do not blindly reuse the old hash or remove verification. Use the tested sampled4 training pool, rank ALL declared candidates at evaluation. Use the fixed final-epoch checkpoint, never tune/select on test data.
14
+
15
+ Known implementation requirements: one shared scalar per question/state/candidate pair, cross-entropy grouped by row_id (decision), NOT group_id (episode). Candidate text must not contain gold labels, scores, provenance or candidate indexes. Accept arbitrary caller-supplied answer names/descriptions or unique answer-string lists, not a fixed class vocabulary. Keep the source recipe's input and truncation logic. Document modifications and diff from source recipe.
16
+
17
+ Known successful environment: Python 3.12, torch 2.12.0+cu126, transformers 5.17.0, datasets 5.0.1, accelerate 1.15.0, kernels 0.16.0, huggingface-hub 1.33.0. Check runtime-versions.json for the remaining pins. FlashAttention kernel reference kernels-community/flash-attn2@f50dc99ed079b35990bc895d43fd353ea0cb376d; use MODERNJEV_ATTN for the script. Kernel repo revisions differ from model repo revisions. Transformers 5.17 rejects kernels 0.17.x and removed warmup_ratio: use warmup_steps=0.03. Materialize lazy Datasets Columns before numerical array comparisons. Keep lazy batch tokenization; no ten-minute upfront full-map pass.
18
+
19
+ Before loading all training candidates, instantiate the ACTUAL production TrainingArguments and run a tiny production-path rehearsal: CUDA forward/backward, optimizer update, validation, save and reload. Compare checkpoint tensors/predictions at matching precision; do not compare BF16 and FP32 with an unrealistic 1e-5 tolerance. Then reload the original base model and reset optimizer/seed for the real 6000-row run, keeping rehearsal results separate. Inspect all hardcoded 60000-row assertions, hub destinations and finalizer references before execution. Log optimizer steps, unique rows covered, elapsed time, GPU type, dependency versions, selected-row hash and checkpoint revision.
20
+
21
+ Evaluate fresh results, not copied results from the source model: next action (1158 test decisions), tool choice (542), and unseen When2Call (3652) AFTER checkpoint freeze. Never train on When2Call or the three single-label task families. Reverify per-task constant-majority references 616/1158, 79/542, 1295/3652; distinguish descriptive test-majority references from allowed-choice training-frequency baselines. Retain per-task untrained ModernBERT scalar-head comparisons where feasible, clearly labeling its random head. No overall or macro accuracy in the model card or summary. Any aggregate fields from the old finalizer must be removed from presentation. Tests must not change the training recipe.
22
+
23
+ Test open answers with fresh 2/3/7/20-choice inputs, reordered choices and renamed labels. Report interface acceptance separately from semantic accuracy. Save per-row predictions, per-task correct/total/baselines, training coverage, exact dependencies, executable recipe, prediction helper, and a useful model card with limitations and examples. Verify saved checkpoint reload and weight difference from the base model. Keep everything private.
24
+
25
+ Proof requirement: this is ONE PROMPT WITH AN EXISTING TESTED RECIPE, not training-code invention from scratch. Record all your tool calls, job IDs/URLs, autonomous fixes/retries, cost estimates versus actual charges, source and new checkpoint revisions, and any additional authorization clicks required. Do not claim success merely because a job launched. Success means actual training of all 6000 decisions, a reloadable fresh checkpoint in the new private repo, completed per-task evaluations, and no follow-up implementation instructions from Codex or the human. If only partial success, state exactly what remains. Monitor and deliver the final result without asking me to send a second 'continue' message.
workflow/recipe-bundle.txt ADDED
@@ -0,0 +1,1143 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Verified recipe source bundle. Existing trained weights and evaluation outcomes are NOT included. Source revision 5b31f4c2778f186cda66499ec6bf1d922e00149d. Extract the delimited files exactly, then document adaptations for 6000 decisions. evaluate_open_labels.py references its old local paths: parameterize dataset/model/helper paths for the new job.
2
+
3
+ ===== FILE: recipe/train.py =====
4
+ """ModernJEV-Decide-Preview — training, evaluation, persistence.
5
+
6
+ Dataset : MaziyarPanahi/AgentToolDecisions-180K @ f2fb14e4ec977c420f376c08785664cd38763d7e
7
+ Base : answerdotai/ModernBERT-base @ 8949b909ec900327062f0ebf497f51aef5e6f0c8
8
+ Scope : task_family in {agent_next_action_type, tool_selection} ONLY (choice primitive).
9
+
10
+ Objective: shared-encoder candidate scalar scorer; per-row_id softmax cross-entropy
11
+ over the row's DECLARED candidates. group_id is the EPISODE, not the decision —
12
+ grouping is ALWAYS by row_id (asserted). Candidates enter the softmax as text
13
+ (label + criterion), so variable tool names need no fixed head and no candidate
14
+ index is exposed to the model. Inputs contain NO gold_label / gold_json /
15
+ gold_score / label_source / source metadata.
16
+
17
+ Training pool (--pool):
18
+ full — every declared candidate of the row joins the softmax.
19
+ sampled4 — gold + up to 3 declared negatives, deterministic per (SEED, epoch,
20
+ row_id). This is ordinary sampled-choice softmax within the same loss;
21
+ EVALUATION ALWAYS RANKS ALL DECLARED CANDIDATES regardless of --pool.
22
+ The pilot benchmarks both objectives and the launch report states which one ran.
23
+
24
+ Modes:
25
+ pilot — benchmarks full vs sampled4 throughput/memory, GPU latency, save check.
26
+ prototype — time-guarded training on the exact stratified subset, interval monitoring
27
+ on a stratified subset of the OFFICIAL validation split (full-val eval
28
+ recorded for the final fixed-epoch checkpoint), ONE final test evaluation on all in-scope
29
+ test rows, baselines (uniform / train-frequency / untrained ModernBERT
30
+ head), shuffled-order label invariance, optional frozen-backbone probe,
31
+ and persistence to --save_dir (default /output/modernjev). No Hub push
32
+ unless --push is passed explicitly.
33
+
34
+ No Space is created anywhere; metrics are metrics.jsonl + stdout only.
35
+ """
36
+ import argparse
37
+ import gzip
38
+ import hashlib
39
+ import json
40
+ import os
41
+ import random
42
+ import time
43
+
44
+ import numpy as np
45
+ import torch
46
+ import torch.nn.functional as F
47
+ from datasets import Dataset, load_dataset
48
+ from torch.utils.data import DataLoader, Dataset as TorchDataset, Sampler
49
+ from transformers import (AutoModelForSequenceClassification, AutoTokenizer,
50
+ Trainer, TrainerCallback, TrainingArguments)
51
+
52
+ DS_ID = "MaziyarPanahi/AgentToolDecisions-180K"
53
+ DS_REV = "f2fb14e4ec977c420f376c08785664cd38763d7e"
54
+ BASE_ID = "answerdotai/ModernBERT-base"
55
+ BASE_REV = "8949b909ec900327062f0ebf497f51aef5e6f0c8"
56
+ FOCUS = ("agent_next_action_type", "tool_selection")
57
+ SEED = 42
58
+ MAX_LEN = 4096
59
+ ATTN_IMPL = os.environ.get("MODERNJEV_ATTN", "sdpa")
60
+ EXPECTED = {"train": 171056, "validation": 2713, "test": 6231}
61
+ EXPECTED_FOCUS_TRAIN = 112973
62
+
63
+ METRICS = []
64
+
65
+
66
+ def log_metric(d):
67
+ d["ts"] = round(time.time(), 1)
68
+ METRICS.append(d)
69
+ print("METRIC " + json.dumps(d, default=str), flush=True)
70
+
71
+
72
+ def save_metrics(path):
73
+ with open(path, "w") as f:
74
+ for d in METRICS:
75
+ f.write(json.dumps(d, default=str) + "\n")
76
+
77
+
78
+ def serialize_state(row):
79
+ """Compact the state. Drops the policy key ONLY when it exactly equals the
80
+ first system message (audited dedupe rule)."""
81
+ state = json.loads(row["state_json"])
82
+ conv = state.get("conversation") or []
83
+ policy = state.get("policy")
84
+ first = conv[0] if conv else None
85
+ dup = (policy is not None and isinstance(first, dict)
86
+ and first.get("role") == "system" and first.get("content") == policy)
87
+ compact = {"available_tools": state.get("available_tools") or [],
88
+ "conversation": conv}
89
+ if policy is not None and not dup:
90
+ compact["policy"] = policy
91
+ return json.dumps(compact, ensure_ascii=False)
92
+
93
+
94
+ def parse_row(r):
95
+ criteria = json.loads(r["criteria_json"])
96
+ keys = list(criteria.keys())
97
+ gold = r["gold_label"]
98
+ return {"row_id": r["row_id"], "group_id": r["group_id"],
99
+ "family": r["task_family"],
100
+ "text_a": r["question_text"] + "\n\nSTATE:\n" + serialize_state(r),
101
+ "cand_keys": keys,
102
+ "cand_texts": [f"{k}: {criteria[k]}" for k in keys],
103
+ "gold_idx": keys.index(gold) if gold in keys else None}
104
+
105
+
106
+ def load_and_prepare(n_rows, seed=SEED):
107
+ t0 = time.time()
108
+ ds = load_dataset(DS_ID, revision=DS_REV)
109
+ for split, n in EXPECTED.items():
110
+ assert len(ds[split]) == n, f"{split}: {len(ds[split])} != {n}"
111
+ train = ds["train"].filter(lambda r: r["task_family"] in FOCUS, num_proc=8)
112
+ assert len(train) == EXPECTED_FOCUS_TRAIN, f"{len(train)} != {EXPECTED_FOCUS_TRAIN}"
113
+
114
+ metadata = train.select_columns(["row_id", "task_family"])[:]
115
+ fam_idx = {f: [] for f in FOCUS}
116
+ for i, fam in enumerate(metadata["task_family"]):
117
+ fam_idx[fam].append(i)
118
+ n_a = n_rows * len(fam_idx["agent_next_action_type"]) // len(train)
119
+ n_t = n_rows - n_a
120
+ rng = random.Random(seed)
121
+ picked = []
122
+ for fam, k in (("agent_next_action_type", n_a), ("tool_selection", n_t)):
123
+ all_ids = metadata["row_id"]
124
+ idxs = sorted(fam_idx[fam], key=lambda i: all_ids[i])
125
+ picked += rng.sample(idxs, k)
126
+ assert len(picked) == n_rows
127
+ picked.sort()
128
+ fields = ["row_id", "group_id", "task_family", "question_text", "state_json", "criteria_json", "gold_label"]
129
+ selected = train.select(picked).select_columns(fields)[:]
130
+ parsed = [parse_row(dict(zip(fields, values))) for values in zip(*(selected[f] for f in fields))]
131
+ for r in parsed:
132
+ assert r["gold_idx"] is not None, f"gold missing in {r['row_id']}"
133
+
134
+ fam_counts = {f: sum(1 for r in parsed if r["family"] == f) for f in FOCUS}
135
+ gold_classes = {f: {} for f in FOCUS}
136
+ for r in parsed:
137
+ g = r["cand_keys"][r["gold_idx"]]
138
+ gold_classes[r["family"]][g] = gold_classes[r["family"]].get(g, 0) + 1
139
+ ids_hash = hashlib.sha256(
140
+ "\n".join(sorted(r["row_id"] for r in parsed)).encode()).hexdigest()
141
+ log_metric({"event": "data_prep", "n_rows": len(parsed), "n_a": n_a, "n_t": n_t,
142
+ "family_counts": fam_counts, "gold_class_counts": {f: gold_classes[f] if f == "agent_next_action_type" else {"distinct": len(gold_classes[f])} for f in FOCUS},
143
+ "selected_row_ids_sha256": ids_hash,
144
+ "seconds": round(time.time() - t0, 1)})
145
+ return parsed, ds, ids_hash, (n_a, n_t)
146
+
147
+
148
+ class LazyPairDataset:
149
+ """Tokenize only requested pairs. Metadata stays small; no up-front Dataset.map."""
150
+ def __init__(self, parsed, tokenizer, tag, max_len, training_pool):
151
+ self.parsed, self.tokenizer, self.tag, self.max_len = parsed, tokenizer, tag, max_len
152
+ self.refs = []
153
+ for row in range(len(parsed)):
154
+ candidates, _ = pool_for(parsed, row, 1, training_pool)
155
+ self.refs.extend((row, candidate) for candidate in candidates)
156
+ self.metadata = Dataset.from_dict({
157
+ "row_idx": [r for r, _ in self.refs],
158
+ "cand_idx": [c for _, c in self.refs],
159
+ # Conservative packing bound; actual batch padding uses actual token lengths.
160
+ "length": [max_len] * len(self.refs),
161
+ })
162
+ self.row_sortlen = np.array([
163
+ min(max_len, max(1, len(r["text_a"]) // 4)) for r in parsed], dtype=np.int64)
164
+ log_metric({"event": "lazy_pairs_ready", "tag": tag,
165
+ "n_rows": len(parsed), "n_pairs": len(self.refs),
166
+ "max_len": max_len, "mapped_pairs": 0})
167
+
168
+ def __len__(self):
169
+ return len(self.refs)
170
+
171
+ def select_columns(self, columns):
172
+ return self.metadata.select_columns(columns)
173
+
174
+ def __getitem__(self, indices):
175
+ scalar = isinstance(indices, (int, np.integer))
176
+ if scalar:
177
+ indices = [int(indices)]
178
+ elif isinstance(indices, slice):
179
+ indices = list(range(*indices.indices(len(self))))
180
+ else:
181
+ indices = list(indices)
182
+ refs = [self.refs[int(i)] for i in indices]
183
+ encoded = self.tokenizer(
184
+ [self.parsed[r]["text_a"] for r, c in refs],
185
+ [self.parsed[r]["cand_texts"][c] for r, c in refs],
186
+ truncation="only_first", max_length=self.max_len,
187
+ padding=False, verbose=False)
188
+ result = {"row_idx": [r for r, c in refs],
189
+ "cand_idx": [c for r, c in refs],
190
+ "input_ids": encoded["input_ids"],
191
+ "attention_mask": encoded["attention_mask"],
192
+ "length": [len(ids) for ids in encoded["input_ids"]],
193
+ "truncated": [bool(e.overflowing) for e in encoded.encodings]}
194
+ return {k: v[0] for k, v in result.items()} if scalar else result
195
+
196
+
197
+ def tokenize_pairs(parsed, tokenizer, tag, max_len=MAX_LEN, num_proc=8, training_pool="full"):
198
+ return LazyPairDataset(parsed, tokenizer, tag, max_len, training_pool)
199
+
200
+
201
+ class RowIndexer:
202
+ def __init__(self, parsed, pair_ds):
203
+ meta = pair_ds.select_columns(["row_idx", "cand_idx", "length"]).with_format("numpy")[:]
204
+ assert np.all(np.diff(meta["row_idx"]) >= 0), "pairs must be row-major"
205
+ self.lookup = {}
206
+ self.row_maxlen = np.zeros(len(parsed), dtype=np.int64)
207
+ for flat, (row, cand, length) in enumerate(zip(meta["row_idx"], meta["cand_idx"], meta["length"])):
208
+ row, cand = int(row), int(cand)
209
+ assert (row, cand) not in self.lookup
210
+ self.lookup[(row, cand)] = flat
211
+ self.row_maxlen[row] = max(self.row_maxlen[row], int(length))
212
+ assert np.all(self.row_maxlen > 0)
213
+ self.row_sortlen = getattr(pair_ds, "row_sortlen", self.row_maxlen)
214
+ def flat(self, row, cand):
215
+ return self.lookup[(row, cand)]
216
+
217
+
218
+ def pool_for(parsed, row, epoch, mode):
219
+ """Training pool for one row: (pool_cand_indices, gold_pos_in_pool).
220
+ full: all declared candidates. sampled4: gold + up to 3 declared negatives,
221
+ deterministic per (SEED, epoch, row_id)."""
222
+ r = parsed[row]
223
+ k = len(r["cand_keys"])
224
+ gold = r["gold_idx"]
225
+ if mode == "full" or k <= 4:
226
+ return list(range(k)), gold
227
+ rng = random.Random(f"{SEED}:{epoch}:{r['row_id']}")
228
+ others = sorted(set(range(k)) - {gold})
229
+ negs = rng.sample(others, min(3, len(others)))
230
+ pool = [gold] + sorted(negs)
231
+ rng.shuffle(pool)
232
+ return pool, pool.index(gold)
233
+
234
+
235
+ class TrainRefs(TorchDataset):
236
+ """Position i -> (row, cand_actual) for the current epoch; gold_pos is
237
+ recomputed in collate from the pool definition (deterministic)."""
238
+
239
+ def __init__(self, sampler):
240
+ self.sampler = sampler
241
+
242
+ def __len__(self):
243
+ return len(self.sampler.refs)
244
+
245
+ def __getitem__(self, i):
246
+ return self.sampler.refs[i]
247
+
248
+
249
+ class WholeRowBatchSampler(Sampler):
250
+ """Yields batches of positions into self.refs. Every row's pooled pairs stay
251
+ in one batch (padding-token budget checked at row boundaries)."""
252
+
253
+ def __init__(self, parsed, indexer, mode, token_budget, max_rows=64,
254
+ seed=SEED, shuffle=True, length_buckets=True):
255
+ self.parsed, self.indexer, self.mode = parsed, indexer, mode
256
+ self.token_budget, self.max_rows = token_budget, max_rows
257
+ self.rng = random.Random(seed)
258
+ self.shuffle = shuffle
259
+ self.length_buckets = length_buckets
260
+ self.epoch = 0
261
+ self.refs = []
262
+ self.rows_cum = 0
263
+ self.pairs_cum = 0
264
+ self.tokens_cum = 0
265
+ self.batches = self._build_epoch()
266
+
267
+ def _build_epoch(self):
268
+ self.epoch += 1
269
+ order = list(range(len(self.parsed)))
270
+ if self.shuffle:
271
+ self.rng.shuffle(order)
272
+ # Sort locally within randomized buckets to reduce padding, preserving all rows.
273
+ if self.length_buckets:
274
+ order = [row for begin in range(0, len(order), 256)
275
+ for row in sorted(order[begin:begin + 256], key=lambda r: int(self.indexer.row_sortlen[r]))]
276
+ refs, batches, cur = [], [], []
277
+ batch_maxlen, batch_rows = 0, 0
278
+ for r in order:
279
+ pool, _ = pool_for(self.parsed, r, self.epoch, self.mode)
280
+ row_len = int(self.indexer.row_maxlen[r])
281
+ new_maxlen = max(batch_maxlen, row_len)
282
+ if cur and (new_maxlen * (len(cur) + len(pool)) > self.token_budget
283
+ or batch_rows >= self.max_rows):
284
+ batches.append(cur)
285
+ cur, batch_maxlen, batch_rows = [], 0, 0
286
+ start = len(refs)
287
+ refs.extend((r, cand) for cand in pool)
288
+ cur.extend(range(start, start + len(pool)))
289
+ batch_maxlen = max(batch_maxlen, row_len)
290
+ batch_rows += 1
291
+ if cur:
292
+ batches.append(cur)
293
+ self.refs = refs
294
+ return batches
295
+
296
+ def __iter__(self):
297
+ return iter(self.batches)
298
+
299
+ def __len__(self):
300
+ return len(self.batches)
301
+
302
+
303
+ def make_collate(pair_ds, indexer, pad_id, parsed, mode):
304
+ def collate(items):
305
+ rows = torch.tensor([it[0] for it in items], dtype=torch.long)
306
+ by_row = {}
307
+ for row_index, candidate_index in items:
308
+ by_row.setdefault(row_index, []).append(candidate_index)
309
+ gold_positions = {}
310
+ for row_index, candidate_indices in by_row.items():
311
+ gold_index = parsed[row_index]["gold_idx"]
312
+ assert candidate_indices.count(gold_index) == 1, "Each row needs exactly one gold candidate"
313
+ gold_positions[row_index] = candidate_indices.index(gold_index)
314
+ golds = torch.tensor([gold_positions[it[0]] for it in items], dtype=torch.long)
315
+ flats = [indexer.flat(it[0], it[1]) for it in items]
316
+ rec = pair_ds[flats]
317
+ maxlen = max(rec["length"])
318
+ B = len(items)
319
+ ids = torch.full((B, maxlen), pad_id, dtype=torch.long)
320
+ att = torch.zeros((B, maxlen), dtype=torch.long)
321
+ for i, (ii, aa) in enumerate(zip(rec["input_ids"], rec["attention_mask"])):
322
+ n = len(ii)
323
+ ids[i, :n] = torch.tensor(ii, dtype=torch.long)
324
+ att[i, :n] = torch.tensor(aa, dtype=torch.long)
325
+ return {"input_ids": ids, "attention_mask": att, "row": rows, "gold": golds,
326
+ "_n_truncated": sum(rec["truncated"]), "_n_at_max": sum(n >= MAX_LEN for n in rec["length"])}
327
+ collate.epoch = 1
328
+ return collate
329
+
330
+
331
+ def grouped_ce(logits, rows, gold):
332
+ """Per-row softmax CE. rows must be contiguous per row (asserted)."""
333
+ uniq_c = torch.unique_consecutive(rows)
334
+ uniq_all = torch.unique(rows)
335
+ assert len(uniq_c) == len(uniq_all), "row pairs not contiguous — grouping unsafe"
336
+ counts = torch.unique_consecutive(rows, return_counts=True)[1].tolist()
337
+ losses = []
338
+ ofs = 0
339
+ for c in counts:
340
+ seg = logits[ofs:ofs + c]
341
+ g = int(gold[ofs].item())
342
+ assert g < c, f"gold {g} out of pool size {c}"
343
+ losses.append(F.cross_entropy(seg.unsqueeze(0),
344
+ torch.tensor([g], device=seg.device)))
345
+ ofs += c
346
+ return torch.stack(losses).mean(), len(counts)
347
+
348
+
349
+ class GroupTrainer(Trainer):
350
+ def __init__(self, *args, **kwargs):
351
+ super().__init__(*args, **kwargs)
352
+ self.model_accepts_loss_kwargs = False
353
+
354
+ def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
355
+ gold = inputs.pop("gold")
356
+ rows = inputs.pop("row")
357
+ if not hasattr(self, "seen_rows"):
358
+ self.seen_rows = set()
359
+ self.rows_processed = 0
360
+ self.pairs_processed = 0
361
+ self.tokens_processed = 0
362
+ self.truncated_pairs = 0
363
+ self.pairs_at_max = 0
364
+ self.truncated_pairs += int(inputs.pop("_n_truncated", 0))
365
+ self.pairs_at_max += int(inputs.pop("_n_at_max", 0))
366
+ batch_rows = torch.unique_consecutive(rows).detach().cpu().tolist()
367
+ self.seen_rows.update(batch_rows)
368
+ self.rows_processed += len(batch_rows)
369
+ self.pairs_processed += len(rows)
370
+ self.tokens_processed += int(inputs["attention_mask"].sum().item())
371
+ out = model(input_ids=inputs["input_ids"],
372
+ attention_mask=inputs["attention_mask"])
373
+ logits = out.logits.squeeze(-1).float()
374
+ loss, _ = grouped_ce(logits, rows, gold)
375
+ return (loss, out) if return_outputs else loss
376
+
377
+ def get_train_dataloader(self):
378
+ workers = self.args.dataloader_num_workers
379
+ return DataLoader(self._refs_ds, batch_sampler=self._sampler,
380
+ collate_fn=self._collate, num_workers=workers,
381
+ persistent_workers=workers > 0,
382
+ prefetch_factor=2 if workers else None,
383
+ pin_memory=self.args.device.type == "cuda")
384
+
385
+
386
+ class EvalPairs(TorchDataset):
387
+ def __init__(self, pair_ds):
388
+ self.pair_ds = pair_ds
389
+
390
+ def __len__(self):
391
+ return len(self.pair_ds)
392
+
393
+ def __getitem__(self, i):
394
+ return i
395
+
396
+
397
+ def make_eval_collate(pair_ds, pad_id):
398
+ def collate(idx_list):
399
+ rec = pair_ds[idx_list]
400
+ maxlen = max(rec["length"])
401
+ B = len(idx_list)
402
+ ids = torch.full((B, maxlen), pad_id, dtype=torch.long)
403
+ att = torch.zeros((B, maxlen), dtype=torch.long)
404
+ for i, (ii, aa) in enumerate(zip(rec["input_ids"], rec["attention_mask"])):
405
+ n = len(ii)
406
+ ids[i, :n] = torch.tensor(ii, dtype=torch.long)
407
+ att[i, :n] = torch.tensor(aa, dtype=torch.long)
408
+ return {"input_ids": ids, "attention_mask": att,
409
+ "row": torch.tensor(rec["row_idx"], dtype=torch.long),
410
+ "_n_truncated": sum(rec["truncated"])}
411
+ return collate
412
+
413
+
414
+ @torch.no_grad()
415
+ def evaluate_rows(model, pair_ds, parsed, device, batch_size=8):
416
+ """Row-grouped accuracy over ALL declared candidates."""
417
+ model.eval()
418
+ scores = [[] for _ in range(len(parsed))]
419
+ n_truncated_pairs = 0
420
+ dl = DataLoader(EvalPairs(pair_ds), batch_size=batch_size,
421
+ collate_fn=make_eval_collate(pair_ds, model.config.pad_token_id), num_workers=0)
422
+ t0 = time.time()
423
+ for batch in dl:
424
+ n_truncated_pairs += int(batch["_n_truncated"])
425
+ with torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
426
+ out = model(input_ids=batch["input_ids"].to(device),
427
+ attention_mask=batch["attention_mask"].to(device))
428
+ lg = out.logits.squeeze(-1).float().cpu()
429
+ for j, r in enumerate(batch["row"].tolist()):
430
+ scores[r].append(lg[j].item())
431
+ correct, n = 0, 0
432
+ per_fam = {f: [0, 0] for f in FOCUS}
433
+ pred_labels = []
434
+ for r, row in enumerate(parsed):
435
+ if not scores[r]:
436
+ continue
437
+ p = min(range(len(scores[r])), key=lambda c: (-scores[r][c], row["cand_keys"][c]))
438
+ ok = int(p == row["gold_idx"])
439
+ correct += ok
440
+ n += 1
441
+ per_fam[row["family"]][0] += ok
442
+ per_fam[row["family"]][1] += 1
443
+ pred_labels.append(row["cand_keys"][p])
444
+ model.train()
445
+ return {"acc": correct / max(n, 1), "n": n,
446
+ "macro_accuracy": sum(per_fam[f][0] / max(per_fam[f][1], 1) for f in FOCUS) / len(FOCUS),
447
+ "per_family": {f: {"acc": per_fam[f][0] / max(per_fam[f][1], 1),
448
+ "n": per_fam[f][1]} for f in FOCUS},
449
+ "pred_labels": pred_labels,
450
+ "n_pairs": len(pair_ds), "n_truncated_pairs": n_truncated_pairs,
451
+ "seconds": round(time.time() - t0, 1)}
452
+
453
+
454
+ def stratified_subset(parsed, n, seed=SEED):
455
+ fam_rows = {f: sorted([r for r in parsed if r["family"] == f],
456
+ key=lambda r: r["row_id"]) for f in FOCUS}
457
+ total = sum(len(v) for v in fam_rows.values())
458
+ out = []
459
+ for f in FOCUS:
460
+ k = n * len(fam_rows[f]) // total
461
+ out += random.Random(seed).sample(fam_rows[f], k)
462
+ return out[:n]
463
+
464
+
465
+ def parse_eval_rows(split_ds):
466
+ parsed, skipped = [], 0
467
+ for r in split_ds:
468
+ row = parse_row(r)
469
+ if row["gold_idx"] is None:
470
+ skipped += 1
471
+ continue
472
+ parsed.append(row)
473
+ return parsed, skipped
474
+
475
+
476
+ def tokenize_eval(parsed, tokenizer):
477
+ pd = tokenize_pairs(parsed, tokenizer, tag="eval", num_proc=4)
478
+ idx = RowIndexer(parsed, pd)
479
+ return pd, idx
480
+
481
+
482
+ def build_baseline_results(parsed_tr, parsed_val, parsed_te):
483
+ """Analytic uniform expected accuracy and allowed-choice training frequency."""
484
+ freq = {f: {} for f in FOCUS}
485
+ for row in parsed_tr:
486
+ g = row["cand_keys"][row["gold_idx"]]
487
+ fam_freq = freq[row["family"]]
488
+ fam_freq[g] = fam_freq.get(g, 0) + 1
489
+ res = {}
490
+ for name, parsed in (("val", parsed_val), ("test", parsed_te)):
491
+ methods = {"uniform_expected": {f: [0., 0] for f in FOCUS},
492
+ "train_frequency": {f: [0., 0] for f in FOCUS}}
493
+ for row in parsed:
494
+ fam = row["family"]
495
+ expected = 1. / len(row["cand_keys"])
496
+ best = min(row["cand_keys"], key=lambda key: (-freq[fam].get(key, 0), key))
497
+ matched = int(best == row["cand_keys"][row["gold_idx"]])
498
+ for method, value in (("uniform_expected", expected), ("train_frequency", matched)):
499
+ methods[method][fam][0] += value
500
+ methods[method][fam][1] += 1
501
+ res[name] = {}
502
+ for method, counts in methods.items():
503
+ per = {f: {"acc": c / n, "n": n} for f, (c, n) in counts.items()}
504
+ res[name][method] = {"acc": sum(c for c, _ in counts.values()) / len(parsed),
505
+ "n": len(parsed), "per_family": per,
506
+ "macro_accuracy": sum(v["acc"] for v in per.values()) / len(FOCUS)}
507
+ return res
508
+
509
+
510
+ @torch.no_grad()
511
+ def gpu_latency(model, pair_ds, device, n=100):
512
+ model.eval()
513
+ lat = []
514
+ for i in range(min(n + 5, len(pair_ds))):
515
+ rec = pair_ds[[i]]
516
+ ids = torch.tensor(rec["input_ids"][0], dtype=torch.long,
517
+ device=device).unsqueeze(0)
518
+ att = torch.tensor(rec["attention_mask"][0], dtype=torch.long,
519
+ device=device).unsqueeze(0)
520
+ if device == "cuda":
521
+ torch.cuda.synchronize()
522
+ t = time.time()
523
+ with torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
524
+ model(input_ids=ids, attention_mask=att)
525
+ if device == "cuda":
526
+ torch.cuda.synchronize()
527
+ if i >= 5:
528
+ lat.append(time.time() - t)
529
+ lat.sort()
530
+ return {"measurement": "one candidate forward pass, excludes tokenization", "warmup_pairs": 5,
531
+ "pair_ms_p50": round(1000 * lat[len(lat) // 2], 1),
532
+ "pair_ms_p95": round(1000 * lat[int(0.95 * len(lat))], 1)}
533
+
534
+
535
+ def shuffled_invariance(model, tokenizer, parsed_subset, device, k=5, batch_size=8):
536
+ """Permute candidate order k times; the predicted LABEL must be identical."""
537
+ rng = random.Random(SEED + 1)
538
+ pd0, _ = tokenize_eval(parsed_subset, tokenizer)
539
+ base_labels = evaluate_rows(model, pd0, parsed_subset, device)["pred_labels"]
540
+ agree, total = 0, 0
541
+ for rep in range(k):
542
+ perm = []
543
+ for r in parsed_subset:
544
+ order = list(range(len(r["cand_texts"])))
545
+ rng.shuffle(order)
546
+ perm.append({**r, "cand_keys": [r["cand_keys"][j] for j in order],
547
+ "cand_texts": [r["cand_texts"][j] for j in order],
548
+ "gold_idx": order.index(r["gold_idx"])})
549
+ pd, _ = tokenize_eval(perm, tokenizer)
550
+ acc = evaluate_rows(model, pd, perm, device, batch_size)
551
+ agree += sum(1 for a, b in zip(base_labels, acc["pred_labels"]) if a == b)
552
+ total += len(base_labels)
553
+ return {"invariance_rate": agree / max(total, 1), "perms": k,
554
+ "n_rows": len(parsed_subset)}
555
+
556
+
557
+ def parse_args():
558
+ p = argparse.ArgumentParser()
559
+ p.add_argument("--mode", choices=["pilot", "prototype"], required=True)
560
+ p.add_argument("--pool", choices=["full", "sampled4"], default="sampled4",
561
+ help="training candidate pool (eval always uses all declared)")
562
+ p.add_argument("--n_rows", type=int, default=60000)
563
+ p.add_argument("--max_steps", type=int, default=25)
564
+ p.add_argument("--pilot_sampled_only", action="store_true")
565
+ p.add_argument("--pilot_length_buckets", action="store_true")
566
+ p.add_argument("--token_budget", type=int, default=32768)
567
+ p.add_argument("--grad_accum", type=int, default=2)
568
+ p.add_argument("--lr", type=float, default=2e-5)
569
+ p.add_argument("--eval_steps", type=int, default=400)
570
+ p.add_argument("--deadline_seconds", type=int, default=6600)
571
+ p.add_argument("--eval_reserve_seconds", type=int, default=1200)
572
+ p.add_argument("--save_dir", default="/output/modernjev")
573
+ p.add_argument("--push", action="store_true", default=False)
574
+ p.add_argument("--hub_model_id", default="OpenMed/ModernJEV-Decide-Preview")
575
+ return p.parse_args()
576
+
577
+
578
+ def prototype_training_args(args, save_dir, device):
579
+ return TrainingArguments(
580
+ output_dir=os.path.join(save_dir, "ckpt"), per_device_train_batch_size=1,
581
+ gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr,
582
+ bf16=device == "cuda", use_cpu=device == "cpu",
583
+ num_train_epochs=1.0, logging_steps=25, logging_first_step=True,
584
+ save_strategy="no", eval_strategy="no", report_to="none",
585
+ seed=SEED, remove_unused_columns=False,
586
+ lr_scheduler_type="linear", warmup_steps=0.03,
587
+ dataloader_num_workers=2 if device == "cuda" else 0)
588
+
589
+
590
+ def load_model():
591
+ return AutoModelForSequenceClassification.from_pretrained(
592
+ BASE_ID, revision=BASE_REV, num_labels=1, attn_implementation=ATTN_IMPL)
593
+
594
+
595
+ def main():
596
+ args = parse_args()
597
+ t_start = time.time()
598
+ torch.manual_seed(SEED)
599
+ random.seed(SEED)
600
+ device = "cuda" if torch.cuda.is_available() else "cpu"
601
+ gpu_name = torch.cuda.get_device_name(0) if device == "cuda" else "cpu"
602
+ import transformers
603
+ log_metric({"event": "env", "mode": args.mode, "device": device, "gpu": gpu_name,
604
+ "torch": torch.__version__, "transformers": transformers.__version__, "attention": ATTN_IMPL})
605
+ assert device == "cuda", "GPU required"
606
+ # Validate the complete prototype argument branch before any dataset work.
607
+ validated_training_args = prototype_training_args(args, args.save_dir, device) if args.mode == "prototype" else None
608
+ log_metric({"event": "training_api_validated", "mode": args.mode})
609
+
610
+ tokenizer = AutoTokenizer.from_pretrained(BASE_ID, revision=BASE_REV)
611
+ pad_id = tokenizer.pad_token_id
612
+ assert pad_id is not None, "tokenizer has no pad token"
613
+
614
+ n_train_rows = 500 if args.mode == "pilot" else args.n_rows
615
+ parsed, ds, ids_hash, (n_a, n_t) = load_and_prepare(n_train_rows)
616
+ if args.mode == "prototype":
617
+ assert ids_hash == "c5b306b0471ba104161051a7241b3bce4b69e1d959ff3e34fb06ef8eb4b077d9", "Selected subset differs from authorized frozen manifest"
618
+ os.makedirs(args.save_dir, exist_ok=True)
619
+ pair_ds = tokenize_pairs(parsed, tokenizer, tag="train",
620
+ training_pool=args.pool if args.mode == "prototype" else "full")
621
+ indexer = RowIndexer(parsed, pair_ds)
622
+ collate = make_collate(pair_ds, indexer, pad_id, parsed, args.pool)
623
+
624
+ if args.mode == "pilot":
625
+ results = {"phases": {}}
626
+ pilot_modes = ("sampled4",) if args.pilot_sampled_only else ("full", "sampled4")
627
+ for phase_i, mode in enumerate(pilot_modes):
628
+ torch.manual_seed(SEED)
629
+ phase_steps = args.max_steps if args.pilot_sampled_only else args.max_steps // 2 + (args.max_steps % 2 if phase_i == 1 else 0)
630
+ model = load_model().to(device)
631
+ sampler = WholeRowBatchSampler(parsed, indexer, mode, args.token_budget, length_buckets=args.pilot_length_buckets)
632
+ targs = TrainingArguments(
633
+ output_dir="/tmp/pilot_" + mode, per_device_train_batch_size=1,
634
+ gradient_accumulation_steps=1, learning_rate=args.lr, bf16=True,
635
+ max_steps=phase_steps, logging_steps=3, save_strategy="no",
636
+ eval_strategy="no", report_to="none", seed=SEED,
637
+ remove_unused_columns=False)
638
+ trainer = GroupTrainer(model=model, args=targs,
639
+ train_dataset=TrainRefs(sampler))
640
+ trainer._sampler = sampler
641
+ trainer._collate = make_collate(pair_ds, indexer, pad_id, parsed, mode)
642
+ trainer._refs_ds = TrainRefs(sampler)
643
+ torch.cuda.reset_peak_memory_stats()
644
+ t0 = time.time()
645
+ trainer.train()
646
+ dt = time.time() - t0
647
+ results["phases"][mode] = {
648
+ "steps": trainer.state.global_step, "seconds": round(dt, 2),
649
+ "steps_per_s": round(trainer.state.global_step / dt, 3),
650
+ "rows_seen": len(trainer.seen_rows),
651
+ "rows_per_s": round(trainer.rows_processed / dt, 2),
652
+ "pairs_seen": trainer.pairs_processed,
653
+ "pairs_per_s": round(trainer.pairs_processed / dt, 1),
654
+ "tokens_per_s": round(trainer.tokens_processed / dt),
655
+ "max_mem_gb": round(torch.cuda.max_memory_allocated() / 1e9, 2)}
656
+ log_metric({"event": "pilot_phase", "pool": mode,
657
+ **results["phases"][mode]})
658
+ checkpoint = os.path.join(args.save_dir, "checkpoint_" + mode)
659
+ model.save_pretrained(checkpoint)
660
+ tokenizer.save_pretrained(checkpoint)
661
+ del trainer
662
+ if mode == "full":
663
+ del model
664
+ torch.cuda.empty_cache()
665
+ results["latency"] = gpu_latency(model, pair_ds, device)
666
+ log_metric({"event": "latency_gpu", **results["latency"]})
667
+ results["env"] = {"gpu": gpu_name, "torch": torch.__version__}
668
+ os.makedirs(args.save_dir, exist_ok=True)
669
+ with open(os.path.join(args.save_dir, "pilot_results.json"), "w") as f:
670
+ json.dump(results, f, indent=2, default=str)
671
+ save_metrics(os.path.join(args.save_dir, "pilot_metrics.jsonl"))
672
+ reload_model = AutoModelForSequenceClassification.from_pretrained(os.path.join(args.save_dir, "checkpoint_sampled4"), attn_implementation=ATTN_IMPL).to(device)
673
+ rec = pair_ds[[0]]
674
+ ids = torch.tensor(rec["input_ids"], device=device)
675
+ att = torch.tensor(rec["attention_mask"], device=device)
676
+ model.eval(); reload_model.eval()
677
+ saved_weights = reload_model.state_dict()
678
+ for name, value in model.state_dict().items():
679
+ assert torch.equal(value, saved_weights[name]), f"Checkpoint changed parameter: {name}"
680
+ with torch.no_grad(), torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
681
+ expected_logits = model(input_ids=ids, attention_mask=att).logits.float()
682
+ actual_logits = reload_model(input_ids=ids, attention_mask=att).logits.float()
683
+ assert torch.isfinite(actual_logits).all()
684
+ assert torch.allclose(expected_logits, actual_logits, atol=1e-4, rtol=1e-4), "Same-precision reload mismatch"
685
+ results["checkpoint_parameter_equality"] = True
686
+ results["checkpoint_reload_verified"] = True
687
+ with open(os.path.join(args.save_dir, "pilot_results.json"), "w") as f:
688
+ json.dump(results, f, indent=2)
689
+ if args.push:
690
+ from huggingface_hub import HfApi
691
+ api = HfApi()
692
+ for filename in ("pilot_results.json", "pilot_metrics.jsonl"):
693
+ api.upload_file(path_or_fileobj=os.path.join(args.save_dir, filename), path_in_repo="pilot/" + filename, repo_id=args.hub_model_id, repo_type="model")
694
+ print("PILOT_DONE", flush=True)
695
+ return
696
+
697
+ # ---------------- PROTOTYPE ----------------
698
+ save_dir = args.save_dir
699
+ os.makedirs(save_dir, exist_ok=True)
700
+ log_metric({"event": "prototype_start", "pool": args.pool,
701
+ "n_rows_selected": len(parsed), "n_a": n_a, "n_t": n_t})
702
+
703
+ val_full, _ = parse_eval_rows(ds["validation"].filter(
704
+ lambda r: r["task_family"] in FOCUS, num_proc=8))
705
+ val_sub = stratified_subset(val_full, 600, seed=SEED)
706
+ val_pd, _ = tokenize_eval(val_sub, tokenizer)
707
+
708
+ model = load_model().to(device)
709
+ sampler = WholeRowBatchSampler(parsed, indexer, args.pool, args.token_budget)
710
+ targs = validated_training_args
711
+ trainer = GroupTrainer(model=model, args=targs, train_dataset=TrainRefs(sampler))
712
+ trainer._sampler = sampler
713
+ trainer._collate = collate
714
+ trainer._refs_ds = TrainRefs(sampler)
715
+
716
+ best = {"acc": -1.0, "step": -1}
717
+
718
+ def run_val(step):
719
+ acc = evaluate_rows(model, val_pd, val_sub, device)
720
+ acc.pop("pred_labels")
721
+ log_metric({"event": "val_acc", "step": step,
722
+ "n_rows": len(val_sub), **acc})
723
+ if acc["acc"] > best["acc"]:
724
+ best.update(acc=acc["acc"], step=step)
725
+
726
+
727
+ class Guards(TrainerCallback):
728
+ def on_step_end(self, targs2, state, control, **kw):
729
+ if state.global_step <= 5 or state.global_step % 100 == 0:
730
+ elapsed = time.time() - t_start
731
+ log_metric({"event": "coverage_progress", "step": state.global_step,
732
+ "rows_seen": len(trainer.seen_rows), "target": len(parsed),
733
+ "elapsed_seconds": round(elapsed, 1)})
734
+ if state.global_step % 1000 == 0:
735
+ latest = os.path.join(save_dir, "latest-checkpoint")
736
+ model.save_pretrained(latest)
737
+ tokenizer.save_pretrained(latest)
738
+ with open(os.path.join(latest, "coverage.json"), "w") as f:
739
+ json.dump({"rows_seen": len(trainer.seen_rows), "step": state.global_step}, f)
740
+ if state.global_step % args.eval_steps == 0 and state.global_step > 0:
741
+ run_val(state.global_step)
742
+ if time.time() - t_start > args.deadline_seconds - args.eval_reserve_seconds:
743
+ control.should_training_stop = True
744
+ log_metric({"event": "time_guard_stop", "step": state.global_step,
745
+ "rows_cum": len(trainer.seen_rows)})
746
+
747
+ trainer.add_callback(Guards())
748
+ t0 = time.time()
749
+ trainer.train()
750
+ train_seconds = time.time() - t0
751
+ rows_trained = len(trainer.seen_rows)
752
+ steps_done = trainer.state.global_step
753
+ log_metric({"event": "train_done", "pool": args.pool,
754
+ "seconds": round(train_seconds, 1), "steps": steps_done,
755
+ "rows_covered": rows_trained, "n_rows_selected": len(parsed),
756
+ "n_pairs_processed": trainer.pairs_processed, "n_truncated_pairs": trainer.truncated_pairs,
757
+ "note": "rows_covered counts unique decisions actually iterated; "
758
+ "no full-epoch guarantee"})
759
+
760
+ # One fixed epoch: final weights are the selected checkpoint.
761
+ # Validation is monitored without rewinding to a partially trained checkpoint.
762
+ model_dir = os.path.join(save_dir, "model")
763
+ model.save_pretrained(model_dir)
764
+ tokenizer.save_pretrained(model_dir)
765
+ with open(os.path.join(save_dir, "training_coverage.json"), "w") as f:
766
+ json.dump({"target": len(parsed), "rows_seen": rows_trained,
767
+ "complete": rows_trained == len(parsed),
768
+ "steps": steps_done, "selection": "final fixed-epoch checkpoint",
769
+ "seen_row_ids": sorted(parsed[i]["row_id"] for i in trainer.seen_rows)}, f)
770
+ if args.push:
771
+ from huggingface_hub import HfApi
772
+ api = HfApi()
773
+ api.upload_folder(folder_path=model_dir, repo_id=args.hub_model_id, repo_type="model",
774
+ commit_message="Persist final prototype before evaluation")
775
+ api.upload_file(path_or_fileobj=os.path.join(save_dir, "training_coverage.json"),
776
+ path_in_repo="training_coverage.json", repo_id=args.hub_model_id, repo_type="model")
777
+ log_metric({"event": "checkpoint_saved_before_evaluation", "rows_covered": rows_trained})
778
+
779
+ results = {"model": "ModernJEV-Decide-Preview",
780
+ "dataset": {"id": DS_ID, "revision": DS_REV},
781
+ "base": {"id": BASE_ID, "revision": BASE_REV},
782
+ "train_pool": args.pool, "max_len": MAX_LEN, "seed": SEED,
783
+ "n_rows_selected": len(parsed), "n_a": n_a, "n_t": n_t,
784
+ "selected_row_ids_sha256": ids_hash,
785
+ "rows_covered": rows_trained, "steps": steps_done,
786
+ "train_seconds": round(train_seconds, 1),
787
+ "lr": args.lr, "token_budget": args.token_budget,
788
+ "grad_accum": args.grad_accum, "validation_monitor": best,
789
+ "input_preparation": "lazy per batch, no upfront map",
790
+ "training_input_stats": {"pairs": trainer.pairs_processed, "truncated": trainer.truncated_pairs, "at_max": trainer.pairs_at_max},
791
+ "checkpoint_selection": "final fixed-epoch checkpoint", "complete_training_coverage": rows_trained == len(parsed),
792
+ "gpu": gpu_name, "torch": torch.__version__, "attention": ATTN_IMPL,
793
+ "transformers": transformers.__version__}
794
+
795
+ val_pd_full, _ = tokenize_eval(val_full, tokenizer)
796
+ results["val_full"] = {k: v for k, v in
797
+ evaluate_rows(model, val_pd_full, val_full, device).items()
798
+ if k != "pred_labels"}
799
+ log_metric({"event": "val_full", **results["val_full"]})
800
+
801
+ test_focus, skipped_te = parse_eval_rows(ds["test"].filter(
802
+ lambda r: r["task_family"] in FOCUS, num_proc=8))
803
+ log_metric({"event": "test_prep", "n_rows": len(test_focus),
804
+ "skipped": skipped_te})
805
+ te_pd, _ = tokenize_eval(test_focus, tokenizer)
806
+ te = evaluate_rows(model, te_pd, test_focus, device)
807
+ results["test"] = {k: v for k, v in te.items() if k != "pred_labels"}
808
+ log_metric({"event": "final_test", **results["test"]})
809
+
810
+ inv_rows = stratified_subset(test_focus, 300, seed=SEED)
811
+ results["shuffled_invariance"] = shuffled_invariance(
812
+ model, tokenizer, inv_rows, device)
813
+ log_metric({"event": "shuffled_invariance", **results["shuffled_invariance"]})
814
+
815
+ results["baselines"] = build_baseline_results(parsed, val_full, test_focus)
816
+ log_metric({"event": "baselines", **results["baselines"]})
817
+
818
+ torch.manual_seed(SEED)
819
+ base_model = load_model().to(device)
820
+ results["baseline_untrained_head"] = {
821
+ "val": {k: v for k, v in evaluate_rows(
822
+ base_model, val_pd_full, val_full, device).items() if k != "pred_labels"},
823
+ "test": {k: v for k, v in evaluate_rows(
824
+ base_model, te_pd, test_focus, device).items() if k != "pred_labels"}}
825
+ log_metric({"event": "baseline_untrained_head",
826
+ **results["baseline_untrained_head"]})
827
+ del base_model
828
+ torch.cuda.empty_cache()
829
+
830
+ results["latency_gpu"] = gpu_latency(model, te_pd, device)
831
+ log_metric({"event": "latency_gpu", **results["latency_gpu"]})
832
+
833
+ results["probe"] = {"omitted": True, "reason": "Budget reserved for full prototype coverage and held-out evaluation"}
834
+
835
+ model_dir = os.path.join(save_dir, "model")
836
+ with open(os.path.join(save_dir, "results.json"), "w") as f:
837
+ json.dump(results, f, indent=2, default=str)
838
+ save_metrics(os.path.join(save_dir, "metrics.jsonl"))
839
+ with gzip.open(os.path.join(save_dir, "selected_row_ids.json.gz"), "wt") as f:
840
+ json.dump({"sha256": ids_hash, "n": len(parsed), "n_a": n_a, "n_t": n_t,
841
+ "pool": args.pool, "rows_covered": rows_trained,
842
+ "row_ids": [r["row_id"] for r in parsed]}, f)
843
+ here = os.path.dirname(os.path.abspath(__file__))
844
+ if os.path.exists(os.path.join(here, "predict.py")):
845
+ import shutil
846
+ shutil.copy(os.path.join(here, "predict.py"),
847
+ os.path.join(save_dir, "predict.py"))
848
+ if args.push:
849
+ from huggingface_hub import HfApi
850
+ api = HfApi()
851
+ assert api.model_info(args.hub_model_id).private, "Private model required"
852
+ for name in ["results.json", "metrics.jsonl", "selected_row_ids.json.gz",
853
+ "predict.py"]:
854
+ api.upload_file(path_or_fileobj=os.path.join(save_dir, name),
855
+ repo_id=args.hub_model_id, repo_type="model",
856
+ path_in_repo=name)
857
+ log_metric({"event": "persisted", "save_dir": save_dir, "pushed": args.push})
858
+ print("PROTOTYPE_DONE", flush=True)
859
+
860
+
861
+ if __name__ == "__main__":
862
+ main()
863
+ ===== END FILE: recipe/train.py =====
864
+
865
+ ===== FILE: predict.py =====
866
+ """Typed-choice inference for ModernJEV-Decide-Preview.
867
+ The encoder scores each declared (state, candidate) pair; it does not generate text.
868
+ """
869
+ import json
870
+ import os
871
+ import torch
872
+ from transformers import AutoModelForSequenceClassification, AutoTokenizer
873
+
874
+ MAX_LEN = 4096
875
+ MODEL_ID = "OpenMed/ModernJEV-Decide-Preview"
876
+
877
+ def serialize_state(state):
878
+ if not isinstance(state, dict):
879
+ raise TypeError("state must be a dict containing conversation and available_tools")
880
+ conv = state.get("conversation") or []
881
+ policy = state.get("policy")
882
+ first = conv[0] if conv else None
883
+ duplicate = (policy is not None and isinstance(first, dict)
884
+ and first.get("role") == "system" and first.get("content") == policy)
885
+ compact = {"available_tools": state.get("available_tools") or [], "conversation": conv}
886
+ if policy is not None and not duplicate:
887
+ compact["policy"] = policy
888
+ return json.dumps(compact, ensure_ascii=False)
889
+
890
+ def normalize_criteria(criteria):
891
+ if isinstance(criteria, list):
892
+ if not criteria or any(not isinstance(v, str) or not v.strip() for v in criteria):
893
+ raise ValueError("Answer list must contain nonempty strings")
894
+ if len(set(criteria)) != len(criteria):
895
+ raise ValueError("Answer list must contain unique strings")
896
+ criteria = {v: v for v in criteria}
897
+ if not isinstance(criteria, dict) or not criteria:
898
+ raise ValueError("criteria must be a nonempty mapping or list of unique answer strings")
899
+ if any(not isinstance(k, str) or not k.strip() or not isinstance(v, str)
900
+ for k, v in criteria.items()):
901
+ raise TypeError("Choice labels must be nonempty strings and descriptions must be strings")
902
+ return criteria
903
+
904
+ class DecisionModel:
905
+ def __init__(self, model_path=MODEL_ID, device=None, revision=None):
906
+ self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
907
+ self.tokenizer = AutoTokenizer.from_pretrained(model_path, revision=revision)
908
+ self.model = AutoModelForSequenceClassification.from_pretrained(
909
+ model_path, revision=revision, attn_implementation="sdpa").to(self.device).eval()
910
+ if self.model.config.num_labels != 1:
911
+ raise ValueError("Expected a trained scalar candidate-scoring head")
912
+
913
+ @torch.inference_mode()
914
+ def decide(self, *, state, question, criteria, candidate_batch_size=8):
915
+ if not isinstance(question, str) or not question.strip():
916
+ raise ValueError("question must be nonempty text")
917
+ criteria = normalize_criteria(criteria)
918
+ if not isinstance(candidate_batch_size, int) or isinstance(candidate_batch_size, bool) or candidate_batch_size < 1:
919
+ raise ValueError("candidate_batch_size must be a positive integer")
920
+ keys = list(criteria)
921
+ text_a = question + "\n\nSTATE:\n" + serialize_state(state)
922
+ text_bs = [f"{k}: {criteria[k]}" for k in keys]
923
+ raw_a_length = len(self.tokenizer(text_a, add_special_tokens=False, verbose=False)["input_ids"])
924
+ special = self.tokenizer.num_special_tokens_to_add(pair=True)
925
+ raw_lengths = [raw_a_length + len(self.tokenizer(t, add_special_tokens=False)["input_ids"]) + special for t in text_bs]
926
+ scores = []
927
+ for begin in range(0, len(keys), candidate_batch_size):
928
+ ts = text_bs[begin:begin + candidate_batch_size]
929
+ encoded = self.tokenizer([text_a] * len(ts), ts, truncation="only_first",
930
+ max_length=MAX_LEN, padding=True, return_tensors="pt", verbose=False).to(self.device)
931
+ with torch.autocast(device_type=self.device.split(":")[0],
932
+ dtype=torch.bfloat16, enabled=self.device.startswith("cuda")):
933
+ logits = self.model(input_ids=encoded["input_ids"],
934
+ attention_mask=encoded["attention_mask"]).logits.squeeze(-1)
935
+ scores.extend(logits.float().cpu().tolist())
936
+ probabilities = torch.softmax(torch.tensor(scores), dim=0).tolist()
937
+ order = sorted(range(len(keys)), key=lambda i: (-scores[i], keys[i]))
938
+ return {"predicted_label": keys[order[0]], "allowed_choices": keys,
939
+ "candidates": [{"label": keys[i], "score": probabilities[i], "raw_score": scores[i],
940
+ "rank": rank + 1} for rank,i in enumerate(order)],
941
+ "truncated": any(n > MAX_LEN for n in raw_lengths),
942
+ "max_sequence_length": MAX_LEN,
943
+ "note": "Scores rank this supplied choice set; they are not calibrated confidence."}
944
+
945
+ _default = None
946
+ def predict_typed(question_text, state_json, criteria_json):
947
+ global _default
948
+ if _default is None:
949
+ _default = DecisionModel(os.environ.get("MODEL_PATH", MODEL_ID))
950
+ state = json.loads(state_json) if isinstance(state_json, str) else state_json
951
+ criteria = json.loads(criteria_json) if isinstance(criteria_json, str) else criteria_json
952
+ return _default.decide(state=state, question=question_text, criteria=criteria)
953
+
954
+ def predict_batch(rows):
955
+ return [predict_typed(r["question_text"], r["state_json"], r["criteria_json"]) for r in rows]
956
+
957
+
958
+ ===== END FILE: predict.py =====
959
+
960
+ ===== FILE: runtime-versions.json =====
961
+ {
962
+ "torch": "2.12.0+cu126",
963
+ "transformers": "5.17.0",
964
+ "datasets": "5.0.1",
965
+ "accelerate": "1.15.0",
966
+ "huggingface-hub": "1.33.0",
967
+ "kernels": "0.16.0"
968
+ }
969
+
970
+ ===== END FILE: runtime-versions.json =====
971
+
972
+ ===== FILE: evaluate_open_labels.py =====
973
+ """Supplemental frozen-checkpoint evaluation; no training or Hub writes.
974
+
975
+ Run --audit-only first. After final checkpoint is saved, supply --model and,
976
+ for a Hub model, --revision. CPU is default; JSONL predictions resume safely.
977
+ """
978
+ import argparse
979
+ import hashlib
980
+ import importlib.util
981
+ import json
982
+ import statistics
983
+ from collections import Counter
984
+ from pathlib import Path
985
+
986
+ ROOT = Path(__file__).resolve().parent
987
+ DATA_REV = 'f2fb14e4ec977c420f376c08785664cd38763d7e'
988
+ CACHE = Path('/private/tmp/atd-ml-intern-hub-cache/MaziyarPanahi___agent_tool_decisions-180_k/default/0.0.0') / DATA_REV
989
+ FAMILIES = ('agent_next_action_type', 'tool_selection', 'when_to_call_tool')
990
+ EXPECTED = {FAMILIES[0]: (1158, 616), FAMILIES[1]: (542, 79), FAMILIES[2]: (3652, 1295)}
991
+
992
+
993
+ def write_json(path, value):
994
+ temp = path.with_suffix('.tmp')
995
+ temp.write_text(json.dumps(value, indent=2) + '\n')
996
+ temp.replace(path)
997
+
998
+
999
+ def load_test():
1000
+ from datasets import Dataset
1001
+ return Dataset.from_file(str(CACHE / 'agent_tool_decisions-180_k-test.arrow'))
1002
+
1003
+
1004
+ def audit(test):
1005
+ report = {'dataset_revision': DATA_REV, 'families': {}}
1006
+ for family in FAMILIES:
1007
+ rows = [r for r in test if r['task_family'] == family]
1008
+ labels = Counter(r['gold_label'] for r in rows)
1009
+ size, correct = EXPECTED[family]
1010
+ assert len(rows) == size and max(labels.values()) == correct
1011
+ candidates = [len(json.loads(r['criteria_json'])) for r in rows]
1012
+ assert all(r['gold_label'] in json.loads(r['criteria_json']) for r in rows)
1013
+ report['families'][family] = {
1014
+ 'n': size, 'majority_correct': correct, 'majority_accuracy': correct / size,
1015
+ 'majority_labels': sorted(k for k,v in labels.items() if v == correct),
1016
+ 'candidate_count': {'median': statistics.median(candidates),
1017
+ 'min': min(candidates), 'max': max(candidates)},
1018
+ 'uniform_expected_accuracy': statistics.mean(1/n for n in candidates),
1019
+ 'reference_definition': 'Descriptive test-set constant-label majority; no model fitting',
1020
+ 'unseen_task_family': family == 'when_to_call_tool',
1021
+ }
1022
+ # Verify When2Call is absent from every training shard, not just the subset.
1023
+ from datasets import Dataset
1024
+ train_families = Counter()
1025
+ selected_ids = set((ROOT / 'selected-row-ids.txt').read_text().splitlines())
1026
+ train_gold = {f: Counter() for f in FAMILIES[:2]}
1027
+ for path in sorted(CACHE.glob('*-train-*.arrow')):
1028
+ columns = Dataset.from_file(str(path)).select_columns(['task_family','row_id','gold_label'])[:]
1029
+ for family, row_id, gold in zip(columns['task_family'], columns['row_id'], columns['gold_label']):
1030
+ train_families[family] += 1
1031
+ if row_id in selected_ids and family in train_gold:
1032
+ train_gold[family][gold] += 1
1033
+ assert sum(train_families.values()) == 171056
1034
+ assert train_families['when_to_call_tool'] == 0
1035
+ assert sum(sum(c.values()) for c in train_gold.values()) == 60000
1036
+ report['when2call_training_rows'] = 0
1037
+ for family, counts in train_gold.items():
1038
+ label = min(counts, key=lambda k: (-counts[k], k))
1039
+ rows = [r for r in test if r['task_family'] == family]
1040
+ report['families'][family]['training_fixed_majority'] = {
1041
+ 'label': label, 'correct': sum(r['gold_label'] == label for r in rows),
1042
+ 'n': len(rows),
1043
+ 'majority_label_absent_from_choices': sum(label not in json.loads(r['criteria_json']) for r in rows),
1044
+ }
1045
+ return report
1046
+
1047
+
1048
+ def open_labels(client):
1049
+ checks = []
1050
+ for n in (2, 3, 7, 20):
1051
+ descriptions = ['Escalate this damaged parcel to a human support agent.',
1052
+ 'Search the product catalog for a new item.']
1053
+ descriptions += [f'Route to unrelated department number {i} for a different request.' for i in range(n-2)]
1054
+ criteria = {f'fresh_choice_{n}_{i}': d for i,d in enumerate(descriptions)}
1055
+ state = {'policy': 'Escalate damaged parcels to a human support agent.',
1056
+ 'conversation': [{'role':'user','content':'My parcel arrived broken. I need help.'}]}
1057
+ question = 'Which of these custom workflow branches should handle this request?'
1058
+ outputs = []
1059
+ variants = (criteria, dict(reversed(list(criteria.items()))),
1060
+ {f'opaque_{n}_{i}': d for i,d in enumerate(descriptions)}, descriptions)
1061
+ for variant in variants:
1062
+ answer = client.decide(state=state, question=question, criteria=variant, candidate_batch_size=4)
1063
+ allowed = list(variant)
1064
+ assert answer['predicted_label'] in allowed
1065
+ assert sorted(c['label'] for c in answer['candidates']) == sorted(allowed)
1066
+ assert len(answer['candidates']) == n
1067
+ outputs.append(answer)
1068
+ original_desc = criteria[outputs[0]['predicted_label']]
1069
+ renamed_desc = variants[2][outputs[2]['predicted_label']]
1070
+ checks.append({'answer_count': n, 'interface_pass': True,
1071
+ 'reorder_same_label': outputs[0]['predicted_label'] == outputs[1]['predicted_label'],
1072
+ 'rename_same_description': original_desc == renamed_desc,
1073
+ 'selected_expected_description': original_desc == descriptions[0],
1074
+ 'outputs': outputs})
1075
+ return {'scalar_head': client.model.config.num_labels,
1076
+ 'interpretation': 'Interface checks and illustrative decisions, not a clinical benchmark or general accuracy estimate.',
1077
+ 'checks': checks}
1078
+
1079
+
1080
+ def main():
1081
+ p = argparse.ArgumentParser()
1082
+ p.add_argument('--audit-only', action='store_true')
1083
+ p.add_argument('--interface-only', action='store_true')
1084
+ p.add_argument('--model')
1085
+ p.add_argument('--revision')
1086
+ p.add_argument('--device', default='cpu')
1087
+ p.add_argument('--output-dir', type=Path, default=ROOT / 'supplemental-evaluation')
1088
+ args = p.parse_args()
1089
+ out = args.output_dir; out.mkdir(parents=True, exist_ok=True)
1090
+ test = load_test(); baseline = audit(test)
1091
+ write_json(out/'per-task-baselines.json', baseline)
1092
+ print(json.dumps(baseline), flush=True)
1093
+ if args.audit_only: return
1094
+ if not args.model: p.error('--model is required for checkpoint evaluation')
1095
+ local = Path(args.model).is_dir()
1096
+ if not local and not args.revision: p.error('Pin --revision for a Hub model')
1097
+ if local:
1098
+ weights = sorted(Path(args.model).glob('*.safetensors'))
1099
+ assert weights, 'Local checkpoint must have saved safetensors'
1100
+ h = hashlib.sha256()
1101
+ for path in weights:
1102
+ with path.open('rb') as f:
1103
+ for chunk in iter(lambda:f.read(8*1024*1024), b''): h.update(chunk)
1104
+ identity = h.hexdigest()
1105
+ else: identity = args.revision
1106
+ manifest = {'model': args.model, 'checkpoint_identity': identity, 'dataset_revision': DATA_REV}
1107
+ if (out/'manifest.json').exists(): assert json.loads((out/'manifest.json').read_text()) == manifest
1108
+ write_json(out/'manifest.json', manifest)
1109
+ spec = importlib.util.spec_from_file_location('inference', ROOT/'predict_open_labels.py')
1110
+ helper = importlib.util.module_from_spec(spec); spec.loader.exec_module(helper)
1111
+ import torch
1112
+ torch.set_num_threads(4)
1113
+ client = helper.DecisionModel(args.model, device=args.device, revision=args.revision)
1114
+ assert client.model.config.num_labels == 1
1115
+ write_json(out/'open-label-checks.json', open_labels(client))
1116
+ if args.interface_only: return
1117
+ rows = [r for r in test if r['task_family'] == 'when_to_call_tool']
1118
+ path = out/'when2call-predictions.jsonl'; previous = {}
1119
+ if path.exists():
1120
+ for line in path.read_text().splitlines():
1121
+ if line: item = json.loads(line); previous[item['row_id']] = item
1122
+ assert set(previous).issubset({r['row_id'] for r in rows})
1123
+ with path.open('a') as f:
1124
+ for row in rows:
1125
+ if row['row_id'] in previous: continue
1126
+ pred = client.decide(state=json.loads(row['state_json']), question=row['question_text'],
1127
+ criteria=json.loads(row['criteria_json']), candidate_batch_size=4)
1128
+ rec = {'row_id': row['row_id'], 'gold': row['gold_label'], **pred}
1129
+ f.write(json.dumps(rec)+'\n'); f.flush(); previous[row['row_id']] = rec
1130
+ if len(previous) % 100 == 0: print(f'When2Call {len(previous)}/{len(rows)}', flush=True)
1131
+ correct = sum(previous[r['row_id']]['predicted_label'] == r['gold_label'] for r in rows)
1132
+ n = len(rows); reference = baseline['families']['when_to_call_tool']['majority_accuracy']
1133
+ result = {'task': 'when_to_call_tool', 'correct': correct, 'n': n, 'accuracy': correct/n,
1134
+ 'majority_reference': reference, 'lift_percentage_points': 100*(correct/n-reference),
1135
+ 'unseen_task_family': True, 'checkpoint_identity': identity,
1136
+ 'truncated_rows': sum(x['truncated'] for x in previous.values())}
1137
+ write_json(out/'when2call-results.json', result)
1138
+ print(json.dumps(result), flush=True)
1139
+
1140
+
1141
+ if __name__ == '__main__': main()
1142
+
1143
+ ===== END FILE: evaluate_open_labels.py =====
workflow/recipe-manifest.json ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "file": "recipe/train.py",
4
+ "sha256": "37774a1a9a7456cf25b2f33b7272a9817be89dc235577e4f63d507ad7128dd8a"
5
+ },
6
+ {
7
+ "file": "predict.py",
8
+ "sha256": "d77ea48b1fadc19b8e77404d6066031d571d01d7c2b6120526499b069bea2fd5"
9
+ },
10
+ {
11
+ "file": "runtime-versions.json",
12
+ "sha256": "a506eb8faa5a457b04834d03160d8054a16c3dba086e2d6d570d8e1dbc360163"
13
+ },
14
+ {
15
+ "file": "evaluate_open_labels.py",
16
+ "sha256": "da0db4b3114ec7d1d0fbfddb54668ef74f9ea940574d9c7f05f679af4ab8151e",
17
+ "source": "previously completed local evaluation runner"
18
+ }
19
+ ]