{"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":"markdown","source":"## 1. Import necessary libraries","metadata":{}},{"cell_type":"code","source":"#necessary imports\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pickle\nimport os\nimport json\nfrom sklearn.model_selection import StratifiedKFold\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torchvision import transforms\nimport torchvision.models as models","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:26:34.596431Z","iopub.execute_input":"2022-11-09T15:26:34.596925Z","iopub.status.idle":"2022-11-09T15:26:34.603602Z","shell.execute_reply.started":"2022-11-09T15:26:34.596894Z","shell.execute_reply":"2022-11-09T15:26:34.602047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Read in the data and check if there is any class imbalance","metadata":{}},{"cell_type":"code","source":"#base path\nBASE_PATH = '../input/cassava-leaf-disease-classification/'","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:26:40.091909Z","iopub.execute_input":"2022-11-09T15:26:40.092665Z","iopub.status.idle":"2022-11-09T15:26:40.101897Z","shell.execute_reply.started":"2022-11-09T15:26:40.092621Z","shell.execute_reply":"2022-11-09T15:26:40.100779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#label_dict holds the class encoding\nlabel_dict = json.load(open(f'{BASE_PATH}label_num_to_disease_map.json'))\n#train_ids hold the image names and the classes they belong to\ntrain_ids = pd.read_csv(f'{BASE_PATH}train.csv')\ntest_ids = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:26:45.538441Z","iopub.execute_input":"2022-11-09T15:26:45.538806Z","iopub.status.idle":"2022-11-09T15:26:45.564586Z","shell.execute_reply.started":"2022-11-09T15:26:45.538768Z","shell.execute_reply":"2022-11-09T15:26:45.563637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#plot the distribution of the labels\ntrain_ids['label'].value_counts().plot(kind='bar');","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:26:48.681518Z","iopub.execute_input":"2022-11-09T15:26:48.681911Z","iopub.status.idle":"2022-11-09T15:26:48.903937Z","shell.execute_reply.started":"2022-11-09T15:26:48.681875Z","shell.execute_reply":"2022-11-09T15:26:48.902998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### There is class imbalance in the dataset, so we need to take care of it. We can use the stratified sampling method to balance the dataset.","metadata":{}},{"cell_type":"markdown","source":"## 3. Create Custom_Dataset and apply transformation to images","metadata":{}},{"cell_type":"code","source":"#apply transformations to the images\ntransform = transforms.Compose([\n                transforms.ToPILImage(),\n                transforms.Resize((224, 224)),\n                transforms.ToTensor(),\n                transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))\n            ])","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:26:53.134654Z","iopub.execute_input":"2022-11-09T15:26:53.135038Z","iopub.status.idle":"2022-11-09T15:26:53.14135Z","shell.execute_reply.started":"2022-11-09T15:26:53.135005Z","shell.execute_reply":"2022-11-09T15:26:53.13988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Define custom image dataset\nclass Custom_Dataset(Dataset):\n    def __init__(self, df, path, transform=None):\n        self.df = df\n        self.path = path\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_name, img_label = self.df.iloc[idx]\n        img_label = torch.tensor(img_label)\n        # print(img_label)\n        # print(img_name)\n        img_path = os.path.join(self.path, img_name)\n        #print(img_path)\n        img = plt.imread(img_path)\n        if self.transform:\n            img = self.transform(img)\n        return img, img_label","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:26:57.126772Z","iopub.execute_input":"2022-11-09T15:26:57.127159Z","iopub.status.idle":"2022-11-09T15:26:57.133981Z","shell.execute_reply.started":"2022-11-09T15:26:57.127127Z","shell.execute_reply":"2022-11-09T15:26:57.132909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Create image dataset using Custom_Dataset and show some images","metadata":{}},{"cell_type":"code","source":"train_img_path = f'{BASE_PATH}train_images/'\ntest_img_path = f'{BASE_PATH}test_images/'\n\ndf_train = train_ids.copy()\ndf_test = test_ids.copy()\ntrain_dataset = Custom_Dataset(df_train, train_img_path, transform=transform)\ntest_dataset = Custom_Dataset(df_test, test_img_path, transform=transform)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:27:01.462577Z","iopub.execute_input":"2022-11-09T15:27:01.46349Z","iopub.status.idle":"2022-11-09T15:27:01.473604Z","shell.execute_reply.started":"2022-11-09T15:27:01.463453Z","shell.execute_reply":"2022-11-09T15:27:01.472454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#show some images\ndef show_images(dataset):\n    #subplot 8 images\n    fig,ax = plt.subplots(2,4, figsize=(16,8))\n    for i in range(2):\n        for j in range(4):\n            idx = i*4+j\n            ax[i,j].imshow(dataset[idx][0].permute(1,2,0).detach().numpy())\n            for key, value in label_dict.items():\n                if int(key) == dataset[idx][1]:\n                    ax[i,j].set_title(value)\n                    break\n            ax[i,j].axis('off')\n    plt.show();","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:27:05.638469Z","iopub.execute_input":"2022-11-09T15:27:05.638829Z","iopub.status.idle":"2022-11-09T15:27:05.646886Z","shell.execute_reply.started":"2022-11-09T15:27:05.638799Z","shell.execute_reply":"2022-11-09T15:27:05.645506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_images(train_dataset)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:27:09.621408Z","iopub.execute_input":"2022-11-09T15:27:09.621758Z","iopub.status.idle":"2022-11-09T15:27:10.848353Z","shell.execute_reply.started":"2022-11-09T15:27:09.621728Z","shell.execute_reply":"2022-11-09T15:27:10.847324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. Define/load model, optimizer, loss function, and set the device to train on","metadata":{}},{"cell_type":"code","source":"#define the device to use\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:27:17.902173Z","iopub.execute_input":"2022-11-09T15:27:17.903417Z","iopub.status.idle":"2022-11-09T15:27:17.978566Z","shell.execute_reply.started":"2022-11-09T15:27:17.90337Z","shell.execute_reply":"2022-11-09T15:27:17.977186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#load Resnet34\nmodel = models.resnet34(pretrained=True)\n#change last layer\nnum_ftrs = model.fc.in_features #number of input features to the last layer\n#len(label_dict) = 5 -> 5 classes\nmodel.fc = nn.Linear(in_features = num_ftrs, out_features = len(label_dict)) \nmodel = model.to(device) #send the model to the device","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:27:20.865128Z","iopub.execute_input":"2022-11-09T15:27:20.865495Z","iopub.status.idle":"2022-11-09T15:27:29.341414Z","shell.execute_reply.started":"2022-11-09T15:27:20.865466Z","shell.execute_reply":"2022-11-09T15:27:29.340268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#loss and optimizer\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:27:34.141676Z","iopub.execute_input":"2022-11-09T15:27:34.142038Z","iopub.status.idle":"2022-11-09T15:27:34.150914Z","shell.execute_reply.started":"2022-11-09T15:27:34.142006Z","shell.execute_reply":"2022-11-09T15:27:34.147533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6. Define the training loop","metadata":{}},{"cell_type":"code","source":"#define the training function\ndef train(model, train_loader, val_loader, criterion, optimizer, num_epochs=10):\n    #train the model\n    for epoch in range(num_epochs):\n        for i, (images, labels) in enumerate(train_loader):\n            images = images.to(device)\n            labels = labels.to(device)\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            if (i+1) % 50 == 0:\n                print(f'Epoch [{epoch+1}/{num_epochs}], Step [{i+1}/{len(train_loader)}], Loss: {loss.item():.4f}')\n    #evaluate the model\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = model(images)\n            _, predicted = torch.max(outputs.data, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n    print(f'Accuracy of the model on the {len(val_ds)} val_set: {100 * correct / total:.2f}%')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 7. Use the Stratified K-fold cross validation and train the model. ","metadata":{}},{"cell_type":"markdown","source":"### We will use straticified sampling method and there will be n folds in total. Using each fold we will create a training and testing datasets and the data loaders,","metadata":{}},{"cell_type":"code","source":"#data loader\nbatch_size = 128\nn_splits = 5\nnum_epochs = 10\n#stratified kfold cross validation\nskf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=42)\nfor fold, (train_index, val_index) in enumerate(skf.split(df_train['image_id'], df_train['label'])):\n    train_ds = Subset(train_dataset, train_index) #this will select only the train_index indices from the whole dataset (train_dataset)\n    val_ds = Subset(train_dataset, val_index)\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=False)\n    val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False)   \n    #train the model\n    print(f'Model training on fold: {fold+1}')\n    train(model, train_loader, val_loader, criterion, optimizer, num_epochs=num_epochs)\n","metadata":{"execution":{"iopub.status.busy":"2022-11-09T15:27:43.960316Z","iopub.execute_input":"2022-11-09T15:27:43.96071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#save model as pickle\nwith open('model.pkl', 'wb') as f:\n    pickle.dump(model, f)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 8. Use to model to test and submit","metadata":{}},{"cell_type":"code","source":"test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#test the model\nfor images, labels in test_loader:\n    images = images.to(device)\n    outputs = model(images)\n    _, predicted = torch.max(outputs.data, 1)\n    print(predicted)\n    break","metadata":{"execution":{"iopub.status.busy":"2022-07-31T11:52:44.187671Z","iopub.execute_input":"2022-07-31T11:52:44.188322Z","iopub.status.idle":"2022-07-31T11:52:44.224276Z","shell.execute_reply.started":"2022-07-31T11:52:44.188285Z","shell.execute_reply":"2022-07-31T11:52:44.223245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['label'] = predicted.cpu().detach().numpy()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T11:52:44.857729Z","iopub.execute_input":"2022-07-31T11:52:44.858573Z","iopub.status.idle":"2022-07-31T11:52:44.864121Z","shell.execute_reply.started":"2022-07-31T11:52:44.858524Z","shell.execute_reply":"2022-07-31T11:52:44.863208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T11:52:45.433745Z","iopub.execute_input":"2022-07-31T11:52:45.434779Z","iopub.status.idle":"2022-07-31T11:52:45.441956Z","shell.execute_reply.started":"2022-07-31T11:52:45.434735Z","shell.execute_reply":"2022-07-31T11:52:45.440883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test","metadata":{"execution":{"iopub.status.busy":"2022-07-31T11:53:10.061735Z","iopub.execute_input":"2022-07-31T11:53:10.062407Z","iopub.status.idle":"2022-07-31T11:53:10.070405Z","shell.execute_reply.started":"2022-07-31T11:53:10.062368Z","shell.execute_reply":"2022-07-31T11:53:10.069437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}