{"cells":[{"metadata":{},"cell_type":"markdown","source":"### Imports"},{"metadata":{"trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nimport glob\n%matplotlib inline\nfrom fastai.vision.all import *\nfrom fastai import *\n\nset_seed(42) #Set random seed to a constant so tests are reproducible\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Set paths"},{"metadata":{"trusted":true},"cell_type":"code","source":"path = Path('../input/cassava-leaf-disease-classification')\npath.ls()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Read training data csv"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df = pd.read_csv(path/'train.csv')\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Add 'train_images/' to image_id column to easily access directory"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df['image_id'] = 'train_images/' + train_df['image_id']\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Add qualitative labels"},{"metadata":{"trusted":true},"cell_type":"code","source":"import json\nwith open(path/'label_num_to_disease_map.json') as json_file:\n    data = json.load(json_file)\n    print(data)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"labels_dic = {0: 'Bacterial Blight',\n1: 'Brown Streak Disease',\n2: 'Green Mottle',\n3: 'Mosaic Disease',\n4: 'Healthy'\n}\ntrain_df['qual_label'] = train_df['label'].map(labels_dic)\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Datablock"},{"metadata":{},"cell_type":"markdown","source":"Functions to obtain x and y - image paths and labels."},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_x(row):\n    return path/row['image_id']\n\ndef get_y(row):\n    return row['label']","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Create data block with validation set of 20%, transforming each item to 448x448px and then randomly cropping batches to 224x224px. Other data augmentation also applied to batches, which should improve the accuracy."},{"metadata":{"trusted":true},"cell_type":"code","source":"CassavaBlock = DataBlock(\n    blocks = (ImageBlock, CategoryBlock), \n    splitter = RandomSplitter(valid_pct=0.2, seed=42),\n    get_x = get_x,\n    get_y = get_y,\n    item_tfms = Resize(448),\n    batch_tfms = [RandomResizedCropGPU(224), *aug_transforms(), Normalize.from_stats(*imagenet_stats)] #Data augmentation\n)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Data loaders. Show 4 images."},{"metadata":{"trusted":true},"cell_type":"code","source":"dls = CassavaBlock.dataloaders(train_df, batch_size=64)\ndls.valid.show_batch(max_n=4, nrows=1)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Training"},{"metadata":{},"cell_type":"markdown","source":"The compeition does not allow internet access. Normally FastAI can obtain weights for ResNet-50 from the Internet, but now it must be done offline. Obtain weights from Kaggle dataset then copy the file over to the directory at which FastAI will expect it."},{"metadata":{"trusted":true},"cell_type":"code","source":"if not os.path.exists('/root/.cache/torch/hub/checkpoints/'):\n        os.makedirs('/root/.cache/torch/hub/checkpoints/')\n!cp '../input/resnet50/resnet50.pth' '/root/.cache/torch/hub/checkpoints/resnet50-19c8e357.pth'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Create the model and fine tune to our data, with 7 epochs."},{"metadata":{"trusted":true},"cell_type":"code","source":"learn = cnn_learner(dls, resnet50, metrics=accuracy, loss_func = LabelSmoothingCrossEntropy(), opt_func = ranger)\nlearn.fine_tune(10)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Plot confusion matrix."},{"metadata":{"trusted":true},"cell_type":"code","source":"interp = ClassificationInterpretation.from_learner(learn)\ninterp.plot_confusion_matrix()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Show worst 5 images in terms of loss - i.e. the images which the model is not predicting well on."},{"metadata":{"trusted":true},"cell_type":"code","source":"interp.plot_top_losses(5, nrows=5)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"From the above, the predictor is confusing Mosaic Disease with other diseases, namely Brown Steak Disease and Green Mottle. An expert in Cassava plants, and plants in general, would be able to provide insight into this - further reading required."},{"metadata":{},"cell_type":"markdown","source":"### Predictions"},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df = pd.read_csv(path/'sample_submission.csv') #Read csv\nsample_df_copy = sample_df.copy() #Make copy so that when uploading original, image ids are unchanged.\nsample_df_copy.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df_copy['image_id'] = 'test_images/' + sample_df_copy['image_id'] #Add path to image ids","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_dl = dls.test_dl(sample_df_copy) #Data loader\ntest_dl.show_batch()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"preds = learn.tta(dl=test_dl, n=8, beta=0) #Predictions for each class (probability)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df['label'] = np.argmax(preds[0], axis=1) #Add prediction to original dataframe - maximum probability","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df.to_csv('submission.csv', index=False) #Dataframe to csv","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}