Download train_structured.py from guilindev/pacman-decision-tiny: direct link, hf CLI and curl.
- Browser
- Download file 3.02 kB
-
https://huggingface.co/guilindev/pacman-decision-tiny/resolve/main/train_structured.py
- Command line
-
hf download hf://guilindev/pacman-decision-tiny/train_structured.py
-
curl -L -o train_structured.py https://huggingface.co/guilindev/pacman-decision-tiny/resolve/main/train_structured.py
3.02 kB
| 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<best: | |
| best=ce;best_state=copy.deepcopy(model.state_dict());best_epoch=epoch | |
| model.load_state_dict(best_state);model.eval() | |
| with torch.no_grad(): | |
| p=model(vx,vm).softmax(-1);acc=float((p.argmax(-1)==vt.argmax(-1)).float().mean()) | |
| perm=torch.tensor([2,0,3,1],device=device) | |
| permuted=model(vx[:,perm],vm[:,perm]).softmax(-1) | |
| equiv_error=float((permuted-p[:,perm]).abs().max()) | |
| assert equiv_error<1e-5 | |
| save_file({k:v.cpu().contiguous() for k,v in model.state_dict().items()},out/'structured_policy.safetensors') | |
| config=dict(architecture='Shared option MLP with mean/max set pooling',features=FEATURES, | |
| parameters=sum(p.numel() for p in model.parameters()),dtype='float32',train_rows=4096,validation_rows=512, | |
| seed=42,epochs=100,selected_epoch=best_epoch,batch=128,lr=.001,weight_decay=.0001, | |
| gpu=torch.cuda.get_device_name(),training_seconds=time.monotonic()-start, | |
| validation_teacher_agreement=acc,validation_ce=best,permutation_max_error=equiv_error, | |
| note='Standalone structured-feature supervised baseline; not an LLM fine-tune',history=history) | |
| (out/'structured_config.json').write_text(json.dumps(config,indent=2)) | |
| print('STRUCTURED_TRAINING_COMPLETE',json.dumps(config),flush=True) | |