MaziyarPanahi's picture
ModernJEV-Decide-Preview: model, results, compute costs and ML Intern workflow
d5b398e
Raw History Blame Contribute Delete
40.5 kB
"""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()