Download main_method/code/train_temporal_cnn.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 5.09 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/train_temporal_cnn.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/train_temporal_cnn.py
-
curl -L -o train_temporal_cnn.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/train_temporal_cnn.py
5.09 kB
| """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 | |
| 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) | |