{"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":"# Plant Pathology 2021 - FGVC8","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y torchtext\n!pip install -q --upgrade torch torchvision\n!pip install -q \"lightning-flash[image]\" \"torchmetrics<0.8\"\n!pip install -q -U timm segmentation-models-pytorch\n\n! pip list | grep torch\n! pip list | grep lightning\n! nvidia-smi -L","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-04-20T14:11:32.950467Z","iopub.execute_input":"2022-04-20T14:11:32.95077Z","iopub.status.idle":"2022-04-20T14:14:14.8737Z","shell.execute_reply.started":"2022-04-20T14:11:32.950698Z","shell.execute_reply":"2022-04-20T14:14:14.872875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data exploration\n\nChecking what data do we have available and what is the labels distribution...","metadata":{}},{"cell_type":"code","source":"%matplotlib inline\n\nimport os\nimport json\nimport pandas as pd\nfrom pprint import pprint\n\nbase_path = '/kaggle/input/plant-pathology-2021-fgvc8-960px'\npath_csv = os.path.join(base_path, 'train.csv')\ntrain_data = pd.read_csv(path_csv)\ndisplay(train_data.head())","metadata":{"execution":{"iopub.status.busy":"2022-04-20T14:14:14.8773Z","iopub.execute_input":"2022-04-20T14:14:14.877581Z","iopub.status.idle":"2022-04-20T14:14:14.929516Z","shell.execute_reply.started":"2022-04-20T14:14:14.877553Z","shell.execute_reply":"2022-04-20T14:14:14.9284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can see that each image can have multiple labels so lets check what is the mos common label count...\n\n*The target classes, a space delimited list of all diseases found in the image.\nUnhealthy leaves with too many diseases to classify visually will have the complex class, and may also have a subset of the diseases identified.*","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\ntrain_data['nb_classes'] = [len(lbs.split(\" \")) for lbs in train_data['labels']]\nlb_hist = dict(zip(range(10), np.bincount(train_data['nb_classes'])))\npprint(lb_hist)","metadata":{"execution":{"iopub.status.busy":"2022-04-20T14:14:14.930878Z","iopub.execute_input":"2022-04-20T14:14:14.931253Z","iopub.status.idle":"2022-04-20T14:14:14.951837Z","shell.execute_reply.started":"2022-04-20T14:14:14.931219Z","shell.execute_reply":"2022-04-20T14:14:14.950909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Browse the label distribution, enrolling all labels in the dataset, so in case an image has two labels both are used in this stat...","metadata":{}},{"cell_type":"code","source":"import itertools\nimport seaborn as sns\n\nlabels_all = list(itertools.chain(*[lbs.split(\" \") for lbs in train_data['labels']]))\ntrain_data['labels_sorted'] = [\" \".join(sorted(lbs.split(\" \"))) for lbs in train_data['labels']]\n\nsns.set()\nax = sns.countplot(y=labels_all, orient='v')\nax.grid()","metadata":{"execution":{"iopub.status.busy":"2022-04-20T14:14:14.95355Z","iopub.execute_input":"2022-04-20T14:14:14.954141Z","iopub.status.idle":"2022-04-20T14:14:16.114944Z","shell.execute_reply.started":"2022-04-20T14:14:14.954103Z","shell.execute_reply":"2022-04-20T14:14:16.114088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Flash finetuning","metadata":{}},{"cell_type":"code","source":"import flash\nimport torch\nimport pytorch_lightning as pl\nfrom flash.image import ImageClassificationData, ImageClassifier","metadata":{"execution":{"iopub.status.busy":"2022-04-20T14:14:16.117813Z","iopub.execute_input":"2022-04-20T14:14:16.118219Z","iopub.status.idle":"2022-04-20T14:14:24.32085Z","shell.execute_reply.started":"2022-04-20T14:14:16.118178Z","shell.execute_reply":"2022-04-20T14:14:24.320027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Load the data","metadata":{}},{"cell_type":"code","source":"datamodule = ImageClassificationData.from_data_frame(\n    \"image\",\n    # list(labels_uq),\n    \"labels\",\n    train_data_frame=train_data,\n    train_images_root=os.path.join(base_path, \"train_images\"),\n    transform_kwargs={\"image_size\": (384, 384)},\n    batch_size=24,\n    num_workers=2,\n    val_split=0.2,\n)\nprint(datamodule.multi_label)","metadata":{"execution":{"iopub.status.busy":"2022-04-20T14:14:24.322455Z","iopub.execute_input":"2022-04-20T14:14:24.322825Z","iopub.status.idle":"2022-04-20T14:14:39.398475Z","shell.execute_reply.started":"2022-04-20T14:14:24.322786Z","shell.execute_reply":"2022-04-20T14:14:39.397539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Build the model","metadata":{}},{"cell_type":"code","source":"model = ImageClassifier(\n    backbone=\"tf_efficientnet_b4_ns\",\n    optimizer=torch.optim.AdamW,\n    learning_rate=0.005,\n    labels=datamodule.labels,\n    multi_label=datamodule.multi_label,\n)","metadata":{"execution":{"iopub.status.busy":"2022-04-20T14:14:39.400086Z","iopub.execute_input":"2022-04-20T14:14:39.400717Z","iopub.status.idle":"2022-04-20T14:15:05.74973Z","shell.execute_reply.started":"2022-04-20T14:14:39.400677Z","shell.execute_reply":"2022-04-20T14:15:05.748961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Create the trainer","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\n\nlogger = pl.loggers.CSVLogger(save_dir='logs/')\n\ntrainer = flash.Trainer(\n    gpus=1,\n    logger=logger,\n    max_epochs=5,\n    precision=16,\n    val_check_interval=0.5,\n    # limit_train_batches=0.1,\n    # limit_val_batches=0.1,\n)","metadata":{"execution":{"iopub.status.busy":"2022-04-20T14:15:05.750973Z","iopub.execute_input":"2022-04-20T14:15:05.751337Z","iopub.status.idle":"2022-04-20T14:15:05.807413Z","shell.execute_reply.started":"2022-04-20T14:15:05.751293Z","shell.execute_reply":"2022-04-20T14:15:05.806477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. Train the model","metadata":{}},{"cell_type":"code","source":"# Train the model\ntrainer.finetune(model, datamodule=datamodule, strategy=('freeze_unfreeze', 1))\n\n# Save it!\ntrainer.save_checkpoint(\"image_classification_model.pt\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-04-20T14:15:05.808921Z","iopub.execute_input":"2022-04-20T14:15:05.809571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nmetrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\ndel metrics[\"step\"]\nmetrics.set_index(\"epoch\", inplace=True)\ndisplay(metrics.dropna(axis=1, how=\"all\").head())\ng = sns.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(12, 4)\nplt.grid()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}