{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"!pip install ../input/easydict/easydict-1.9-py2.py3-none-any.whl\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"from sklearn.model_selection import GroupKFold, StratifiedKFold\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport random\n\nimport os\nimport sys\n\n\nsys.path.insert(1, '../input/snapmix/')\nimport glob\nfrom utils import load_checkpoint,get_train_setting\nfrom trainer.comm_test import validate\nfrom trainer.comm_train import train\nimport networks.resnet_ft as resnet_ft\nfrom easydict import EasyDict as edict \nimport torch\nimport torch.nn as nn\nfrom datasets.cassava import ImageLoader\nfrom datasets.tfs import get_cassava_transform\nfrom torch.utils import data\nimport time\nfrom tqdm import tqdm\nimport copy\nfrom collections import Counter\n\ndef set_env(seed=0):\n    # set seeding\n    random.seed(seed)\n    np.random.seed(seed) # cpu vars\n    torch.manual_seed(seed) # cpu  vars\n    torch.cuda.manual_seed(seed) # cpu  vars\n    torch.cuda.manual_seed_all(seed) # gpu vars\n    \ndef predict(model,testloader,midlevel=False):   \n    model.eval()\n    time_start = time.time()\n    pbar = tqdm(testloader, dynamic_ncols=True, total=len(testloader))\n    pres = []\n    for idx, (input, _) in enumerate(pbar):\n\n        input = input.cuda()\n\n        if conf.tta is None:\n            output,_,moutput = model(input)\n        else:\n            bs, ncrops, c, h, w = input.size()\n            output,_,moutput = model(input.view(-1,c,h,w))\n            output = output.view(bs, ncrops, -1).mean(1)\n            moutput = moutput.view(bs, ncrops, -1).mean(1)\n        if midlevel:\n            foutput = output + moutput\n        else:\n            foutput = output\n        pre = torch.argmax(foutput,dim=1)\n        pres.append(pre)\n    pres = torch.cat(pres)\n    \n    return pres\n\ndef get_dataset(conf):\n\n    datadir = 'data/cassava'\n\n    if conf and 'datadir' in conf:\n        datadir = conf.datadir\n\n\n    trainpd,valpd = None,None\n\n    testimgdir = datadir + '/test_images'\n    testfile = glob.glob(testimgdir+'/*.jpg')\n    testfile = [os.path.basename(fn) for fn in testfile]\n    testpd = pd.DataFrame(testfile, columns =['image_id']) \n    testpd['label'] = 0\n\n  \n    if 'foldid' in conf:\n        traindata = pd.read_csv(datadir+'/train.csv')\n        folds = StratifiedKFold(n_splits=5).split(np.arange(traindata.shape[0]), traindata.label.values)\n        trainidx,validx = list(folds)[conf.foldid]\n        trainpd = traindata.loc[trainidx,:].reset_index(drop=True)\n        valpd = traindata.loc[validx,:].reset_index(drop=True)\n    imgdir = datadir + '/train_images'\n    transform_train,transform_test = get_cassava_transform(conf)\n    ds_train = ImageLoader(imgdir, train=True, transform=transform_train,pdata=trainpd)\n    \n    \n\n    if conf.tta is None or conf.tta == 2:\n        \n        ds_val = ImageLoader(imgdir, train=False, transform=transform_test,pdata=valpd,tta=conf.tta)\n        ds_test = ImageLoader(testimgdir, train=False, transform=transform_test,pdata=testpd,tta=conf.tta)\n    else:\n        ds_test = ImageLoader(testimgdir, train=False, transform=transform_train,pdata=testpd,tta=conf.tta)\n        ds_val = ImageLoader(imgdir, train=False, transform=transform_train,pdata=valpd,tta=conf.tta)\n\n    return ds_train,ds_val,ds_test,testpd\n\ndef get_params(model,conf=None):\n\n    if conf is not None and 'prams_group' in  conf:\n        prams_group = conf.prams_group\n        lr_group = conf.lr_group\n        params = []\n        for pram,lr in zip(prams_group,lr_group):\n            params.append({'params':model.module.get_params(pram),'lr': lr})\n\n        return params\n\n    return model.parameters()\n\n ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"set_env(seed=0)\nfoldid = 2\n\nconf = edict({\n    'depth':50,\n    'pretrained':True,\n    'num_class':5,\n    'midlevel':False,\n    'datadir':'../input/cassava-leaf-disease-classification',\n    'dataset':'cassava',\n    'testing':False,\n    'tta': None,\n    'foldid':foldid,\n    'cropsize':448,\n    'netname':'resnet50',\n    'net_type':'resnet_ft',\n    'prams_group':['ftlayer','freshlayer'],\n    'lr_group':[0.001,0.01],\n    'lrstep':[20],\n    'lr': 0.001,\n    'epochs':25,\n    'lrgamma':0.1,\n    'criterion':'CrossEntropyLoss',\n    'reduction':'none',\n    'momentum':0.9,\n    'weight_decay':1e-4,\n    'pretrained':True,\n    'mixmethod': 'snapmix',\n    'prob': 1,\n    'beta': 5}\n   )\n\n\n\nds_train,ds_val,ds_test,testpd = get_dataset(conf)\ntrain_loader =data.DataLoader(ds_train, batch_size=16, shuffle= True, num_workers=8, pin_memory=True)\nval_loader =data.DataLoader(ds_val, batch_size=32, shuffle= False, num_workers=8, pin_memory=True)\ntest_loader =data.DataLoader(ds_test, batch_size=32, shuffle= False, num_workers=16, pin_memory=True)\n\n\nmodel = eval(conf.net_type).get_net(conf)\nmodel = nn.DataParallel(model).cuda()\n\noptimizer = torch.optim.SGD(get_params(model,conf),conf.lr,momentum=conf.momentum,weight_decay=conf.weight_decay,nesterov=True)\ncriterion = nn.CrossEntropyLoss(reduction='none').cuda()\n\nscheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer,milestones=conf.lrstep, gamma=conf.lrgamma, last_epoch=-1)\n\noutfile = './'+conf.netname+'_f'+str(foldid)+'.pt'\n\nif not os.path.isfile('../input/resfold2/'+outfile):\n    best_score = 0.\n    ## ------main loop-----\n    for epoch in range(0, conf.epochs):  \n        lr = optimizer.param_groups[0]['lr']\n        print(\"Epoch: [{} | {} LR: {}\".format(epoch+1,conf.epochs,lr))\n        tmp_loss = train(train_loader, model, criterion, optimizer, conf)\n        scheduler.step()\n        infostr = {'Epoch:  {}   train_loss: {}'.format(epoch+1,tmp_loss)}\n        print(infostr)\n        with torch.no_grad():\n            val_score,val_loss,mscore,ascore = validate(val_loader, model,criterion, conf)\n            comscore = val_score\n            if conf.midlevel:\n                comscore = ascore\n            is_best = comscore > best_score\n            best_score = max(comscore,best_score)\n            infostr = {'Epoch:  {:.4f}   loss: {:.4f},gs: {:.4f},bs:{:.4f},ms: {:.4f},as:{:.4f}'.format(epoch+1,val_loss,val_score,best_score,mscore,ascore)}\n            print(infostr)\n            if is_best:\n                mdict = {'state_dict': model.module.state_dict(),'epoch':epoch}\n                torch.save(mdict,outfile)\nelse:\n    outfile = '../input/resfold2/'+outfile\n\nload_checkpoint(model,outfile)\nwith torch.no_grad():\n    pres = predict(model,test_loader,conf.midlevel).cpu().numpy()\ntestpd['label'] = pres\ntestpd.to_csv('submission.csv', index=False)\nprint(testpd.head())\n                    \n                    ","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":4}