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)