Text Classification
Transformers
Safetensors
English
modernbert
encoder
decision-model
tool-routing
agentic
preview
text-embeddings-inference
Instructions to use MaziyarPanahi/ModernJEV-Decide-Preview with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use MaziyarPanahi/ModernJEV-Decide-Preview with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="MaziyarPanahi/ModernJEV-Decide-Preview")# Load model directly from transformers import AutoTokenizer, AutoModelForSequenceClassification tokenizer = AutoTokenizer.from_pretrained("MaziyarPanahi/ModernJEV-Decide-Preview") model = AutoModelForSequenceClassification.from_pretrained("MaziyarPanahi/ModernJEV-Decide-Preview", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download workflow/recipe-bundle.txt from MaziyarPanahi/ModernJEV-Decide-Preview: direct link, hf CLI and curl.
- Browser
- Download file 55.6 kB
-
https://huggingface.co/MaziyarPanahi/ModernJEV-Decide-Preview/resolve/main/workflow/recipe-bundle.txt
- Command line
-
hf download hf://MaziyarPanahi/ModernJEV-Decide-Preview/workflow/recipe-bundle.txt
-
curl -L -o recipe-bundle.txt https://huggingface.co/MaziyarPanahi/ModernJEV-Decide-Preview/resolve/main/workflow/recipe-bundle.txt
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 ===== | |