"""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): # 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('--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()