File size: 5,919 Bytes
795f737 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | """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()
|