| """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 |
|
|
|
|
| @torch.inference_mode() |
| 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): |
| |
| 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 |
| |
| 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('--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) |
| 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() |
|
|