"""Replay final temporal CNN training, preserving the scored loss/augmentation/RNG logic.""" import argparse,json,random,time from pathlib import Path import numpy as np import torch from torch.utils.data import DataLoader import config as C from dataset import PASTISPatchDataset,train_patch_ids,val_patch_ids,crop_ij,pad_collate from temporal_cnn import make_model from training_feature_clusters import metrics_cm from augmentation import aug_dihedral,aug_temporal_subsample,aug_frame_dropout,aug_spectral,lr_factor OUT=Path(__file__).resolve().parents[1]/"runs" CEREALS=[2,3,4,6,10,11,17,18] def write_json(path,data): path.parent.mkdir(parents=True,exist_ok=True);path.write_text(json.dumps(data,indent=2)) def loader(ids,batch,shuffle=False): return DataLoader(PASTISPatchDataset(ids),batch_size=batch,shuffle=shuffle,num_workers=2,pin_memory=True,persistent_workers=True,collate_fn=pad_collate) def target_for(y,classes): target=torch.full_like(y,-100) for i,c in enumerate(classes):target[y==c]=i return target @torch.inference_mode() def validate(model,dl,classes): model.eval();cm=np.zeros((20,20),np.int64) for (x,pos,_),y in dl: x,pos=x.cuda(),pos.cuda() for i,j in C.FT_CROP_IJ: xc,yc=crop_ij(x,y,i,j,4) with torch.autocast('cuda',dtype=torch.float16):pred=model(xc,pos).argmax(1) pred=np.array(classes)[pred.cpu().numpy().ravel()];gt=yc.numpy().ravel() valid=np.isin(gt,classes);cm+=np.bincount(20*gt[valid]+pred[valid],minlength=400).reshape(20,20) result=metrics_cm(cm) if len(classes)<18:result['selection_score']=float(np.array(result['per_class_iou'])[np.array(classes)-1].mean()) else:result['selection_score']=result['miou'] return result def train(args): classes=CEREALS if args.kind=='cereal' else list(range(1,19)) folder=OUT/f'32X32_{args.kind}_seed{args.seed}';folder.mkdir(parents=True,exist_ok=True) model=make_model(args.kind,len(classes)).cuda();optimizer=torch.optim.AdamW(model.parameters(),lr=1e-3,weight_decay=.01) torch.manual_seed(args.seed);random.seed(args.seed) scaler=torch.amp.GradScaler('cuda');td=loader(train_patch_ids(),2,True);vd=loader(val_patch_ids(),2) start=0;best=-1. if (folder/'latest.tar').exists(): ck=torch.load(folder/'latest.tar',weights_only=False);model.load_state_dict(ck['state']);optimizer.load_state_dict(ck['optimizer']);scaler.load_state_dict(ck['scaler']);start=ck['epoch']+1;best=ck['best'];torch.set_rng_state(ck['rng']);torch.cuda.set_rng_state_all(ck['cuda_rng']);random.setstate(ck['python_rng']) for epoch in range(start,args.epochs): model.train();started=time.monotonic();total=0.;steps=0 for (x,pos,_),y in td: x,pos,y=x.cuda(),pos.cuda(),y.cuda() for i,j in C.FT_CROP_IJ: xc,yc=crop_ij(x,y,i,j,4);xc,yc=aug_dihedral(xc,yc);xc,pc=aug_temporal_subsample(xc,pos);xc=aug_frame_dropout(xc);xc=aug_spectral(xc) target=target_for(yc,classes);optimizer.zero_grad(set_to_none=True) factor=lr_factor(epoch*76+steps,args.epochs*76,3*76,.01) for g in optimizer.param_groups:g['lr']=1e-3*factor with torch.autocast('cuda',dtype=torch.float16):logits=model(xc,pc) loss=torch.nn.functional.cross_entropy(logits.float(),target,ignore_index=-100,label_smoothing=.05) if (target!=-100).any() else logits.float().sum()*0 if not torch.isfinite(loss):raise RuntimeError('Nonfinite specialist loss') scaler.scale(loss).backward();scaler.unscale_(optimizer);torch.nn.utils.clip_grad_norm_(model.parameters(),1.);scaler.step(optimizer);scaler.update();total+=loss.item();steps+=1 assert steps==76 metric=validate(model,vd,classes) if metric['selection_score']>best: best=metric['selection_score'];torch.save({'state':model.state_dict(),'classes':classes,'kind':args.kind,'seed':args.seed,'epoch':epoch,'validation':metric},folder/'best.tar') torch.save({'state':model.state_dict(),'optimizer':optimizer.state_dict(),'scaler':scaler.state_dict(),'epoch':epoch,'best':best,'rng':torch.get_rng_state(),'cuda_rng':torch.cuda.get_rng_state_all(),'python_rng':random.getstate()},folder/'latest.tar') write_json(folder/'progress.json',{'epoch_completed':epoch+1,'epochs':args.epochs,'best_selection_score':best,'seconds_per_epoch':time.monotonic()-started}) print('epoch',epoch+1,'loss',total/steps,'selection',metric['selection_score'],'best',best,'seconds',time.monotonic()-started,flush=True) if __name__=='__main__': ap=argparse.ArgumentParser();ap.add_argument('--kind',choices=['cnn_control','attention'],default='cnn_control');ap.add_argument('--seed',type=int,default=3407);ap.add_argument('--epochs',type=int,default=100);args=ap.parse_args() torch.set_num_threads(4);torch.manual_seed(args.seed);random.seed(args.seed);np.random.seed(args.seed);torch.backends.cuda.matmul.allow_tf32=True;torch.backends.cudnn.allow_tf32=True;train(args)