"""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()