import copy,json,time from pathlib import Path import torch from safetensors.torch import save_file from structured_policy import StructuredPolicy,FEATURES,encode torch.manual_seed(42); torch.set_num_threads(4) device='cuda';out=Path('release/structured');out.mkdir(parents=True,exist_ok=True) def load(split,n): rows=[json.loads(x) for x in open('data/sft/'+split+'.jsonl')][:n] x=torch.zeros(len(rows),4,15);mask=torch.zeros(len(rows),4,dtype=torch.bool);target=torch.zeros(len(rows),4) for i,r in enumerate(rows): u=json.loads(r['messages'][1]['content'].split('\n\nRequested field:')[0]);state=json.loads(u['context']) k=len(r['keys']);x[i,:k]=torch.tensor(encode(state,r['keys']));mask[i,:k]=True;target[i,:k]=torch.tensor(r['target']) return x.to(device),mask.to(device),target.to(device) x,mask,target=load('train',4096);vx,vm,vt=load('val',512) model=StructuredPolicy().to(device) opt=torch.optim.AdamW(model.parameters(),lr=.001,weight_decay=.0001) best=float('inf');history=[];start=time.monotonic() for epoch in range(1,101): model.train();order=torch.randperm(len(x),device=device) for i in range(0,len(x),128): j=order[i:i+128];logp=model(x[j],mask[j]).log_softmax(-1) loss=-(target[j]*logp).sum(-1).mean();loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(),1.);opt.step();opt.zero_grad(set_to_none=True) if epoch%10==0: model.eval() with torch.no_grad(): logits=model(vx,vm);ce=float(-(vt*logits.log_softmax(-1)).sum(-1).mean()) acc=float((logits.argmax(-1)==vt.argmax(-1)).float().mean()) row=dict(epoch=epoch,validation_ce=ce,teacher_agreement=acc);history.append(row);print(row,flush=True) if ce