File size: 5,085 Bytes
71d64bb | 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 | """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)
|