Add trained contextual action candidate with browser evaluation and explicit real-web limits
c8e5620 verified Download baim/train.py from devildasdf/devils-agent: direct link, hf CLI and curl.
- Browser
- Download file 6.03 kB
-
https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/train.py
- Command line
-
hf download hf://devildasdf/devils-agent/baim/train.py
-
curl -L -o train.py https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/train.py
6.03 kB
| """Deterministic CPU training, validation-only selection, calibrated test report.""" | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import random | |
| import time | |
| import torch | |
| from torch import nn | |
| from .features import encode, fit_vocab | |
| from .model import PointerPolicy | |
| from .synthetic import load | |
| def logits(model, inputs, batch_size=64): | |
| output = [model(*(x[start:start+batch_size] for x in inputs)) | |
| for start in range(0,len(inputs[0]),batch_size)] | |
| return tuple(torch.cat([row[i] for row in output]) for i in range(2)) | |
| def calibrate(predictions, labels): | |
| # Fit one scalar temperature on validation only; test labels are never used. | |
| temperatures = torch.logspace(-1,1,81) | |
| losses = [nn.functional.cross_entropy(predictions / t, labels).item() for t in temperatures] | |
| return float(temperatures[losses.index(min(losses))]) | |
| def metrics(action_logits, target_logits, labels, targets, temperatures): | |
| action = action_logits.argmax(-1) | |
| target = target_logits.argmax(-1) | |
| joint = (action == labels) & (target == targets) | |
| ap = (action_logits/temperatures[0]).softmax(-1).max(-1).values | |
| tp = (target_logits/temperatures[1]).softmax(-1).max(-1).values | |
| # Marginals are calibrated separately. Do not call their product calibrated. | |
| def ece(prob, correct): | |
| total = 0.0 | |
| for low in torch.arange(0,1,.1): | |
| selected = (prob >= low) & (prob < low+.1 if low < .9 else prob <= 1) | |
| if selected.any(): | |
| total += float(selected.float().mean() * (prob[selected].mean()-correct[selected].float().mean()).abs()) | |
| return total | |
| return dict(samples=len(labels), action_accuracy=float((action==labels).float().mean()), | |
| target_accuracy=float((target==targets).float().mean()), | |
| joint_step_accuracy=float(joint.float().mean()), | |
| action_ece=ece(ap,action==labels), target_ece=ece(tp,target==targets), | |
| candidate_recall=float((targets>=0).float().mean())) | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('--data',default='datasets/synthetic-v1') | |
| parser.add_argument('--output',default='models/v000') | |
| parser.add_argument('--encoder',choices=['mean','gru','transformer'],default='mean') | |
| parser.add_argument('--no-lexical',action='store_true') | |
| parser.add_argument('--contextual-action',action='store_true') | |
| parser.add_argument('--epochs',type=int,default=20) | |
| parser.add_argument('--width',type=int,default=64) | |
| parser.add_argument('--seed',type=int,default=1729) | |
| args = parser.parse_args() | |
| torch.set_num_threads(2) | |
| torch.set_num_interop_threads(1) | |
| torch.manual_seed(args.seed) | |
| random.seed(args.seed) | |
| torch.use_deterministic_algorithms(True) | |
| root = Path(args.output) | |
| root.mkdir(parents=True,exist_ok=True) | |
| train = load(Path(args.data)/'train.jsonl') | |
| validation = load(Path(args.data)/'validation.jsonl') | |
| vocab = fit_vocab(train) | |
| inputs,actions,targets,_ = encode(train,vocab) | |
| valid_inputs,valid_actions,valid_targets,_ = encode(validation,vocab) | |
| if (targets < 0).any(): | |
| raise ValueError('training targets missing after retrieval') | |
| model = PointerPolicy(vocab_size=len(vocab),width=args.width,encoder=args.encoder, | |
| lexical_features=not args.no_lexical,contextual_action=args.contextual_action) | |
| optimizer = torch.optim.AdamW(model.parameters(),lr=.002,weight_decay=.01) | |
| history, best, best_state = [], -1, None | |
| started = time.perf_counter() | |
| for epoch in range(args.epochs): | |
| model.train() | |
| order = torch.randperm(len(train)) | |
| losses = [] | |
| for start in range(0,len(train),64): | |
| indices = order[start:start+64] | |
| a,t = model(*(x[indices] for x in inputs)) | |
| loss = nn.functional.cross_entropy(a,actions[indices]) + nn.functional.cross_entropy(t,targets[indices]) | |
| optimizer.zero_grad() | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(),1.0) | |
| optimizer.step() | |
| losses.append(loss.item()) | |
| model.eval() | |
| va,vt = logits(model,valid_inputs) | |
| result = metrics(va,vt,valid_actions,valid_targets,(1,1)) | |
| score = result['joint_step_accuracy'] | |
| if score > best: | |
| best = score | |
| best_state = {key:value.detach().clone() for key,value in model.state_dict().items()} | |
| row = dict(epoch=epoch+1,loss=sum(losses)/len(losses),validation=result) | |
| history.append(row) | |
| print(json.dumps(row),flush=True) | |
| model.load_state_dict(best_state) | |
| model.eval() | |
| va,vt = logits(model,valid_inputs) | |
| temperatures = [calibrate(va,valid_actions),calibrate(vt,valid_targets)] | |
| model.save_pretrained(root) | |
| (root/'vocab.json').write_text(json.dumps(vocab),encoding='utf-8') | |
| (root/'calibration.json').write_text(json.dumps(dict(temperatures=temperatures,split='validation')),encoding='utf-8') | |
| evaluations = {} | |
| for split in ['validation','test','novel_wording']: | |
| rows = load(Path(args.data)/f'{split}.jsonl') | |
| x,a,t,_ = encode(rows,vocab) | |
| la,lt = logits(model,x) | |
| evaluations[split] = metrics(la,lt,a,t,temperatures) | |
| report = dict(architecture=vars(args),parameter_count=sum(p.numel() for p in model.parameters()), | |
| threads=2,device='cpu',training_seconds=time.perf_counter()-started, | |
| dataset_manifest=json.loads((Path(args.data)/'manifest.json').read_text()), | |
| history=history,evaluation=evaluations, | |
| limitations='Synthetic single-step action/target prediction; not arbitrary-site task success.', | |
| production_promoted=False) | |
| (root/'training-report.json').write_text(json.dumps(report,indent=2),encoding='utf-8') | |
| print(json.dumps(dict(parameter_count=report['parameter_count'],evaluation=evaluations),indent=2)) | |
| if __name__ == '__main__': | |
| main() | |