{"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":"from torch.utils.data import Dataset, DataLoader\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torchvision import models\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport cv2\nimport pandas as pd\n\nfrom tqdm import tqdm\n\nimport glob\nimport os\nimport time\nfrom IPython import display as ipd\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix\nimport seaborn as sn\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport wandb\nfrom torchvision import datasets, models, transforms\n","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:30.113821Z","iopub.execute_input":"2022-07-26T14:49:30.114219Z","iopub.status.idle":"2022-07-26T14:49:34.383326Z","shell.execute_reply.started":"2022-07-26T14:49:30.114134Z","shell.execute_reply":"2022-07-26T14:49:34.382213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"CFG = {\n    \"batch_size\" : 32,\n    \"num_epochs\" : 40,\n    'pretrained' : True,\n    'init_lr' : 0.001,\n    'weight_decay' : 0,\n    'device' : torch.device('cuda' if torch.cuda.is_available() else 'cpu'),\n    'model_name' : 'resnet34',\n    'min_lr' : 0.001,\n    'max_lr' : 0.01,\n    'patience' : 0,\n    'gamma' : 0.1,\n    'momentum' : 0.9,\n    'optimizer' : 'SGD',\n    'lr_scheduler' : 'OneCycleLR',\n}","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:49:42.168480Z","iopub.execute_input":"2022-07-26T14:49:42.168832Z","iopub.status.idle":"2022-07-26T14:49:42.242552Z","shell.execute_reply.started":"2022-07-26T14:49:42.168804Z","shell.execute_reply":"2022-07-26T14:49:42.241410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Set up data","metadata":{}},{"cell_type":"code","source":"def get_train_val(train_path):\n    labels = os.listdir(train_path)\n    train_data = []\n    val_data = []\n    for label in labels:\n        images_path = glob.glob(f'{train_path}/{label}/*.jpg')\n        images_path = [path.replace(\"\\\\\", \"/\") for path in images_path]\n\n        train_paths, val_paths = train_test_split(images_path, test_size=0.2, random_state=43)#42\n        val_data += val_paths\n        train_data += train_paths\n    return train_data, val_data\n\ntrain_path = '../input/paddy-disease-classification/train_images'\ntrain_data, val_data = get_train_val(train_path)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:51:24.405065Z","iopub.execute_input":"2022-07-26T14:51:24.405410Z","iopub.status.idle":"2022-07-26T14:51:25.205715Z","shell.execute_reply.started":"2022-07-26T14:51:24.405382Z","shell.execute_reply":"2022-07-26T14:51:25.204180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data augmentation\nUse Albumentations library for data augmentation","metadata":{}},{"cell_type":"code","source":"train_transforms = A.Compose([\n    A.OneOf([\n        A.Rotate(30, p=1),\n        A.Rotate(-30, p=1),\n        A.HorizontalFlip(p=1),\n        A.VerticalFlip(p=1),\n        # A.CenterCrop(height=480,width=480,p=1),\n        # A.Blur(p=1),\n        A.ColorJitter(brightness=0.05, contrast=0.05, saturation=0.05, hue=0, p=1),\n        # A.RandomShadow(),\n        A.RandomBrightnessContrast(p=1),\n\n    ], p=0.8),\n    A.Resize(224, 224),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2(),\n])\ntrain_transforms_base = A.Compose([\n    A.Resize(480, 480),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2(),\n])\n\nval_transforms = A.Compose([\n    A.Resize(224, 224),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2(),\n])","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:50:10.562286Z","iopub.execute_input":"2022-07-26T14:50:10.562649Z","iopub.status.idle":"2022-07-26T14:50:10.572657Z","shell.execute_reply.started":"2022-07-26T14:50:10.562619Z","shell.execute_reply":"2022-07-26T14:50:10.571656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"labels = {\n    'bacterial_leaf_blight': 0,\n    'bacterial_leaf_streak': 1,\n    'bacterial_panicle_blight': 2,\n    'blast': 3,\n    'brown_spot': 4,\n    'dead_heart': 5,\n    'downy_mildew': 6,\n    'hispa': 7,\n    'normal': 8,\n    'tungro': 9\n    \n}\nnum_classes = len(labels.keys())\none_hot_encoding = F.one_hot(torch.arange(0, num_classes) % num_classes, num_classes=num_classes)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:50:29.886033Z","iopub.execute_input":"2022-07-26T14:50:29.886595Z","iopub.status.idle":"2022-07-26T14:50:29.917995Z","shell.execute_reply.started":"2022-07-26T14:50:29.886553Z","shell.execute_reply":"2022-07-26T14:50:29.916561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I built my own custom Pytorch dataset, but you can opt to use ImageFolder dataset class from Pytorch, which will do the same for you and will be much easier to use.","metadata":{}},{"cell_type":"code","source":"class PaddyDiseaseClassificationDataset(Dataset):\n    def __init__(self, data, dataset_name='', transforms=None, albumentations_transform = True):\n        self.image_paths = data\n        self.transforms = transforms\n        self.name = dataset_name\n        self.distribution = self.get_distribution()\n        self.albumentations_transform = albumentations_transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        \n        img = cv2.imread(img_path)[:, :, ::-1]  # convert it to rgb\n        img = img.astype('float32')\n        #img /= 255  # scale img to [0, 1]\n        \n\n        # print(img_path.split(\"/\")[-2])\n        # print(img_path)\n        label = labels[img_path.split(\"/\")[-2]]\n        if self.transforms is not None:\n            if label == 3 or label ==7:\n                    transform = self.transforms[0]\n            else:\n                transform = self.transforms[-1]\n\n            if self.albumentations_transform == True:\n                img = transform(image=img)['image']\n            else:\n                img1 = cv2.imread(img_path)\n                img = Image.fromarray(img1)\n                img = transform(img)\n\n\n        encoded_label = one_hot_encoding[label]\n        encoded_label = encoded_label.type(torch.FloatTensor)\n\n        return img, encoded_label\n\n    def get_distribution(self):\n        distribution = {}\n        splitted_paths = [path.split(\"/\") for path in self.image_paths]\n        for splitted_path in splitted_paths:\n            if splitted_path[-2] not in distribution:\n                distribution[splitted_path[-2]] = 0\n            else:\n                distribution[splitted_path[-2]] += 1\n        print(f\"Distribution of {self.name} dataset: {distribution}\")\n        return distribution","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:50:39.940326Z","iopub.execute_input":"2022-07-26T14:50:39.941253Z","iopub.status.idle":"2022-07-26T14:50:39.954894Z","shell.execute_reply.started":"2022-07-26T14:50:39.941215Z","shell.execute_reply":"2022-07-26T14:50:39.953837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = PaddyDiseaseClassificationDataset(train_data, dataset_name='train', transforms=[train_transforms], albumentations_transform = True)#The base always the last one\nval_dataset = PaddyDiseaseClassificationDataset(val_data, dataset_name='validation', transforms=[val_transforms], albumentations_transform = True)\ntrain_dl = DataLoader(train_dataset, batch_size=CFG['batch_size'], shuffle=True)\nval_dl = DataLoader(val_dataset, batch_size=CFG['batch_size'], shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:51:39.414647Z","iopub.execute_input":"2022-07-26T14:51:39.415344Z","iopub.status.idle":"2022-07-26T14:51:39.450422Z","shell.execute_reply.started":"2022-07-26T14:51:39.415299Z","shell.execute_reply":"2022-07-26T14:51:39.449459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"from torch import nn\nimport torch\nfrom torchvision import models\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, OneCycleLR\n\nclass CustomModel(nn.Module):\n    def __init__(self, num_classes, model_name, pretrained=True):\n        super(CustomModel, self).__init__()\n        if model_name == 'efficientnet_b1':\n            self.model = models.efficientnet_b1(pretrained=pretrained)\n            in_features = self.model.classifier[1].in_features\n            self.model.classifier[1] = nn.Linear(in_features=in_features, out_features=num_classes, bias=True)\n        elif model_name == 'resnet34':\n            self.model = models.resnet34(pretrained=pretrained)\n            in_features = self.model.fc.in_features\n            self.model.fc = nn.Linear(in_features=in_features, out_features=num_classes, bias=True)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:52:05.588603Z","iopub.execute_input":"2022-07-26T14:52:05.589500Z","iopub.status.idle":"2022-07-26T14:52:05.598354Z","shell.execute_reply.started":"2022-07-26T14:52:05.589444Z","shell.execute_reply":"2022-07-26T14:52:05.597420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CustomModel(num_classes, 'resnet34', pretrained=CFG['pretrained'])\nmodel = model.to(CFG['device'])\n","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:52:14.326855Z","iopub.execute_input":"2022-07-26T14:52:14.327194Z","iopub.status.idle":"2022-07-26T14:52:20.551060Z","shell.execute_reply.started":"2022-07-26T14:52:14.327165Z","shell.execute_reply":"2022-07-26T14:52:20.550031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimizer","metadata":{}},{"cell_type":"code","source":"if CFG['optimizer'] == 'Adam':\n    optimizer = optim.Adam(model.parameters(), lr=CFG['init_lr'], weight_decay=CFG['weight_decay'])\nelif CFG['optimizer'] == 'SGD':\n    optimizer = optim.SGD(model.parameters(), lr=CFG['init_lr'], momentum=CFG['momentum'], weight_decay=CFG['weight_decay'], nesterov=True)\n\nif CFG['lr_scheduler'] == None:\n    scheduler = None\nelif CFG['lr_scheduler'] == 'OneCycleLR':\n    scheduler = OneCycleLR(optimizer, max_lr=CFG['max_lr'], steps_per_epoch=len(train_dl), epochs=CFG['num_epochs'])\nelif CFG['lr_scheduler'] == 'ReduceLROnPlateau':\n    scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG['gamma'], patience=CFG['patience'], min_lr=CFG['min_lr'])\n","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:52:32.388095Z","iopub.execute_input":"2022-07-26T14:52:32.388446Z","iopub.status.idle":"2022-07-26T14:52:32.397185Z","shell.execute_reply.started":"2022-07-26T14:52:32.388418Z","shell.execute_reply":"2022-07-26T14:52:32.396260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss function","metadata":{}},{"cell_type":"code","source":"def criterion(inputs, targets):\n    loss = F.cross_entropy(inputs, targets)\n    return loss","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:52:36.764179Z","iopub.execute_input":"2022-07-26T14:52:36.764551Z","iopub.status.idle":"2022-07-26T14:52:36.769096Z","shell.execute_reply.started":"2022-07-26T14:52:36.764519Z","shell.execute_reply":"2022-07-26T14:52:36.768207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, dataloader, optimizer, epoch):\n    model.train()\n    running_loss = 0.0\n    dataset_size = 0\n    # use tqdm to track progress\n    with tqdm(dataloader, unit=\"batch\") as tepoch:\n        tepoch.set_description(f\"Epoch {epoch} train\")\n        # Iterate over data.\n        for inputs, targets in tepoch:\n            inputs = inputs.to(CFG['device'])\n            targets = targets.to(CFG['device'])\n            # zero the parameter gradients\n            optimizer.zero_grad()\n            # forward\n            outputs = model(inputs)\n            # loss\n            loss = criterion(outputs, targets)\n            # backward\n            loss.backward()\n            optimizer.step()\n            # calculate epoch loss\n            dataset_size += inputs.size(0)\n            running_loss += loss.item() * inputs.size(0)\n            epoch_loss = running_loss / dataset_size\n            # get current learning rate\n            current_lr = optimizer.param_groups[0]['lr']\n            # print statistics\n            tepoch.set_postfix(loss=epoch_loss, lr=current_lr)\n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:54:41.353962Z","iopub.execute_input":"2022-07-26T14:54:41.354364Z","iopub.status.idle":"2022-07-26T14:54:41.376941Z","shell.execute_reply.started":"2022-07-26T14:54:41.354332Z","shell.execute_reply":"2022-07-26T14:54:41.375411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"since = time.time()\nbest_loss = 100000\nfor epoch in range(CFG['num_epochs']):\n    loss = train_epoch(model, train_dl, optimizer, epoch)\n    # save best model\n    if loss < best_loss:\n        best_loss = loss\n        torch.save(model.state_dict(), \"./best.pt\")\ntime_elapsed = time.time() - since\nprint('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-26T14:54:42.566514Z","iopub.execute_input":"2022-07-26T14:54:42.566886Z","iopub.status.idle":"2022-07-26T16:05:44.904183Z","shell.execute_reply.started":"2022-07-26T14:54:42.566857Z","shell.execute_reply":"2022-07-26T16:05:44.902501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission\nLoad best trained model","metadata":{}},{"cell_type":"code","source":"model_path = \"./best.pt\"\nmodel = CustomModel(num_classes=num_classes, model_name=CFG['model_name'], pretrained=False)\nmodel.load_state_dict(torch.load(model_path))\nmodel = model.to(CFG['device'])\nmodel.eval()\nprint(\"Model ready\")","metadata":{"execution":{"iopub.status.busy":"2022-07-26T16:05:48.511408Z","iopub.execute_input":"2022-07-26T16:05:48.511821Z","iopub.status.idle":"2022-07-26T16:05:49.082774Z","shell.execute_reply.started":"2022-07-26T16:05:48.511789Z","shell.execute_reply":"2022-07-26T16:05:49.081507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"reverse_labels = dict((v, k) for k, v in labels.items())\nimages_path = glob.glob('/kaggle/input/paddy-disease-classification/test_images/*.jpg')\n\nsubmission = []\nfor img_path in tqdm(images_path):\n    # process image\n    img = cv2.imread(img_path)[:, :, ::-1]  # convert it to rgb\n    img = img.astype('float32')\n    img = val_transforms(image=img)['image'] # apply same transforms of validation set\n    img = img[None, ...].to(CFG['device']) # add batch dimension to image and use device\n    # predict \n    pred = model(img) \n    pred = torch.max(pred, dim=1)[1]\n    label = reverse_labels[pred.item()]\n    submission.append([img_path.split(\"/\")[-1], label])\nsubmission = pd.DataFrame(submission, columns=['image_id', 'label'])\nsubmission.to_csv(\"./submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T16:05:49.748239Z","iopub.execute_input":"2022-07-26T16:05:49.749019Z","iopub.status.idle":"2022-07-26T16:07:16.281247Z","shell.execute_reply.started":"2022-07-26T16:05:49.748958Z","shell.execute_reply":"2022-07-26T16:07:16.280174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}