{"cells":[{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true,"_kg_hide-output":true},"cell_type":"code","source":"!pip install ../input/timm-0-1-30/timm-0.1.30-py3-none-any.whl","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nfrom PIL import Image, ImageFilter\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torch.optim import *\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nfrom torchvision import models\nimport time\nfrom tqdm import tqdm\nimport random\nimport timm\nimport sys\nsys.path.append('../input/autoaug')\nfrom auto_augment import AutoAugment, Cutout","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = True","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"seed_everything(0)\n\nnum_classes = 5\nbs = 64\nlr = 5e-4\nIMG_SIZE = 224","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_path = '../input/cassava-leaf-disease-classification/train_images/'\ntest_path = '../input/cassava-leaf-disease-classification/test_images/'\n\ntrain_csv = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\nsample = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_csv.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class MyDataset(Dataset):\n    \n    def __init__(self, dataframe, transform=None, test=False):\n        self.df = dataframe\n        self.transform = transform\n        self.test = test\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        \n        label = self.df.label.values[idx]\n        p = self.df.image_id.values[idx]\n        \n        if self.test == False:\n            p_path = train_path + p\n        else:\n            p_path = test_path + p\n            \n        image = cv2.imread(p_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = transforms.ToPILImage()(image)\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.Resize((IMG_SIZE,IMG_SIZE)),\n    transforms.RandomHorizontalFlip(),\n    AutoAugment(),\n    transforms.ToTensor()\n])\n\ntest_transform = transforms.Compose([\n    transforms.Resize((IMG_SIZE,IMG_SIZE)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor()\n])\n\n\ntestset      = MyDataset(sample, transform=test_transform, test=True)\ntest_loader  = DataLoader(testset, batch_size=bs, shuffle=False, num_workers=4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class AverageMeter:\n    \"\"\"\n    Computes and stores the average and current value\n    \"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_model(model, epoch):\n    model.train() \n    \n    losses = AverageMeter()\n    accs = AverageMeter()\n    \n    tk = tqdm(train_loader, total=len(train_loader), position=0, leave=True)\n    for idx, (imgs, labels) in enumerate(tk):\n        imgs_train, labels_train = imgs.cuda(), labels.cuda().long()\n        output_train = model(imgs_train)\n\n        loss = criterion(output_train, labels_train)\n        \n        optimizer.zero_grad() \n        loss.backward()\n        optimizer.step() \n        \n        accs.update((output_train.argmax(1)==labels_train).sum().item()/imgs_train.size(0),imgs_train.size(0))\n        losses.update(loss.item(), imgs_train.size(0))\n\n        tk.set_postfix(loss=losses.avg,acc=accs.avg)\n        \n    return losses.avg\n\n\ndef test_model(model):    \n    model.eval()\n    \n    losses = AverageMeter()\n    accs = AverageMeter()\n    \n    with torch.no_grad():\n        tk = tqdm(val_loader, total=len(val_loader), position=0, leave=True)\n        for idx, (imgs, labels) in enumerate(tk):\n            imgs_valid, labels_valid = imgs.cuda(), labels.cuda().long()\n            output_valid = model(imgs_valid)\n            \n            loss = criterion(output_valid, labels_valid)\n\n            losses.update(loss.item(), imgs_valid.size(0))\n            accs.update((output_valid.argmax(1)==labels_valid).sum().item()/imgs_valid.size(0),imgs_valid.size(0))\n            \n            tk.set_postfix(loss=losses.avg,acc=accs.avg)\n\n            \n    return losses.avg,accs.avg","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# train_df, val_df = train_test_split(train_csv,test_size=0.2,stratify=train_csv.label)\n\n# trainset = MyDataset(train_df, transform=train_transform)\n# train_loader = DataLoader(trainset, batch_size=bs, shuffle=True, num_workers=4)\n\n# valset = MyDataset(val_df, transform=test_transform)\n# val_loader = DataLoader(valset, batch_size=bs, shuffle=False, num_workers=4)\n\n# model = timm.create_model('tf_efficientnet_b0_ns', pretrained=True, num_classes=num_classes)\n# model.cuda()\n\n# optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0)\n# criterion = nn.CrossEntropyLoss()\n# scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, \\\n#                                                        patience=1, verbose=True, min_lr=1e-5)\n\n# best_acc = 0\n# n_epochs = 10\n\n# for epoch in range(n_epochs):\n#     train_loss = train_model(model, epoch)\n#     val_loss, acc = test_model(model)\n\n#     if acc > best_acc:\n#         best_acc = acc\n#         torch.save(model.state_dict(), 'weight.pt')\n\n#     print('current_val_acc:', acc, 'best_val_acc:', best_acc)\n\n#     scheduler.step(acc)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-output":true},"cell_type":"code","source":"model = timm.create_model('tf_efficientnet_b0_ns', pretrained=False, num_classes=num_classes)\nmodel.cuda()\n\nmodel.load_state_dict(torch.load('../input/cassava-b0/weight.pt'))\nmodel.eval()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_pred = []\n\nwith torch.no_grad():\n    for i, data in enumerate(tqdm(test_loader, position=0, leave=True)):\n        images, _ = data\n        images = images.cuda()\n\n        pred = model(images)\n\n        pred = pred.argmax(1).cpu().detach().numpy().astype('int')\n\n        test_pred.extend(pred)\n\nsample.label = test_pred\nsample.to_csv('submission.csv',index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample","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}