{"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 os\nimport random\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nimport torchvision.models as models\nfrom torchvision import transforms\nfrom PIL import Image\nfrom PIL import Image\nimport torch\nfrom torchvision import datasets, models, transforms\nimport torch.nn as nn\nfrom torch.nn import functional as F\nimport torch.optim as optim\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-07-17T11:04:29.616178Z","iopub.execute_input":"2021-07-17T11:04:29.616535Z","iopub.status.idle":"2021-07-17T11:04:29.624532Z","shell.execute_reply.started":"2021-07-17T11:04:29.616503Z","shell.execute_reply":"2021-07-17T11:04:29.622157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATH = '../input/cassava-leaf-disease-classification'\ntrain_csv = f'{PATH}/train.csv'\ntrain_dir = f'{PATH}/train_images/'\ntest_dir = f'{PATH}/test_images/'\ntrain_df = pd.read_csv(train_csv)\ntrain_df.head()\n\n\ntrain_bs = 16\nvalid_bs = 32\nimg_size = 224","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.189394Z","iopub.execute_input":"2021-07-17T11:04:30.189744Z","iopub.status.idle":"2021-07-17T11:04:30.213973Z","shell.execute_reply.started":"2021-07-17T11:04:30.189713Z","shell.execute_reply":"2021-07-17T11:04:30.213209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_df, Y_df = train_df['image_id'].values, train_df['label'].values\nX_test = os.listdir(test_dir)\n","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.215617Z","iopub.execute_input":"2021-07-17T11:04:30.215967Z","iopub.status.idle":"2021-07-17T11:04:30.222782Z","shell.execute_reply.started":"2021-07-17T11:04:30.215931Z","shell.execute_reply":"2021-07-17T11:04:30.22201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\nX_train, X_valid, Y_train, Y_valid = train_test_split(X_df, Y_df, test_size = 0.25, shuffle = True)","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.224544Z","iopub.execute_input":"2021-07-17T11:04:30.225087Z","iopub.status.idle":"2021-07-17T11:04:30.232834Z","shell.execute_reply.started":"2021-07-17T11:04:30.225049Z","shell.execute_reply":"2021-07-17T11:04:30.231972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, Dir, Files, Labels, Transform):\n        self.dir = Dir\n        self.files = Files\n        self.labels = Labels\n        self.transform = Transform\n\n    def __len__(self):\n        return len(self.files)\n\n    def __getitem__(self, idx):\n        img = Image.open(os.path.join(self.dir, self.files[idx]))\n        if \"train\" in self.dir:\n            return self.transform(img), self.labels[idx]\n        if \"test\" in self.dir:\n            return self.transform(img), self.files[idx]","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.234976Z","iopub.execute_input":"2021-07-17T11:04:30.235216Z","iopub.status.idle":"2021-07-17T11:04:30.245166Z","shell.execute_reply.started":"2021-07-17T11:04:30.235186Z","shell.execute_reply":"2021-07-17T11:04:30.244378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndata_transforms = {\n    'train':\n    transforms.Compose([\n        transforms.Resize((224,224)),\n        transforms.RandomAffine(0, shear=10, scale=(0.8,1.2)),\n        transforms.RandomHorizontalFlip(),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                 std=[0.229, 0.224, 0.225])\n    ]),\n    'validation':\n    transforms.Compose([\n        transforms.Resize((224,224)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                 std=[0.229, 0.224, 0.225])\n    ]),\n}\n","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.246339Z","iopub.execute_input":"2021-07-17T11:04:30.246701Z","iopub.status.idle":"2021-07-17T11:04:30.254049Z","shell.execute_reply.started":"2021-07-17T11:04:30.246668Z","shell.execute_reply":"2021-07-17T11:04:30.252934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_data = CassavaDataset(train_dir, X_train, Y_train, Transform)\n# valid_data = CassavaDataset(train_dir, X_valid, Y_valid, Transform)\n# test_data = CassavaDataset(test_dir, X_test, None, Transform)\n\n# train_loader = DataLoader(train_data, batch_size = train_bs, shuffle = True, num_workers = 4)\n# valid_loader = DataLoader(valid_data, batch_size = valid_bs, shuffle = True, num_workers = 4)\n# test_loader = DataLoader(test_data, batch_size = 1, shuffle = True, num_workers = 4)","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.255326Z","iopub.execute_input":"2021-07-17T11:04:30.255692Z","iopub.status.idle":"2021-07-17T11:04:30.264366Z","shell.execute_reply.started":"2021-07-17T11:04:30.255659Z","shell.execute_reply":"2021-07-17T11:04:30.263538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n\nimage_datasets = {\n    'train': \n        CassavaDataset(train_dir, X_train, Y_train, data_transforms['train']),\n    'validation': \n        CassavaDataset(train_dir, X_valid, Y_valid, data_transforms['validation'])\n}\ndataloaders = {\n    'train':\n        DataLoader(image_datasets['train'], batch_size = train_bs, shuffle = True, num_workers = 4),\n    'validation':\n        DataLoader(image_datasets['validation'], batch_size = valid_bs, shuffle = True, num_workers = 4)\n}\n\n","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.306499Z","iopub.execute_input":"2021-07-17T11:04:30.306731Z","iopub.status.idle":"2021-07-17T11:04:30.313596Z","shell.execute_reply.started":"2021-07-17T11:04:30.306709Z","shell.execute_reply":"2021-07-17T11:04:30.312706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = models.resnet50(pretrained=True).to(device)\n    \nfor param in model.parameters():\n    param.requires_grad = False   \n    \nmodel.fc = nn.Sequential(\n               nn.Linear(2048, 5, bias = True))\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.316147Z","iopub.execute_input":"2021-07-17T11:04:30.316692Z","iopub.status.idle":"2021-07-17T11:04:30.994726Z","shell.execute_reply.started":"2021-07-17T11:04:30.316658Z","shell.execute_reply":"2021-07-17T11:04:30.993886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.fc.parameters())","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:30.996148Z","iopub.execute_input":"2021-07-17T11:04:30.996502Z","iopub.status.idle":"2021-07-17T11:04:31.000613Z","shell.execute_reply.started":"2021-07-17T11:04:30.996465Z","shell.execute_reply":"2021-07-17T11:04:30.999812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, criterion, optimizer, num_epochs=30):\n    for epoch in range(num_epochs):\n        print('Epoch {}/{}'.format(epoch+1, num_epochs))\n        print('-' * 10)\n\n        for phase in ['train', 'validation']:\n            if phase == 'train':\n                model.train()\n            else:\n                model.eval()\n\n            running_loss = 0.0\n            running_corrects = 0\n\n            for inputs, labels in dataloaders[phase]:\n                inputs = inputs.to(device)\n                labels = labels.to(device) \n\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n                if phase == 'train':\n                    optimizer.zero_grad()\n                    loss.backward()\n                    optimizer.step()\n\n                _, preds = torch.max(outputs, 1)\n                running_loss += loss.item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels.data)\n\n            epoch_loss = running_loss / len(image_datasets[phase])\n            epoch_acc = running_corrects.double() / len(image_datasets[phase])\n            \n            print('{} loss: {:.4f}, acc: {:.4f}'.format(phase,\n                                                      epoch_loss,\n                                                        epoch_acc))\n    return model","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:31.002208Z","iopub.execute_input":"2021-07-17T11:04:31.002706Z","iopub.status.idle":"2021-07-17T11:04:31.01286Z","shell.execute_reply.started":"2021-07-17T11:04:31.00267Z","shell.execute_reply":"2021-07-17T11:04:31.012154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_trained = train_model(model, criterion, optimizer, num_epochs=30)","metadata":{"execution":{"iopub.status.busy":"2021-07-17T11:04:31.01425Z","iopub.execute_input":"2021-07-17T11:04:31.014735Z","iopub.status.idle":"2021-07-17T13:01:15.822888Z","shell.execute_reply.started":"2021-07-17T11:04:31.014699Z","shell.execute_reply":"2021-07-17T13:01:15.821693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}