{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport json\nfrom PIL import Image\nimport os\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import StratifiedKFold\n\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\nfrom albumentations import (\n    Compose, OneOf, Normalize, Resize, RandomResizedCrop, RandomCrop, HorizontalFlip, VerticalFlip, \n    RandomBrightness, RandomContrast, RandomBrightnessContrast, Rotate, ShiftScaleRotate, Cutout, \n    IAAAdditiveGaussianNoise, Transpose\n    )\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-07-06T09:58:53.951290Z","iopub.execute_input":"2021-07-06T09:58:53.951635Z","iopub.status.idle":"2021-07-06T09:58:57.859075Z","shell.execute_reply.started":"2021-07-06T09:58:53.951607Z","shell.execute_reply":"2021-07-06T09:58:57.858095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n\npackage_path = '../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master'\nsys.path.append(package_path)","metadata":{"execution":{"iopub.status.busy":"2021-07-06T09:59:56.063909Z","iopub.execute_input":"2021-07-06T09:59:56.064268Z","iopub.status.idle":"2021-07-06T09:59:56.068604Z","shell.execute_reply.started":"2021-07-06T09:59:56.064239Z","shell.execute_reply":"2021-07-06T09:59:56.067647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nTRAIN_DIR = '../input/cassava-leaf-disease-classification/train_images/'\nTEST_DIR = '../input/cassava-leaf-disease-classification/test_images/'","metadata":{"execution":{"iopub.status.busy":"2021-07-06T09:59:54.328294Z","iopub.execute_input":"2021-07-06T09:59:54.330904Z","iopub.status.idle":"2021-07-06T09:59:54.339371Z","shell.execute_reply.started":"2021-07-06T09:59:54.330852Z","shell.execute_reply":"2021-07-06T09:59:54.338591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = json.load(open(\"../input/cassava-leaf-disease-classification/label_num_to_disease_map.json\"))\ntrain = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\nsample = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')\n\nX, Y = train['image_id'].values, train['label'].values\nX_test = [name for name in (os.listdir(TEST_DIR))]","metadata":{"execution":{"iopub.status.busy":"2021-07-06T09:59:58.683640Z","iopub.execute_input":"2021-07-06T09:59:58.684086Z","iopub.status.idle":"2021-07-06T09:59:58.720549Z","shell.execute_reply.started":"2021-07-06T09:59:58.684053Z","shell.execute_reply":"2021-07-06T09:59:58.719691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED=777\n    \nBATCH = 32\nEPOCHS = 10\n\nLR = 0.0001\nIM_SIZE = 256\n\nN_FOLD=5\nTARGET_COL='label'\n    \nMODELNAME='efficientnet-b0'","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:00:00.578164Z","iopub.execute_input":"2021-07-06T10:00:00.578520Z","iopub.status.idle":"2021-07-06T10:00:00.583545Z","shell.execute_reply.started":"2021-07-06T10:00:00.578490Z","shell.execute_reply":"2021-07-06T10:00:00.582658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folds = train.copy()\nFold = StratifiedKFold(n_splits=N_FOLD, shuffle=True, random_state=SEED)\nfor n, (train_index, val_index) in enumerate(Fold.split(folds, folds[TARGET_COL])):\n    folds.loc[val_index, 'fold'] = int(n)\nfolds['fold'] = folds['fold'].astype(int)","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:00:02.503957Z","iopub.execute_input":"2021-07-06T10:00:02.504475Z","iopub.status.idle":"2021-07-06T10:00:02.527363Z","shell.execute_reply.started":"2021-07-06T10:00:02.504442Z","shell.execute_reply":"2021-07-06T10:00:02.526166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainData(Dataset):\n    def __init__(self, Dir, FNames, Labels, Transform):\n        self.dir = Dir\n        self.fnames = FNames\n        self.transform = Transform\n        self.lbs = Labels\n        \n    def __len__(self):\n        return len(self.fnames)\n\n    def __getitem__(self, index):\n        x = Image.open(os.path.join(self.dir, self.fnames[index]))  \n        return self.transform(x), self.lbs[index] \n        \nclass TestData(Dataset):\n    def __init__(self, Dir, FNames, Transform):\n        self.dir = Dir\n        self.fnames = FNames\n        self.transform = Transform\n        \n    def __len__(self):\n        return len(self.fnames)\n\n    def __getitem__(self, index):\n        x = Image.open(os.path.join(self.dir, self.fnames[index]))     \n        return self.transform(x), self.fnames[index]","metadata":{"execution":{"iopub.status.busy":"2021-07-06T09:59:35.696326Z","iopub.execute_input":"2021-07-06T09:59:35.696817Z","iopub.status.idle":"2021-07-06T09:59:35.706152Z","shell.execute_reply.started":"2021-07-06T09:59:35.696783Z","shell.execute_reply":"2021-07-06T09:59:35.705145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = transforms.Compose(\n    [transforms.RandomResizedCrop((IM_SIZE, IM_SIZE), scale=(0.8, 1.0)),\n     transforms.RandomRotation(90),\n     transforms.RandomHorizontalFlip(p=0.5),\n     transforms.RandomVerticalFlip(p=0.5),\n     transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.1, hue=0),\n     transforms.ToTensor(),\n     transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n    ])\n\ntest_transform = transforms.Compose(\n    [transforms.Resize((IM_SIZE, IM_SIZE)),\n     transforms.ToTensor(),\n     transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))])","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:00:05.856635Z","iopub.execute_input":"2021-07-06T10:00:05.857043Z","iopub.status.idle":"2021-07-06T10:00:05.864662Z","shell.execute_reply.started":"2021-07-06T10:00:05.857006Z","shell.execute_reply":"2021-07-06T10:00:05.863805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trn_idx = folds[folds['fold'] != 0].index\nval_idx = folds[folds['fold'] == 0].index\n\nX_train, Y_train = X[trn_idx], Y[trn_idx]\nX_val, Y_val = X[val_idx], Y[val_idx]\n\ntrainset = TrainData(TRAIN_DIR, X_train, Y_train, train_transform)\ntrainloader = DataLoader(trainset,\n                         batch_size=BATCH,\n                         shuffle=True,\n                         num_workers=4)\n\nvalidset = TrainData(TRAIN_DIR, X_val, Y_val, test_transform)\nvalidloader = DataLoader(validset,\n                         batch_size=BATCH,\n                         shuffle=False,\n                         num_workers=4)\n\ntestset = TestData(TEST_DIR, X_test, test_transform)\ntestloader = DataLoader(testset,\n                        batch_size=BATCH,\n                        shuffle=False,\n                        num_workers=4)","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:00:08.544839Z","iopub.execute_input":"2021-07-06T10:00:08.545372Z","iopub.status.idle":"2021-07-06T10:00:08.580452Z","shell.execute_reply.started":"2021-07-06T10:00:08.545330Z","shell.execute_reply":"2021-07-06T10:00:08.579678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from efficientnet_pytorch import EfficientNet\n\nclass enetv2(nn.Module):\n    def __init__(self, out_dim=1, ModelName=\"efficientnet-b0\"):\n        super(enetv2, self).__init__()\n        self.basemodel = EfficientNet.from_name(MODELNAME)\n        for param in self.basemodel.parameters():\n            param.requires_grad = False\n        self.myfc = nn.Linear(self.basemodel._fc.in_features, out_dim)\n        self.basemodel._fc = nn.Identity()        \n            \n    def extract(self, x):\n        return self.basemodel(x)\n\n    def forward(self, x):\n        x = self.basemodel(x)\n        x = self.myfc(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:00:11.090310Z","iopub.execute_input":"2021-07-06T10:00:11.090852Z","iopub.status.idle":"2021-07-06T10:00:11.140050Z","shell.execute_reply.started":"2021-07-06T10:00:11.090802Z","shell.execute_reply":"2021-07-06T10:00:11.139176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = enetv2(5, MODELNAME)\ncheckpoint = torch.load(\"../input/efficientnet-pytorch/efficientnet-b0-08094119.pth\", map_location=DEVICE)\nmodel.load_state_dict(checkpoint, strict=False)\nmodel = model.to(DEVICE)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)\nscheduler = ReduceLROnPlateau(optimizer, mode='min', verbose=True)","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:00:13.603280Z","iopub.execute_input":"2021-07-06T10:00:13.603849Z","iopub.status.idle":"2021-07-06T10:00:14.278130Z","shell.execute_reply.started":"2021-07-06T10:00:13.603806Z","shell.execute_reply":"2021-07-06T10:00:14.277279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, train_loader):\n    model.train()\n    running_loss = 0\n    correct = 0\n    total = 0\n    \n    for batch_idx, (images, labels) in enumerate(train_loader):\n        images = images.to(DEVICE)\n        labels = labels.to(DEVICE)\n        \n        optimizer.zero_grad() \n        outputs = model(images)\n        \n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predict = torch.max(outputs.data, 1)\n        correct += (predict == labels).sum().item()\n        total += labels.size(0)\n        \n    train_loss = running_loss / len(train_loader)\n    train_acc = correct / total\n    \n    return train_loss, train_acc\n\ndef valid(model, valid_loader):\n    model.eval()\n    running_loss = 0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        \n        for batch_idx, (images, labels) in enumerate(valid_loader):\n            images = images.to(DEVICE)\n            labels = labels.to(DEVICE)\n            \n            outputs = model(images)\n            \n            loss = criterion(outputs, labels)\n            running_loss += loss.item()\n            \n            _, predict = torch.max(outputs.data, 1)\n            correct += (predict == labels).sum().item()\n            total += labels.size(0)\n            \n    val_loss = running_loss / len(valid_loader)\n    val_acc = correct / total\n    \n    return val_loss, val_acc","metadata":{"execution":{"iopub.status.busy":"2021-07-05T15:11:23.456215Z","iopub.execute_input":"2021-07-05T15:11:23.456673Z","iopub.status.idle":"2021-07-05T15:11:23.467279Z","shell.execute_reply.started":"2021-07-05T15:11:23.456616Z","shell.execute_reply":"2021-07-05T15:11:23.466541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_score = 0.\n\nfor epoch_idx in range(EPOCHS):\n\n    train_loss, train_acc = train(model, trainloader)\n    valid_loss, valid_acc = valid(model, validloader)\n    \n    # model save\n    if valid_acc > best_score:\n        best_score = valid_acc\n\n        torch.save({'model': model.state_dict()},\n                    OUTPUT_DIR+f'{MODELNAME}_best.pth')\n        \n        print('model saved')\n        \n    # rl scheduler\n    scheduler.step(valid_loss)\n\n    print('Epoch: {} |train_loss: {:.3f} valid loss: {:.3f} train_acc: {:.3f} valid_acc: {:.3f}'.format(epoch_idx, train_loss, valid_loss, train_acc, valid_acc))","metadata":{"execution":{"iopub.status.busy":"2021-07-05T15:11:23.468415Z","iopub.execute_input":"2021-07-05T15:11:23.468886Z","iopub.status.idle":"2021-07-05T22:07:52.897582Z","shell.execute_reply.started":"2021-07-05T15:11:23.468839Z","shell.execute_reply":"2021-07-05T22:07:52.894297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load(OUTPUT_DIR+f'{MODELNAME}_best.pth'), strict=False)\ns_ls = []\n\nwith torch.no_grad():\n    model.eval()\n    for image, fname in testloader: \n        image = image.to(DEVICE)\n        \n        logits = model(image)        \n        ps = torch.exp(logits)        \n        _, top_class = ps.topk(1, dim=1)\n        \n        for i, pred in enumerate(top_class):\n            s_ls.append([fname[i], pred.item()])","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:05:15.948513Z","iopub.execute_input":"2021-07-06T10:05:15.949024Z","iopub.status.idle":"2021-07-06T10:05:16.500397Z","shell.execute_reply.started":"2021-07-06T10:05:15.948992Z","shell.execute_reply":"2021-07-06T10:05:16.499460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.DataFrame.from_records(s_ls, columns=['image_id', 'label'])\nsub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2021-07-06T10:05:19.756007Z","iopub.execute_input":"2021-07-06T10:05:19.756597Z","iopub.status.idle":"2021-07-06T10:05:19.770287Z","shell.execute_reply.started":"2021-07-06T10:05:19.756546Z","shell.execute_reply":"2021-07-06T10:05:19.769339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.head()","metadata":{"execution":{"iopub.status.busy":"2021-07-05T22:07:53.440379Z","iopub.execute_input":"2021-07-05T22:07:53.440881Z","iopub.status.idle":"2021-07-05T22:07:53.487285Z","shell.execute_reply.started":"2021-07-05T22:07:53.440830Z","shell.execute_reply":"2021-07-05T22:07:53.486134Z"},"trusted":true},"execution_count":null,"outputs":[]}]}