{"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":"## A look at the data","metadata":{}},{"cell_type":"markdown","source":"Let's start out by setting up our environment by importing the required modules and setting a random seed:","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport os\nimport pandas as pd\nfrom fastai.vision.all import *","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(999,reproducible=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_path = Path('../input/cassava-leaf-disease-classification')\nos.listdir(dataset_path)","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(dataset_path/'train.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['path'] = train_df['image_id'].map(lambda x:dataset_path/'train_images'/x)\ntrain_df = train_df.drop(columns=['image_id'])\ntrain_df = train_df.sample(frac=1).reset_index(drop=True) #shuffle dataframe\ntrain_df.head(10)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len_df = len(train_df)\nprint(f\"There are {len_df} images\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#delete part of label 3 data. Makes result worse!!! \n#0.3 frac gives about 76% accuracy, \n#0.7 frac gives about 82% accuracy\n\n#filtered_train_df = train_df[train_df['label'] != 3]\n#label3_df = train_df[train_df['label'] == 3]\n#label3_df = label3_df.sample(frac = 0.7)\n\n#result = pd.concat([filtered_train_df,label3_df])\n#train_df = result.sample(frac=1).reset_index(drop=True) #shuffle dataframe\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"All dataset >21,000 images\n\nThe distribution of the different classes:","metadata":{}},{"cell_type":"code","source":"train_df['label'].hist(figsize = (10, 5))\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Categories\n0. Cassava Bacterial Blight (CBB)\n1. Cassava Brown Streak Disease (CBSD)\n2. Cassava Green Mottle (CGM)\n3. Cassava Mosaic Disease (CMD)\n4. Healthy\n","metadata":{}},{"cell_type":"markdown","source":"## Data loading\n* The item transforms performs a large crop on each of the images.\n* The batch transforms performs random resized crop to 224 and also apply other standard augmentations (in `aug_tranforms`) at the batch level on the GPU.\n* The batch size is set to 256 here.\n","metadata":{}},{"cell_type":"code","source":"item_tfms = RandomResizedCrop(460, min_scale=0.75, ratio=(1.,1.))\nbatch_tfms = [*aug_transforms(size=224, max_warp=0), Normalize.from_stats(*imagenet_stats)]\nbs=256","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = ImageDataLoaders.from_df(train_df, #pass in train DataFrame\n                               valid_pct=0.2, #80-20 train-validation random split\n                               seed=999, #seed\n                               label_col=0, #label is in the first column of the DataFrame\n                               fn_col=1, #filename/path is in the second column of the DataFrame\n                               bs=bs, #pass in batch size\n                               item_tfms=item_tfms, #pass in item_tfms\n                               batch_tfms=batch_tfms) #pass in batch_tfms","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls.show_batch()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model training:","metadata":{}},{"cell_type":"code","source":"# Making pretrained weights work without needing to find the default filename\nif 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'","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn = cnn_learner(dls, \n                    resnet50, \n                    loss_func = LabelSmoothingCrossEntropy(), \n                    metrics = [accuracy], \n                    #cbs=MixUp()\n                   ).to_native_fp16()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.lr_find()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"`learn.fine_tune` trains frozen pretrained model for a single epoch (using one-cycle training), then train the whole pretrained model for several epochs using one-cycle training.","metadata":{}},{"cell_type":"code","source":"learn.fine_tune(5,base_lr=1e-2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.recorder.plot_loss()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Remove CBS to prevent mess up with learn.get_preds, predict, etc\n#learn.remove_cbs([MixUp])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Put the model back to fp32, and export the model ","metadata":{}},{"cell_type":"code","source":"learn = learn.to_native_fp32()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.export()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Checking the confusion matrix:**","metadata":{}},{"cell_type":"code","source":"interp = ClassificationInterpretation.from_learner(learn)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"interp.plot_confusion_matrix()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}