{"cells":[{"metadata":{"_uuid":"12d546ab-f82c-49ea-b282-0ebc373fbe53","_cell_guid":"2bfbde8e-6c8a-48d4-8ab0-33bfb7e12be6","trusted":true},"cell_type":"markdown","source":"### Library import"},{"metadata":{"_uuid":"f99acd7c-4ef5-4773-8bda-0f59ee783601","_cell_guid":"c8123ddc-070c-40d3-ac9b-fd0581380ffb","trusted":true},"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","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import sys\n\npackage_path = '../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master'\nsys.path.append(package_path)\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"7e723b78-5a12-487f-8dc5-d85ba9be934e","_cell_guid":"fb0a4a8e-887a-4220-ab9e-f54956e4a0d5","trusted":true},"cell_type":"markdown","source":"### Directory"},{"metadata":{"_uuid":"43c7d8b3-6d00-468a-9e0a-3c45531ee1ff","_cell_guid":"8dd75c2e-5275-4d6b-ad27-9b1d4bc070c8","trusted":true},"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/'","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"1d008f0b-5471-4d4e-97b5-a436350fd24b","_cell_guid":"ce997916-289b-4fa4-80f5-28bb33c55ead","trusted":true},"cell_type":"markdown","source":"### Data loading"},{"metadata":{"_uuid":"ac02b3da-2a7e-4f66-8069-9dae31b13118","_cell_guid":"1d8d72cb-8d57-43ae-a59e-abb74a653f44","trusted":true},"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))]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train['label'].value_counts()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"292353bf-d7a8-4c01-9487-1ebebac849ef","_cell_guid":"bc3510da-a896-4690-8e1c-6438d7c511f1","trusted":true},"cell_type":"markdown","source":"### Config"},{"metadata":{"_uuid":"9a031aae-b9c7-4d18-aa42-2fdea1863114","_cell_guid":"69713d26-5161-4d88-82d5-d8f4022f2974","trusted":true},"cell_type":"code","source":"class Conf:\n    seed=777\n    \n    BATCH = 32\n    EPOCHS = 10\n\n    LR = 0.0001\n    IM_SIZE = 256\n\n    n_fold=5\n    target_col='label'\n    \n    modelname='efficientnet-b0'","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"2d74feb6-c5b8-4576-babf-dcbd7cff0fec","_cell_guid":"9963df50-ee8c-4788-9b5a-b6480e2df6a7","trusted":true},"cell_type":"markdown","source":"### CV split"},{"metadata":{"_uuid":"47118400-5e44-4e88-bec7-fa8549a0d4ba","_cell_guid":"f7f05753-51bc-4287-ac84-e3ab255cb6fa","trusted":true},"cell_type":"code","source":"folds = train.copy()\nFold = StratifiedKFold(n_splits=Conf.n_fold, shuffle=True, random_state=Conf.seed)\nfor n, (train_index, val_index) in enumerate(Fold.split(folds, folds[Conf.target_col])):\n    folds.loc[val_index, 'fold'] = int(n)\nfolds['fold'] = folds['fold'].astype(int)\nprint(folds.groupby(['fold', Conf.target_col]).size())","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"15ec2dd5-c739-4e34-86ae-a4c02a6d9055","_cell_guid":"0c6e03d6-cee4-4db9-9dc0-a8f9e562c793","trusted":true},"cell_type":"markdown","source":"### Dataset"},{"metadata":{"_uuid":"872ab572-903a-493a-9f13-47356e72dea7","_cell_guid":"4d5828a3-9c3e-4760-a439-68f8b1e6e5b6","trusted":true},"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]","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e9f72b3a-01fc-4274-9f30-00207f8efdb8","_cell_guid":"a38d9685-b054-4fe0-b0f1-d6e5e6ba17b9","trusted":true},"cell_type":"markdown","source":"### Augmentation"},{"metadata":{"_uuid":"d83f9d02-9a87-4d92-b7f0-36976ecf02b2","_cell_guid":"981bee20-627c-4fbd-85c0-d8a9773084d4","trusted":true},"cell_type":"code","source":"train_transform = transforms.Compose(\n    [transforms.RandomResizedCrop((Conf.IM_SIZE, Conf.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((Conf.IM_SIZE, Conf.IM_SIZE)),\n     transforms.ToTensor(),\n     transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))])","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"bd1edb8c-1818-4d9a-b6d5-0976bb08c886","_cell_guid":"ce1cdc28-1823-4910-9965-672d96d685f8","trusted":true},"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=Conf.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=Conf.BATCH,\n                         shuffle=False,\n                         num_workers=4)\n\ntestset = TestData(TEST_DIR, X_test, test_transform)\ntestloader = DataLoader(testset,\n                        batch_size=Conf.BATCH,\n                        shuffle=False,\n                        num_workers=4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"Y_train","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"3a93a5d2-56a8-4b63-b90a-d51cddba1b89","_cell_guid":"94c0f3fe-0b8d-4b68-86c5-a8ccb895adb9","trusted":true},"cell_type":"markdown","source":"### Model"},{"metadata":{"_uuid":"6331756c-e445-4790-8ab3-172f9258cea4","_cell_guid":"2c0986a4-8954-43c2-bcad-01f698c65cb3","trusted":true},"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(Conf.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","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = enetv2(5, Conf.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=Conf.LR)\nscheduler = ReduceLROnPlateau(optimizer, mode='min', verbose=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for param in model.parameters():\n    print(param)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"ffad6512-da06-4032-bd80-abed87df0153","_cell_guid":"cde2ce31-1bf1-4135-ae50-9b4066744b95","trusted":true},"cell_type":"markdown","source":"# Training"},{"metadata":{"trusted":true},"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","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"best_score = 0.\n\nfor epoch_idx in range(Conf.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'{Conf.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))","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"4a1427ef-dc5e-44b6-9842-7d2f47716d7c","_cell_guid":"e52056ff-1b34-4d1a-855c-4f9371f66253","trusted":true},"cell_type":"code","source":"model.load_state_dict(torch.load(OUTPUT_DIR+f'{Conf.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()])","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"f9559f2d-dad1-4082-969d-03fd91b9c7ed","_cell_guid":"ce51405e-96f7-46ff-8b91-7849dde2ce95","trusted":true},"cell_type":"code","source":"sub = pd.DataFrame.from_records(s_ls, columns=['image_id', 'label'])\nsub.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sub.head()","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}