ModernJEV-Decide-Preview / workflow /recipe-bundle.txt
MaziyarPanahi's picture
ModernJEV-Decide-Preview: model, results, compute costs and ML Intern workflow
d5b398e
Raw History Blame Contribute Delete
55.6 kB
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.
===== FILE: recipe/train.py =====
"""ModernJEV-Decide-Preview — training, evaluation, persistence.
Dataset : MaziyarPanahi/AgentToolDecisions-180K @ f2fb14e4ec977c420f376c08785664cd38763d7e
Base : answerdotai/ModernBERT-base @ 8949b909ec900327062f0ebf497f51aef5e6f0c8
Scope : task_family in {agent_next_action_type, tool_selection} ONLY (choice primitive).
Objective: shared-encoder candidate scalar scorer; per-row_id softmax cross-entropy
over the row's DECLARED candidates. group_id is the EPISODE, not the decision —
grouping is ALWAYS by row_id (asserted). Candidates enter the softmax as text
(label + criterion), so variable tool names need no fixed head and no candidate
index is exposed to the model. Inputs contain NO gold_label / gold_json /
gold_score / label_source / source metadata.
Training pool (--pool):
full — every declared candidate of the row joins the softmax.
sampled4 — gold + up to 3 declared negatives, deterministic per (SEED, epoch,
row_id). This is ordinary sampled-choice softmax within the same loss;
EVALUATION ALWAYS RANKS ALL DECLARED CANDIDATES regardless of --pool.
The pilot benchmarks both objectives and the launch report states which one ran.
Modes:
pilot — benchmarks full vs sampled4 throughput/memory, GPU latency, save check.
prototype — time-guarded training on the exact stratified subset, interval monitoring
on a stratified subset of the OFFICIAL validation split (full-val eval
recorded for the final fixed-epoch checkpoint), ONE final test evaluation on all in-scope
test rows, baselines (uniform / train-frequency / untrained ModernBERT
head), shuffled-order label invariance, optional frozen-backbone probe,
and persistence to --save_dir (default /output/modernjev). No Hub push
unless --push is passed explicitly.
No Space is created anywhere; metrics are metrics.jsonl + stdout only.
"""
import argparse
import gzip
import hashlib
import json
import os
import random
import time
import numpy as np
import torch
import torch.nn.functional as F
from datasets import Dataset, load_dataset
from torch.utils.data import DataLoader, Dataset as TorchDataset, Sampler
from transformers import (AutoModelForSequenceClassification, AutoTokenizer,
Trainer, TrainerCallback, TrainingArguments)
DS_ID = "MaziyarPanahi/AgentToolDecisions-180K"
DS_REV = "f2fb14e4ec977c420f376c08785664cd38763d7e"
BASE_ID = "answerdotai/ModernBERT-base"
BASE_REV = "8949b909ec900327062f0ebf497f51aef5e6f0c8"
FOCUS = ("agent_next_action_type", "tool_selection")
SEED = 42
MAX_LEN = 4096
ATTN_IMPL = os.environ.get("MODERNJEV_ATTN", "sdpa")
EXPECTED = {"train": 171056, "validation": 2713, "test": 6231}
EXPECTED_FOCUS_TRAIN = 112973
METRICS = []
def log_metric(d):
d["ts"] = round(time.time(), 1)
METRICS.append(d)
print("METRIC " + json.dumps(d, default=str), flush=True)
def save_metrics(path):
with open(path, "w") as f:
for d in METRICS:
f.write(json.dumps(d, default=str) + "\n")
def serialize_state(row):
"""Compact the state. Drops the policy key ONLY when it exactly equals the
first system message (audited dedupe rule)."""
state = json.loads(row["state_json"])
conv = state.get("conversation") or []
policy = state.get("policy")
first = conv[0] if conv else None
dup = (policy is not None and isinstance(first, dict)
and first.get("role") == "system" and first.get("content") == policy)
compact = {"available_tools": state.get("available_tools") or [],
"conversation": conv}
if policy is not None and not dup:
compact["policy"] = policy
return json.dumps(compact, ensure_ascii=False)
def parse_row(r):
criteria = json.loads(r["criteria_json"])
keys = list(criteria.keys())
gold = r["gold_label"]
return {"row_id": r["row_id"], "group_id": r["group_id"],
"family": r["task_family"],
"text_a": r["question_text"] + "\n\nSTATE:\n" + serialize_state(r),
"cand_keys": keys,
"cand_texts": [f"{k}: {criteria[k]}" for k in keys],
"gold_idx": keys.index(gold) if gold in keys else None}
def load_and_prepare(n_rows, seed=SEED):
t0 = time.time()
ds = load_dataset(DS_ID, revision=DS_REV)
for split, n in EXPECTED.items():
assert len(ds[split]) == n, f"{split}: {len(ds[split])} != {n}"
train = ds["train"].filter(lambda r: r["task_family"] in FOCUS, num_proc=8)
assert len(train) == EXPECTED_FOCUS_TRAIN, f"{len(train)} != {EXPECTED_FOCUS_TRAIN}"
metadata = train.select_columns(["row_id", "task_family"])[:]
fam_idx = {f: [] for f in FOCUS}
for i, fam in enumerate(metadata["task_family"]):
fam_idx[fam].append(i)
n_a = n_rows * len(fam_idx["agent_next_action_type"]) // len(train)
n_t = n_rows - n_a
rng = random.Random(seed)
picked = []
for fam, k in (("agent_next_action_type", n_a), ("tool_selection", n_t)):
all_ids = metadata["row_id"]
idxs = sorted(fam_idx[fam], key=lambda i: all_ids[i])
picked += rng.sample(idxs, k)
assert len(picked) == n_rows
picked.sort()
fields = ["row_id", "group_id", "task_family", "question_text", "state_json", "criteria_json", "gold_label"]
selected = train.select(picked).select_columns(fields)[:]
parsed = [parse_row(dict(zip(fields, values))) for values in zip(*(selected[f] for f in fields))]
for r in parsed:
assert r["gold_idx"] is not None, f"gold missing in {r['row_id']}"
fam_counts = {f: sum(1 for r in parsed if r["family"] == f) for f in FOCUS}
gold_classes = {f: {} for f in FOCUS}
for r in parsed:
g = r["cand_keys"][r["gold_idx"]]
gold_classes[r["family"]][g] = gold_classes[r["family"]].get(g, 0) + 1
ids_hash = hashlib.sha256(
"\n".join(sorted(r["row_id"] for r in parsed)).encode()).hexdigest()
log_metric({"event": "data_prep", "n_rows": len(parsed), "n_a": n_a, "n_t": n_t,
"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},
"selected_row_ids_sha256": ids_hash,
"seconds": round(time.time() - t0, 1)})
return parsed, ds, ids_hash, (n_a, n_t)
class LazyPairDataset:
"""Tokenize only requested pairs. Metadata stays small; no up-front Dataset.map."""
def __init__(self, parsed, tokenizer, tag, max_len, training_pool):
self.parsed, self.tokenizer, self.tag, self.max_len = parsed, tokenizer, tag, max_len
self.refs = []
for row in range(len(parsed)):
candidates, _ = pool_for(parsed, row, 1, training_pool)
self.refs.extend((row, candidate) for candidate in candidates)
self.metadata = Dataset.from_dict({
"row_idx": [r for r, _ in self.refs],
"cand_idx": [c for _, c in self.refs],
# Conservative packing bound; actual batch padding uses actual token lengths.
"length": [max_len] * len(self.refs),
})
self.row_sortlen = np.array([
min(max_len, max(1, len(r["text_a"]) // 4)) for r in parsed], dtype=np.int64)
log_metric({"event": "lazy_pairs_ready", "tag": tag,
"n_rows": len(parsed), "n_pairs": len(self.refs),
"max_len": max_len, "mapped_pairs": 0})
def __len__(self):
return len(self.refs)
def select_columns(self, columns):
return self.metadata.select_columns(columns)
def __getitem__(self, indices):
scalar = isinstance(indices, (int, np.integer))
if scalar:
indices = [int(indices)]
elif isinstance(indices, slice):
indices = list(range(*indices.indices(len(self))))
else:
indices = list(indices)
refs = [self.refs[int(i)] for i in indices]
encoded = self.tokenizer(
[self.parsed[r]["text_a"] for r, c in refs],
[self.parsed[r]["cand_texts"][c] for r, c in refs],
truncation="only_first", max_length=self.max_len,
padding=False, verbose=False)
result = {"row_idx": [r for r, c in refs],
"cand_idx": [c for r, c in refs],
"input_ids": encoded["input_ids"],
"attention_mask": encoded["attention_mask"],
"length": [len(ids) for ids in encoded["input_ids"]],
"truncated": [bool(e.overflowing) for e in encoded.encodings]}
return {k: v[0] for k, v in result.items()} if scalar else result
def tokenize_pairs(parsed, tokenizer, tag, max_len=MAX_LEN, num_proc=8, training_pool="full"):
return LazyPairDataset(parsed, tokenizer, tag, max_len, training_pool)
class RowIndexer:
def __init__(self, parsed, pair_ds):
meta = pair_ds.select_columns(["row_idx", "cand_idx", "length"]).with_format("numpy")[:]
assert np.all(np.diff(meta["row_idx"]) >= 0), "pairs must be row-major"
self.lookup = {}
self.row_maxlen = np.zeros(len(parsed), dtype=np.int64)
for flat, (row, cand, length) in enumerate(zip(meta["row_idx"], meta["cand_idx"], meta["length"])):
row, cand = int(row), int(cand)
assert (row, cand) not in self.lookup
self.lookup[(row, cand)] = flat
self.row_maxlen[row] = max(self.row_maxlen[row], int(length))
assert np.all(self.row_maxlen > 0)
self.row_sortlen = getattr(pair_ds, "row_sortlen", self.row_maxlen)
def flat(self, row, cand):
return self.lookup[(row, cand)]
def pool_for(parsed, row, epoch, mode):
"""Training pool for one row: (pool_cand_indices, gold_pos_in_pool).
full: all declared candidates. sampled4: gold + up to 3 declared negatives,
deterministic per (SEED, epoch, row_id)."""
r = parsed[row]
k = len(r["cand_keys"])
gold = r["gold_idx"]
if mode == "full" or k <= 4:
return list(range(k)), gold
rng = random.Random(f"{SEED}:{epoch}:{r['row_id']}")
others = sorted(set(range(k)) - {gold})
negs = rng.sample(others, min(3, len(others)))
pool = [gold] + sorted(negs)
rng.shuffle(pool)
return pool, pool.index(gold)
class TrainRefs(TorchDataset):
"""Position i -> (row, cand_actual) for the current epoch; gold_pos is
recomputed in collate from the pool definition (deterministic)."""
def __init__(self, sampler):
self.sampler = sampler
def __len__(self):
return len(self.sampler.refs)
def __getitem__(self, i):
return self.sampler.refs[i]
class WholeRowBatchSampler(Sampler):
"""Yields batches of positions into self.refs. Every row's pooled pairs stay
in one batch (padding-token budget checked at row boundaries)."""
def __init__(self, parsed, indexer, mode, token_budget, max_rows=64,
seed=SEED, shuffle=True, length_buckets=True):
self.parsed, self.indexer, self.mode = parsed, indexer, mode
self.token_budget, self.max_rows = token_budget, max_rows
self.rng = random.Random(seed)
self.shuffle = shuffle
self.length_buckets = length_buckets
self.epoch = 0
self.refs = []
self.rows_cum = 0
self.pairs_cum = 0
self.tokens_cum = 0
self.batches = self._build_epoch()
def _build_epoch(self):
self.epoch += 1
order = list(range(len(self.parsed)))
if self.shuffle:
self.rng.shuffle(order)
# Sort locally within randomized buckets to reduce padding, preserving all rows.
if self.length_buckets:
order = [row for begin in range(0, len(order), 256)
for row in sorted(order[begin:begin + 256], key=lambda r: int(self.indexer.row_sortlen[r]))]
refs, batches, cur = [], [], []
batch_maxlen, batch_rows = 0, 0
for r in order:
pool, _ = pool_for(self.parsed, r, self.epoch, self.mode)
row_len = int(self.indexer.row_maxlen[r])
new_maxlen = max(batch_maxlen, row_len)
if cur and (new_maxlen * (len(cur) + len(pool)) > self.token_budget
or batch_rows >= self.max_rows):
batches.append(cur)
cur, batch_maxlen, batch_rows = [], 0, 0
start = len(refs)
refs.extend((r, cand) for cand in pool)
cur.extend(range(start, start + len(pool)))
batch_maxlen = max(batch_maxlen, row_len)
batch_rows += 1
if cur:
batches.append(cur)
self.refs = refs
return batches
def __iter__(self):
return iter(self.batches)
def __len__(self):
return len(self.batches)
def make_collate(pair_ds, indexer, pad_id, parsed, mode):
def collate(items):
rows = torch.tensor([it[0] for it in items], dtype=torch.long)
by_row = {}
for row_index, candidate_index in items:
by_row.setdefault(row_index, []).append(candidate_index)
gold_positions = {}
for row_index, candidate_indices in by_row.items():
gold_index = parsed[row_index]["gold_idx"]
assert candidate_indices.count(gold_index) == 1, "Each row needs exactly one gold candidate"
gold_positions[row_index] = candidate_indices.index(gold_index)
golds = torch.tensor([gold_positions[it[0]] for it in items], dtype=torch.long)
flats = [indexer.flat(it[0], it[1]) for it in items]
rec = pair_ds[flats]
maxlen = max(rec["length"])
B = len(items)
ids = torch.full((B, maxlen), pad_id, dtype=torch.long)
att = torch.zeros((B, maxlen), dtype=torch.long)
for i, (ii, aa) in enumerate(zip(rec["input_ids"], rec["attention_mask"])):
n = len(ii)
ids[i, :n] = torch.tensor(ii, dtype=torch.long)
att[i, :n] = torch.tensor(aa, dtype=torch.long)
return {"input_ids": ids, "attention_mask": att, "row": rows, "gold": golds,
"_n_truncated": sum(rec["truncated"]), "_n_at_max": sum(n >= MAX_LEN for n in rec["length"])}
collate.epoch = 1
return collate
def grouped_ce(logits, rows, gold):
"""Per-row softmax CE. rows must be contiguous per row (asserted)."""
uniq_c = torch.unique_consecutive(rows)
uniq_all = torch.unique(rows)
assert len(uniq_c) == len(uniq_all), "row pairs not contiguous — grouping unsafe"
counts = torch.unique_consecutive(rows, return_counts=True)[1].tolist()
losses = []
ofs = 0
for c in counts:
seg = logits[ofs:ofs + c]
g = int(gold[ofs].item())
assert g < c, f"gold {g} out of pool size {c}"
losses.append(F.cross_entropy(seg.unsqueeze(0),
torch.tensor([g], device=seg.device)))
ofs += c
return torch.stack(losses).mean(), len(counts)
class GroupTrainer(Trainer):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.model_accepts_loss_kwargs = False
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
gold = inputs.pop("gold")
rows = inputs.pop("row")
if not hasattr(self, "seen_rows"):
self.seen_rows = set()
self.rows_processed = 0
self.pairs_processed = 0
self.tokens_processed = 0
self.truncated_pairs = 0
self.pairs_at_max = 0
self.truncated_pairs += int(inputs.pop("_n_truncated", 0))
self.pairs_at_max += int(inputs.pop("_n_at_max", 0))
batch_rows = torch.unique_consecutive(rows).detach().cpu().tolist()
self.seen_rows.update(batch_rows)
self.rows_processed += len(batch_rows)
self.pairs_processed += len(rows)
self.tokens_processed += int(inputs["attention_mask"].sum().item())
out = model(input_ids=inputs["input_ids"],
attention_mask=inputs["attention_mask"])
logits = out.logits.squeeze(-1).float()
loss, _ = grouped_ce(logits, rows, gold)
return (loss, out) if return_outputs else loss
def get_train_dataloader(self):
workers = self.args.dataloader_num_workers
return DataLoader(self._refs_ds, batch_sampler=self._sampler,
collate_fn=self._collate, num_workers=workers,
persistent_workers=workers > 0,
prefetch_factor=2 if workers else None,
pin_memory=self.args.device.type == "cuda")
class EvalPairs(TorchDataset):
def __init__(self, pair_ds):
self.pair_ds = pair_ds
def __len__(self):
return len(self.pair_ds)
def __getitem__(self, i):
return i
def make_eval_collate(pair_ds, pad_id):
def collate(idx_list):
rec = pair_ds[idx_list]
maxlen = max(rec["length"])
B = len(idx_list)
ids = torch.full((B, maxlen), pad_id, dtype=torch.long)
att = torch.zeros((B, maxlen), dtype=torch.long)
for i, (ii, aa) in enumerate(zip(rec["input_ids"], rec["attention_mask"])):
n = len(ii)
ids[i, :n] = torch.tensor(ii, dtype=torch.long)
att[i, :n] = torch.tensor(aa, dtype=torch.long)
return {"input_ids": ids, "attention_mask": att,
"row": torch.tensor(rec["row_idx"], dtype=torch.long),
"_n_truncated": sum(rec["truncated"])}
return collate
@torch.no_grad()
def evaluate_rows(model, pair_ds, parsed, device, batch_size=8):
"""Row-grouped accuracy over ALL declared candidates."""
model.eval()
scores = [[] for _ in range(len(parsed))]
n_truncated_pairs = 0
dl = DataLoader(EvalPairs(pair_ds), batch_size=batch_size,
collate_fn=make_eval_collate(pair_ds, model.config.pad_token_id), num_workers=0)
t0 = time.time()
for batch in dl:
n_truncated_pairs += int(batch["_n_truncated"])
with torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
out = model(input_ids=batch["input_ids"].to(device),
attention_mask=batch["attention_mask"].to(device))
lg = out.logits.squeeze(-1).float().cpu()
for j, r in enumerate(batch["row"].tolist()):
scores[r].append(lg[j].item())
correct, n = 0, 0
per_fam = {f: [0, 0] for f in FOCUS}
pred_labels = []
for r, row in enumerate(parsed):
if not scores[r]:
continue
p = min(range(len(scores[r])), key=lambda c: (-scores[r][c], row["cand_keys"][c]))
ok = int(p == row["gold_idx"])
correct += ok
n += 1
per_fam[row["family"]][0] += ok
per_fam[row["family"]][1] += 1
pred_labels.append(row["cand_keys"][p])
model.train()
return {"acc": correct / max(n, 1), "n": n,
"macro_accuracy": sum(per_fam[f][0] / max(per_fam[f][1], 1) for f in FOCUS) / len(FOCUS),
"per_family": {f: {"acc": per_fam[f][0] / max(per_fam[f][1], 1),
"n": per_fam[f][1]} for f in FOCUS},
"pred_labels": pred_labels,
"n_pairs": len(pair_ds), "n_truncated_pairs": n_truncated_pairs,
"seconds": round(time.time() - t0, 1)}
def stratified_subset(parsed, n, seed=SEED):
fam_rows = {f: sorted([r for r in parsed if r["family"] == f],
key=lambda r: r["row_id"]) for f in FOCUS}
total = sum(len(v) for v in fam_rows.values())
out = []
for f in FOCUS:
k = n * len(fam_rows[f]) // total
out += random.Random(seed).sample(fam_rows[f], k)
return out[:n]
def parse_eval_rows(split_ds):
parsed, skipped = [], 0
for r in split_ds:
row = parse_row(r)
if row["gold_idx"] is None:
skipped += 1
continue
parsed.append(row)
return parsed, skipped
def tokenize_eval(parsed, tokenizer):
pd = tokenize_pairs(parsed, tokenizer, tag="eval", num_proc=4)
idx = RowIndexer(parsed, pd)
return pd, idx
def build_baseline_results(parsed_tr, parsed_val, parsed_te):
"""Analytic uniform expected accuracy and allowed-choice training frequency."""
freq = {f: {} for f in FOCUS}
for row in parsed_tr:
g = row["cand_keys"][row["gold_idx"]]
fam_freq = freq[row["family"]]
fam_freq[g] = fam_freq.get(g, 0) + 1
res = {}
for name, parsed in (("val", parsed_val), ("test", parsed_te)):
methods = {"uniform_expected": {f: [0., 0] for f in FOCUS},
"train_frequency": {f: [0., 0] for f in FOCUS}}
for row in parsed:
fam = row["family"]
expected = 1. / len(row["cand_keys"])
best = min(row["cand_keys"], key=lambda key: (-freq[fam].get(key, 0), key))
matched = int(best == row["cand_keys"][row["gold_idx"]])
for method, value in (("uniform_expected", expected), ("train_frequency", matched)):
methods[method][fam][0] += value
methods[method][fam][1] += 1
res[name] = {}
for method, counts in methods.items():
per = {f: {"acc": c / n, "n": n} for f, (c, n) in counts.items()}
res[name][method] = {"acc": sum(c for c, _ in counts.values()) / len(parsed),
"n": len(parsed), "per_family": per,
"macro_accuracy": sum(v["acc"] for v in per.values()) / len(FOCUS)}
return res
@torch.no_grad()
def gpu_latency(model, pair_ds, device, n=100):
model.eval()
lat = []
for i in range(min(n + 5, len(pair_ds))):
rec = pair_ds[[i]]
ids = torch.tensor(rec["input_ids"][0], dtype=torch.long,
device=device).unsqueeze(0)
att = torch.tensor(rec["attention_mask"][0], dtype=torch.long,
device=device).unsqueeze(0)
if device == "cuda":
torch.cuda.synchronize()
t = time.time()
with torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
model(input_ids=ids, attention_mask=att)
if device == "cuda":
torch.cuda.synchronize()
if i >= 5:
lat.append(time.time() - t)
lat.sort()
return {"measurement": "one candidate forward pass, excludes tokenization", "warmup_pairs": 5,
"pair_ms_p50": round(1000 * lat[len(lat) // 2], 1),
"pair_ms_p95": round(1000 * lat[int(0.95 * len(lat))], 1)}
def shuffled_invariance(model, tokenizer, parsed_subset, device, k=5, batch_size=8):
"""Permute candidate order k times; the predicted LABEL must be identical."""
rng = random.Random(SEED + 1)
pd0, _ = tokenize_eval(parsed_subset, tokenizer)
base_labels = evaluate_rows(model, pd0, parsed_subset, device)["pred_labels"]
agree, total = 0, 0
for rep in range(k):
perm = []
for r in parsed_subset:
order = list(range(len(r["cand_texts"])))
rng.shuffle(order)
perm.append({**r, "cand_keys": [r["cand_keys"][j] for j in order],
"cand_texts": [r["cand_texts"][j] for j in order],
"gold_idx": order.index(r["gold_idx"])})
pd, _ = tokenize_eval(perm, tokenizer)
acc = evaluate_rows(model, pd, perm, device, batch_size)
agree += sum(1 for a, b in zip(base_labels, acc["pred_labels"]) if a == b)
total += len(base_labels)
return {"invariance_rate": agree / max(total, 1), "perms": k,
"n_rows": len(parsed_subset)}
def parse_args():
p = argparse.ArgumentParser()
p.add_argument("--mode", choices=["pilot", "prototype"], required=True)
p.add_argument("--pool", choices=["full", "sampled4"], default="sampled4",
help="training candidate pool (eval always uses all declared)")
p.add_argument("--n_rows", type=int, default=60000)
p.add_argument("--max_steps", type=int, default=25)
p.add_argument("--pilot_sampled_only", action="store_true")
p.add_argument("--pilot_length_buckets", action="store_true")
p.add_argument("--token_budget", type=int, default=32768)
p.add_argument("--grad_accum", type=int, default=2)
p.add_argument("--lr", type=float, default=2e-5)
p.add_argument("--eval_steps", type=int, default=400)
p.add_argument("--deadline_seconds", type=int, default=6600)
p.add_argument("--eval_reserve_seconds", type=int, default=1200)
p.add_argument("--save_dir", default="/output/modernjev")
p.add_argument("--push", action="store_true", default=False)
p.add_argument("--hub_model_id", default="OpenMed/ModernJEV-Decide-Preview")
return p.parse_args()
def prototype_training_args(args, save_dir, device):
return TrainingArguments(
output_dir=os.path.join(save_dir, "ckpt"), per_device_train_batch_size=1,
gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr,
bf16=device == "cuda", use_cpu=device == "cpu",
num_train_epochs=1.0, logging_steps=25, logging_first_step=True,
save_strategy="no", eval_strategy="no", report_to="none",
seed=SEED, remove_unused_columns=False,
lr_scheduler_type="linear", warmup_steps=0.03,
dataloader_num_workers=2 if device == "cuda" else 0)
def load_model():
return AutoModelForSequenceClassification.from_pretrained(
BASE_ID, revision=BASE_REV, num_labels=1, attn_implementation=ATTN_IMPL)
def main():
args = parse_args()
t_start = time.time()
torch.manual_seed(SEED)
random.seed(SEED)
device = "cuda" if torch.cuda.is_available() else "cpu"
gpu_name = torch.cuda.get_device_name(0) if device == "cuda" else "cpu"
import transformers
log_metric({"event": "env", "mode": args.mode, "device": device, "gpu": gpu_name,
"torch": torch.__version__, "transformers": transformers.__version__, "attention": ATTN_IMPL})
assert device == "cuda", "GPU required"
# Validate the complete prototype argument branch before any dataset work.
validated_training_args = prototype_training_args(args, args.save_dir, device) if args.mode == "prototype" else None
log_metric({"event": "training_api_validated", "mode": args.mode})
tokenizer = AutoTokenizer.from_pretrained(BASE_ID, revision=BASE_REV)
pad_id = tokenizer.pad_token_id
assert pad_id is not None, "tokenizer has no pad token"
n_train_rows = 500 if args.mode == "pilot" else args.n_rows
parsed, ds, ids_hash, (n_a, n_t) = load_and_prepare(n_train_rows)
if args.mode == "prototype":
assert ids_hash == "c5b306b0471ba104161051a7241b3bce4b69e1d959ff3e34fb06ef8eb4b077d9", "Selected subset differs from authorized frozen manifest"
os.makedirs(args.save_dir, exist_ok=True)
pair_ds = tokenize_pairs(parsed, tokenizer, tag="train",
training_pool=args.pool if args.mode == "prototype" else "full")
indexer = RowIndexer(parsed, pair_ds)
collate = make_collate(pair_ds, indexer, pad_id, parsed, args.pool)
if args.mode == "pilot":
results = {"phases": {}}
pilot_modes = ("sampled4",) if args.pilot_sampled_only else ("full", "sampled4")
for phase_i, mode in enumerate(pilot_modes):
torch.manual_seed(SEED)
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)
model = load_model().to(device)
sampler = WholeRowBatchSampler(parsed, indexer, mode, args.token_budget, length_buckets=args.pilot_length_buckets)
targs = TrainingArguments(
output_dir="/tmp/pilot_" + mode, per_device_train_batch_size=1,
gradient_accumulation_steps=1, learning_rate=args.lr, bf16=True,
max_steps=phase_steps, logging_steps=3, save_strategy="no",
eval_strategy="no", report_to="none", seed=SEED,
remove_unused_columns=False)
trainer = GroupTrainer(model=model, args=targs,
train_dataset=TrainRefs(sampler))
trainer._sampler = sampler
trainer._collate = make_collate(pair_ds, indexer, pad_id, parsed, mode)
trainer._refs_ds = TrainRefs(sampler)
torch.cuda.reset_peak_memory_stats()
t0 = time.time()
trainer.train()
dt = time.time() - t0
results["phases"][mode] = {
"steps": trainer.state.global_step, "seconds": round(dt, 2),
"steps_per_s": round(trainer.state.global_step / dt, 3),
"rows_seen": len(trainer.seen_rows),
"rows_per_s": round(trainer.rows_processed / dt, 2),
"pairs_seen": trainer.pairs_processed,
"pairs_per_s": round(trainer.pairs_processed / dt, 1),
"tokens_per_s": round(trainer.tokens_processed / dt),
"max_mem_gb": round(torch.cuda.max_memory_allocated() / 1e9, 2)}
log_metric({"event": "pilot_phase", "pool": mode,
**results["phases"][mode]})
checkpoint = os.path.join(args.save_dir, "checkpoint_" + mode)
model.save_pretrained(checkpoint)
tokenizer.save_pretrained(checkpoint)
del trainer
if mode == "full":
del model
torch.cuda.empty_cache()
results["latency"] = gpu_latency(model, pair_ds, device)
log_metric({"event": "latency_gpu", **results["latency"]})
results["env"] = {"gpu": gpu_name, "torch": torch.__version__}
os.makedirs(args.save_dir, exist_ok=True)
with open(os.path.join(args.save_dir, "pilot_results.json"), "w") as f:
json.dump(results, f, indent=2, default=str)
save_metrics(os.path.join(args.save_dir, "pilot_metrics.jsonl"))
reload_model = AutoModelForSequenceClassification.from_pretrained(os.path.join(args.save_dir, "checkpoint_sampled4"), attn_implementation=ATTN_IMPL).to(device)
rec = pair_ds[[0]]
ids = torch.tensor(rec["input_ids"], device=device)
att = torch.tensor(rec["attention_mask"], device=device)
model.eval(); reload_model.eval()
saved_weights = reload_model.state_dict()
for name, value in model.state_dict().items():
assert torch.equal(value, saved_weights[name]), f"Checkpoint changed parameter: {name}"
with torch.no_grad(), torch.autocast(device_type=device, dtype=torch.bfloat16, enabled=device == "cuda"):
expected_logits = model(input_ids=ids, attention_mask=att).logits.float()
actual_logits = reload_model(input_ids=ids, attention_mask=att).logits.float()
assert torch.isfinite(actual_logits).all()
assert torch.allclose(expected_logits, actual_logits, atol=1e-4, rtol=1e-4), "Same-precision reload mismatch"
results["checkpoint_parameter_equality"] = True
results["checkpoint_reload_verified"] = True
with open(os.path.join(args.save_dir, "pilot_results.json"), "w") as f:
json.dump(results, f, indent=2)
if args.push:
from huggingface_hub import HfApi
api = HfApi()
for filename in ("pilot_results.json", "pilot_metrics.jsonl"):
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")
print("PILOT_DONE", flush=True)
return
# ---------------- PROTOTYPE ----------------
save_dir = args.save_dir
os.makedirs(save_dir, exist_ok=True)
log_metric({"event": "prototype_start", "pool": args.pool,
"n_rows_selected": len(parsed), "n_a": n_a, "n_t": n_t})
val_full, _ = parse_eval_rows(ds["validation"].filter(
lambda r: r["task_family"] in FOCUS, num_proc=8))
val_sub = stratified_subset(val_full, 600, seed=SEED)
val_pd, _ = tokenize_eval(val_sub, tokenizer)
model = load_model().to(device)
sampler = WholeRowBatchSampler(parsed, indexer, args.pool, args.token_budget)
targs = validated_training_args
trainer = GroupTrainer(model=model, args=targs, train_dataset=TrainRefs(sampler))
trainer._sampler = sampler
trainer._collate = collate
trainer._refs_ds = TrainRefs(sampler)
best = {"acc": -1.0, "step": -1}
def run_val(step):
acc = evaluate_rows(model, val_pd, val_sub, device)
acc.pop("pred_labels")
log_metric({"event": "val_acc", "step": step,
"n_rows": len(val_sub), **acc})
if acc["acc"] > best["acc"]:
best.update(acc=acc["acc"], step=step)
class Guards(TrainerCallback):
def on_step_end(self, targs2, state, control, **kw):
if state.global_step <= 5 or state.global_step % 100 == 0:
elapsed = time.time() - t_start
log_metric({"event": "coverage_progress", "step": state.global_step,
"rows_seen": len(trainer.seen_rows), "target": len(parsed),
"elapsed_seconds": round(elapsed, 1)})
if state.global_step % 1000 == 0:
latest = os.path.join(save_dir, "latest-checkpoint")
model.save_pretrained(latest)
tokenizer.save_pretrained(latest)
with open(os.path.join(latest, "coverage.json"), "w") as f:
json.dump({"rows_seen": len(trainer.seen_rows), "step": state.global_step}, f)
if state.global_step % args.eval_steps == 0 and state.global_step > 0:
run_val(state.global_step)
if time.time() - t_start > args.deadline_seconds - args.eval_reserve_seconds:
control.should_training_stop = True
log_metric({"event": "time_guard_stop", "step": state.global_step,
"rows_cum": len(trainer.seen_rows)})
trainer.add_callback(Guards())
t0 = time.time()
trainer.train()
train_seconds = time.time() - t0
rows_trained = len(trainer.seen_rows)
steps_done = trainer.state.global_step
log_metric({"event": "train_done", "pool": args.pool,
"seconds": round(train_seconds, 1), "steps": steps_done,
"rows_covered": rows_trained, "n_rows_selected": len(parsed),
"n_pairs_processed": trainer.pairs_processed, "n_truncated_pairs": trainer.truncated_pairs,
"note": "rows_covered counts unique decisions actually iterated; "
"no full-epoch guarantee"})
# One fixed epoch: final weights are the selected checkpoint.
# Validation is monitored without rewinding to a partially trained checkpoint.
model_dir = os.path.join(save_dir, "model")
model.save_pretrained(model_dir)
tokenizer.save_pretrained(model_dir)
with open(os.path.join(save_dir, "training_coverage.json"), "w") as f:
json.dump({"target": len(parsed), "rows_seen": rows_trained,
"complete": rows_trained == len(parsed),
"steps": steps_done, "selection": "final fixed-epoch checkpoint",
"seen_row_ids": sorted(parsed[i]["row_id"] for i in trainer.seen_rows)}, f)
if args.push:
from huggingface_hub import HfApi
api = HfApi()
api.upload_folder(folder_path=model_dir, repo_id=args.hub_model_id, repo_type="model",
commit_message="Persist final prototype before evaluation")
api.upload_file(path_or_fileobj=os.path.join(save_dir, "training_coverage.json"),
path_in_repo="training_coverage.json", repo_id=args.hub_model_id, repo_type="model")
log_metric({"event": "checkpoint_saved_before_evaluation", "rows_covered": rows_trained})
results = {"model": "ModernJEV-Decide-Preview",
"dataset": {"id": DS_ID, "revision": DS_REV},
"base": {"id": BASE_ID, "revision": BASE_REV},
"train_pool": args.pool, "max_len": MAX_LEN, "seed": SEED,
"n_rows_selected": len(parsed), "n_a": n_a, "n_t": n_t,
"selected_row_ids_sha256": ids_hash,
"rows_covered": rows_trained, "steps": steps_done,
"train_seconds": round(train_seconds, 1),
"lr": args.lr, "token_budget": args.token_budget,
"grad_accum": args.grad_accum, "validation_monitor": best,
"input_preparation": "lazy per batch, no upfront map",
"training_input_stats": {"pairs": trainer.pairs_processed, "truncated": trainer.truncated_pairs, "at_max": trainer.pairs_at_max},
"checkpoint_selection": "final fixed-epoch checkpoint", "complete_training_coverage": rows_trained == len(parsed),
"gpu": gpu_name, "torch": torch.__version__, "attention": ATTN_IMPL,
"transformers": transformers.__version__}
val_pd_full, _ = tokenize_eval(val_full, tokenizer)
results["val_full"] = {k: v for k, v in
evaluate_rows(model, val_pd_full, val_full, device).items()
if k != "pred_labels"}
log_metric({"event": "val_full", **results["val_full"]})
test_focus, skipped_te = parse_eval_rows(ds["test"].filter(
lambda r: r["task_family"] in FOCUS, num_proc=8))
log_metric({"event": "test_prep", "n_rows": len(test_focus),
"skipped": skipped_te})
te_pd, _ = tokenize_eval(test_focus, tokenizer)
te = evaluate_rows(model, te_pd, test_focus, device)
results["test"] = {k: v for k, v in te.items() if k != "pred_labels"}
log_metric({"event": "final_test", **results["test"]})
inv_rows = stratified_subset(test_focus, 300, seed=SEED)
results["shuffled_invariance"] = shuffled_invariance(
model, tokenizer, inv_rows, device)
log_metric({"event": "shuffled_invariance", **results["shuffled_invariance"]})
results["baselines"] = build_baseline_results(parsed, val_full, test_focus)
log_metric({"event": "baselines", **results["baselines"]})
torch.manual_seed(SEED)
base_model = load_model().to(device)
results["baseline_untrained_head"] = {
"val": {k: v for k, v in evaluate_rows(
base_model, val_pd_full, val_full, device).items() if k != "pred_labels"},
"test": {k: v for k, v in evaluate_rows(
base_model, te_pd, test_focus, device).items() if k != "pred_labels"}}
log_metric({"event": "baseline_untrained_head",
**results["baseline_untrained_head"]})
del base_model
torch.cuda.empty_cache()
results["latency_gpu"] = gpu_latency(model, te_pd, device)
log_metric({"event": "latency_gpu", **results["latency_gpu"]})
results["probe"] = {"omitted": True, "reason": "Budget reserved for full prototype coverage and held-out evaluation"}
model_dir = os.path.join(save_dir, "model")
with open(os.path.join(save_dir, "results.json"), "w") as f:
json.dump(results, f, indent=2, default=str)
save_metrics(os.path.join(save_dir, "metrics.jsonl"))
with gzip.open(os.path.join(save_dir, "selected_row_ids.json.gz"), "wt") as f:
json.dump({"sha256": ids_hash, "n": len(parsed), "n_a": n_a, "n_t": n_t,
"pool": args.pool, "rows_covered": rows_trained,
"row_ids": [r["row_id"] for r in parsed]}, f)
here = os.path.dirname(os.path.abspath(__file__))
if os.path.exists(os.path.join(here, "predict.py")):
import shutil
shutil.copy(os.path.join(here, "predict.py"),
os.path.join(save_dir, "predict.py"))
if args.push:
from huggingface_hub import HfApi
api = HfApi()
assert api.model_info(args.hub_model_id).private, "Private model required"
for name in ["results.json", "metrics.jsonl", "selected_row_ids.json.gz",
"predict.py"]:
api.upload_file(path_or_fileobj=os.path.join(save_dir, name),
repo_id=args.hub_model_id, repo_type="model",
path_in_repo=name)
log_metric({"event": "persisted", "save_dir": save_dir, "pushed": args.push})
print("PROTOTYPE_DONE", flush=True)
if __name__ == "__main__":
main()
===== END FILE: recipe/train.py =====
===== FILE: predict.py =====
"""Typed-choice inference for ModernJEV-Decide-Preview.
The encoder scores each declared (state, candidate) pair; it does not generate text.
"""
import json
import os
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
MAX_LEN = 4096
MODEL_ID = "OpenMed/ModernJEV-Decide-Preview"
def serialize_state(state):
if not isinstance(state, dict):
raise TypeError("state must be a dict containing conversation and available_tools")
conv = state.get("conversation") or []
policy = state.get("policy")
first = conv[0] if conv else None
duplicate = (policy is not None and isinstance(first, dict)
and first.get("role") == "system" and first.get("content") == policy)
compact = {"available_tools": state.get("available_tools") or [], "conversation": conv}
if policy is not None and not duplicate:
compact["policy"] = policy
return json.dumps(compact, ensure_ascii=False)
def normalize_criteria(criteria):
if isinstance(criteria, list):
if not criteria or any(not isinstance(v, str) or not v.strip() for v in criteria):
raise ValueError("Answer list must contain nonempty strings")
if len(set(criteria)) != len(criteria):
raise ValueError("Answer list must contain unique strings")
criteria = {v: v for v in criteria}
if not isinstance(criteria, dict) or not criteria:
raise ValueError("criteria must be a nonempty mapping or list of unique answer strings")
if any(not isinstance(k, str) or not k.strip() or not isinstance(v, str)
for k, v in criteria.items()):
raise TypeError("Choice labels must be nonempty strings and descriptions must be strings")
return criteria
class DecisionModel:
def __init__(self, model_path=MODEL_ID, device=None, revision=None):
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
self.tokenizer = AutoTokenizer.from_pretrained(model_path, revision=revision)
self.model = AutoModelForSequenceClassification.from_pretrained(
model_path, revision=revision, attn_implementation="sdpa").to(self.device).eval()
if self.model.config.num_labels != 1:
raise ValueError("Expected a trained scalar candidate-scoring head")
@torch.inference_mode()
def decide(self, *, state, question, criteria, candidate_batch_size=8):
if not isinstance(question, str) or not question.strip():
raise ValueError("question must be nonempty text")
criteria = normalize_criteria(criteria)
if not isinstance(candidate_batch_size, int) or isinstance(candidate_batch_size, bool) or candidate_batch_size < 1:
raise ValueError("candidate_batch_size must be a positive integer")
keys = list(criteria)
text_a = question + "\n\nSTATE:\n" + serialize_state(state)
text_bs = [f"{k}: {criteria[k]}" for k in keys]
raw_a_length = len(self.tokenizer(text_a, add_special_tokens=False, verbose=False)["input_ids"])
special = self.tokenizer.num_special_tokens_to_add(pair=True)
raw_lengths = [raw_a_length + len(self.tokenizer(t, add_special_tokens=False)["input_ids"]) + special for t in text_bs]
scores = []
for begin in range(0, len(keys), candidate_batch_size):
ts = text_bs[begin:begin + candidate_batch_size]
encoded = self.tokenizer([text_a] * len(ts), ts, truncation="only_first",
max_length=MAX_LEN, padding=True, return_tensors="pt", verbose=False).to(self.device)
with torch.autocast(device_type=self.device.split(":")[0],
dtype=torch.bfloat16, enabled=self.device.startswith("cuda")):
logits = self.model(input_ids=encoded["input_ids"],
attention_mask=encoded["attention_mask"]).logits.squeeze(-1)
scores.extend(logits.float().cpu().tolist())
probabilities = torch.softmax(torch.tensor(scores), dim=0).tolist()
order = sorted(range(len(keys)), key=lambda i: (-scores[i], keys[i]))
return {"predicted_label": keys[order[0]], "allowed_choices": keys,
"candidates": [{"label": keys[i], "score": probabilities[i], "raw_score": scores[i],
"rank": rank + 1} for rank,i in enumerate(order)],
"truncated": any(n > MAX_LEN for n in raw_lengths),
"max_sequence_length": MAX_LEN,
"note": "Scores rank this supplied choice set; they are not calibrated confidence."}
_default = None
def predict_typed(question_text, state_json, criteria_json):
global _default
if _default is None:
_default = DecisionModel(os.environ.get("MODEL_PATH", MODEL_ID))
state = json.loads(state_json) if isinstance(state_json, str) else state_json
criteria = json.loads(criteria_json) if isinstance(criteria_json, str) else criteria_json
return _default.decide(state=state, question=question_text, criteria=criteria)
def predict_batch(rows):
return [predict_typed(r["question_text"], r["state_json"], r["criteria_json"]) for r in rows]
===== END FILE: predict.py =====
===== FILE: runtime-versions.json =====
{
"torch": "2.12.0+cu126",
"transformers": "5.17.0",
"datasets": "5.0.1",
"accelerate": "1.15.0",
"huggingface-hub": "1.33.0",
"kernels": "0.16.0"
}
===== END FILE: runtime-versions.json =====
===== FILE: evaluate_open_labels.py =====
"""Supplemental frozen-checkpoint evaluation; no training or Hub writes.
Run --audit-only first. After final checkpoint is saved, supply --model and,
for a Hub model, --revision. CPU is default; JSONL predictions resume safely.
"""
import argparse
import hashlib
import importlib.util
import json
import statistics
from collections import Counter
from pathlib import Path
ROOT = Path(__file__).resolve().parent
DATA_REV = 'f2fb14e4ec977c420f376c08785664cd38763d7e'
CACHE = Path('/private/tmp/atd-ml-intern-hub-cache/MaziyarPanahi___agent_tool_decisions-180_k/default/0.0.0') / DATA_REV
FAMILIES = ('agent_next_action_type', 'tool_selection', 'when_to_call_tool')
EXPECTED = {FAMILIES[0]: (1158, 616), FAMILIES[1]: (542, 79), FAMILIES[2]: (3652, 1295)}
def write_json(path, value):
temp = path.with_suffix('.tmp')
temp.write_text(json.dumps(value, indent=2) + '\n')
temp.replace(path)
def load_test():
from datasets import Dataset
return Dataset.from_file(str(CACHE / 'agent_tool_decisions-180_k-test.arrow'))
def audit(test):
report = {'dataset_revision': DATA_REV, 'families': {}}
for family in FAMILIES:
rows = [r for r in test if r['task_family'] == family]
labels = Counter(r['gold_label'] for r in rows)
size, correct = EXPECTED[family]
assert len(rows) == size and max(labels.values()) == correct
candidates = [len(json.loads(r['criteria_json'])) for r in rows]
assert all(r['gold_label'] in json.loads(r['criteria_json']) for r in rows)
report['families'][family] = {
'n': size, 'majority_correct': correct, 'majority_accuracy': correct / size,
'majority_labels': sorted(k for k,v in labels.items() if v == correct),
'candidate_count': {'median': statistics.median(candidates),
'min': min(candidates), 'max': max(candidates)},
'uniform_expected_accuracy': statistics.mean(1/n for n in candidates),
'reference_definition': 'Descriptive test-set constant-label majority; no model fitting',
'unseen_task_family': family == 'when_to_call_tool',
}
# Verify When2Call is absent from every training shard, not just the subset.
from datasets import Dataset
train_families = Counter()
selected_ids = set((ROOT / 'selected-row-ids.txt').read_text().splitlines())
train_gold = {f: Counter() for f in FAMILIES[:2]}
for path in sorted(CACHE.glob('*-train-*.arrow')):
columns = Dataset.from_file(str(path)).select_columns(['task_family','row_id','gold_label'])[:]
for family, row_id, gold in zip(columns['task_family'], columns['row_id'], columns['gold_label']):
train_families[family] += 1
if row_id in selected_ids and family in train_gold:
train_gold[family][gold] += 1
assert sum(train_families.values()) == 171056
assert train_families['when_to_call_tool'] == 0
assert sum(sum(c.values()) for c in train_gold.values()) == 60000
report['when2call_training_rows'] = 0
for family, counts in train_gold.items():
label = min(counts, key=lambda k: (-counts[k], k))
rows = [r for r in test if r['task_family'] == family]
report['families'][family]['training_fixed_majority'] = {
'label': label, 'correct': sum(r['gold_label'] == label for r in rows),
'n': len(rows),
'majority_label_absent_from_choices': sum(label not in json.loads(r['criteria_json']) for r in rows),
}
return report
def open_labels(client):
checks = []
for n in (2, 3, 7, 20):
descriptions = ['Escalate this damaged parcel to a human support agent.',
'Search the product catalog for a new item.']
descriptions += [f'Route to unrelated department number {i} for a different request.' for i in range(n-2)]
criteria = {f'fresh_choice_{n}_{i}': d for i,d in enumerate(descriptions)}
state = {'policy': 'Escalate damaged parcels to a human support agent.',
'conversation': [{'role':'user','content':'My parcel arrived broken. I need help.'}]}
question = 'Which of these custom workflow branches should handle this request?'
outputs = []
variants = (criteria, dict(reversed(list(criteria.items()))),
{f'opaque_{n}_{i}': d for i,d in enumerate(descriptions)}, descriptions)
for variant in variants:
answer = client.decide(state=state, question=question, criteria=variant, candidate_batch_size=4)
allowed = list(variant)
assert answer['predicted_label'] in allowed
assert sorted(c['label'] for c in answer['candidates']) == sorted(allowed)
assert len(answer['candidates']) == n
outputs.append(answer)
original_desc = criteria[outputs[0]['predicted_label']]
renamed_desc = variants[2][outputs[2]['predicted_label']]
checks.append({'answer_count': n, 'interface_pass': True,
'reorder_same_label': outputs[0]['predicted_label'] == outputs[1]['predicted_label'],
'rename_same_description': original_desc == renamed_desc,
'selected_expected_description': original_desc == descriptions[0],
'outputs': outputs})
return {'scalar_head': client.model.config.num_labels,
'interpretation': 'Interface checks and illustrative decisions, not a clinical benchmark or general accuracy estimate.',
'checks': checks}
def main():
p = argparse.ArgumentParser()
p.add_argument('--audit-only', action='store_true')
p.add_argument('--interface-only', action='store_true')
p.add_argument('--model')
p.add_argument('--revision')
p.add_argument('--device', default='cpu')
p.add_argument('--output-dir', type=Path, default=ROOT / 'supplemental-evaluation')
args = p.parse_args()
out = args.output_dir; out.mkdir(parents=True, exist_ok=True)
test = load_test(); baseline = audit(test)
write_json(out/'per-task-baselines.json', baseline)
print(json.dumps(baseline), flush=True)
if args.audit_only: return
if not args.model: p.error('--model is required for checkpoint evaluation')
local = Path(args.model).is_dir()
if not local and not args.revision: p.error('Pin --revision for a Hub model')
if local:
weights = sorted(Path(args.model).glob('*.safetensors'))
assert weights, 'Local checkpoint must have saved safetensors'
h = hashlib.sha256()
for path in weights:
with path.open('rb') as f:
for chunk in iter(lambda:f.read(8*1024*1024), b''): h.update(chunk)
identity = h.hexdigest()
else: identity = args.revision
manifest = {'model': args.model, 'checkpoint_identity': identity, 'dataset_revision': DATA_REV}
if (out/'manifest.json').exists(): assert json.loads((out/'manifest.json').read_text()) == manifest
write_json(out/'manifest.json', manifest)
spec = importlib.util.spec_from_file_location('inference', ROOT/'predict_open_labels.py')
helper = importlib.util.module_from_spec(spec); spec.loader.exec_module(helper)
import torch
torch.set_num_threads(4)
client = helper.DecisionModel(args.model, device=args.device, revision=args.revision)
assert client.model.config.num_labels == 1
write_json(out/'open-label-checks.json', open_labels(client))
if args.interface_only: return
rows = [r for r in test if r['task_family'] == 'when_to_call_tool']
path = out/'when2call-predictions.jsonl'; previous = {}
if path.exists():
for line in path.read_text().splitlines():
if line: item = json.loads(line); previous[item['row_id']] = item
assert set(previous).issubset({r['row_id'] for r in rows})
with path.open('a') as f:
for row in rows:
if row['row_id'] in previous: continue
pred = client.decide(state=json.loads(row['state_json']), question=row['question_text'],
criteria=json.loads(row['criteria_json']), candidate_batch_size=4)
rec = {'row_id': row['row_id'], 'gold': row['gold_label'], **pred}
f.write(json.dumps(rec)+'\n'); f.flush(); previous[row['row_id']] = rec
if len(previous) % 100 == 0: print(f'When2Call {len(previous)}/{len(rows)}', flush=True)
correct = sum(previous[r['row_id']]['predicted_label'] == r['gold_label'] for r in rows)
n = len(rows); reference = baseline['families']['when_to_call_tool']['majority_accuracy']
result = {'task': 'when_to_call_tool', 'correct': correct, 'n': n, 'accuracy': correct/n,
'majority_reference': reference, 'lift_percentage_points': 100*(correct/n-reference),
'unseen_task_family': True, 'checkpoint_identity': identity,
'truncated_rows': sum(x['truncated'] for x in previous.values())}
write_json(out/'when2call-results.json', result)
print(json.dumps(result), flush=True)
if __name__ == '__main__': main()
===== END FILE: evaluate_open_labels.py =====