TSRDA / main_method /code /train_temporal_cnn.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
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
@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)