{"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":"# Explore 🔎 data...\n\nLets see what annotation and images we have :)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-02-15T09:52:20.579025Z","iopub.execute_input":"2022-02-15T09:52:20.579391Z","iopub.status.idle":"2022-02-15T09:52:20.583893Z","shell.execute_reply.started":"2022-02-15T09:52:20.579358Z","shell.execute_reply":"2022-02-15T09:52:20.583227Z"}}},{"cell_type":"code","source":"! ls -l /kaggle/input/herbarium-2022-fgvc9","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:20:16.81492Z","iopub.execute_input":"2022-02-25T20:20:16.815566Z","iopub.status.idle":"2022-02-25T20:20:17.500345Z","shell.execute_reply.started":"2022-02-25T20:20:16.815478Z","shell.execute_reply":"2022-02-25T20:20:17.499519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Loading the train and test meta","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nimport pandas as pd\nimport seaborn as sn\nimport matplotlib.pyplot as plt\nfrom pprint import pprint\n\nsn.set()\n\nPATH_DATASET = \"/kaggle/input/herbarium-2022-fgvc9\"\n\nwith open(os.path.join(PATH_DATASET, \"train_metadata.json\")) as fp:\n    train_data = json.load(fp)\n\nwith open(os.path.join(PATH_DATASET, \"test_metadata.json\")) as fp:\n    test_data = json.load(fp)\n\npprint(train_data.keys())\npprint(len(test_data))","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:20:17.50224Z","iopub.execute_input":"2022-02-25T20:20:17.502754Z","iopub.status.idle":"2022-02-25T20:20:31.854295Z","shell.execute_reply.started":"2022-02-25T20:20:17.502714Z","shell.execute_reply":"2022-02-25T20:20:31.853385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Brief visualisations","metadata":{}},{"cell_type":"code","source":"train_annotations = pd.DataFrame(train_data['annotations'])\ndisplay(train_annotations.head(3))\n\naxs = train_annotations[[\"genus_id\", \"institution_id\", \"category_id\"]].hist(bins=100, sharey=True, figsize=(8, 8), grid=True, layout=(3, 1))\n_= [ax.set_yscale('log') for ax in axs[0]]","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2022-02-25T20:20:31.855458Z","iopub.execute_input":"2022-02-25T20:20:31.855717Z","iopub.status.idle":"2022-02-25T20:20:35.321611Z","shell.execute_reply.started":"2022-02-25T20:20:31.855671Z","shell.execute_reply":"2022-02-25T20:20:35.320936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_categories = pd.DataFrame(train_data['categories']).set_index(\"category_id\")\ndisplay(train_categories.head())\n# (train_categories.index - train_categories.category_id).hist()","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-25T20:20:35.323616Z","iopub.execute_input":"2022-02-25T20:20:35.32484Z","iopub.status.idle":"2022-02-25T20:20:35.358377Z","shell.execute_reply.started":"2022-02-25T20:20:35.324798Z","shell.execute_reply":"2022-02-25T20:20:35.357566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_genera = pd.DataFrame(train_data['genera']).set_index(\"genus_id\")\ndisplay(train_genera.head())","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-25T20:20:35.362481Z","iopub.execute_input":"2022-02-25T20:20:35.362734Z","iopub.status.idle":"2022-02-25T20:20:35.376567Z","shell.execute_reply.started":"2022-02-25T20:20:35.362701Z","shell.execute_reply":"2022-02-25T20:20:35.375711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_institutions = pd.DataFrame(train_data['institutions']).set_index(\"institution_id\")\ndisplay(train_institutions.head())","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-25T20:20:35.377729Z","iopub.execute_input":"2022-02-25T20:20:35.378239Z","iopub.status.idle":"2022-02-25T20:20:35.393057Z","shell.execute_reply.started":"2022-02-25T20:20:35.378203Z","shell.execute_reply":"2022-02-25T20:20:35.392352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_images = pd.DataFrame(train_data['images']).set_index(\"image_id\")\ndisplay(train_images.head())","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-25T20:20:35.394168Z","iopub.execute_input":"2022-02-25T20:20:35.394474Z","iopub.status.idle":"2022-02-25T20:20:36.251701Z","shell.execute_reply.started":"2022-02-25T20:20:35.394438Z","shell.execute_reply":"2022-02-25T20:20:36.249736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_distances = pd.DataFrame(train_data['distances'])\ndisplay(train_distances.head())\n\nfig = plt.figure(figsize=(18, 18))\nheat = train_distances.pivot(index=\"genus_id_y\", columns=\"genus_id_x\", values=\"distance\")\n_= sn.heatmap(heat, ax=fig.gca())","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2022-02-25T20:20:36.254421Z","iopub.execute_input":"2022-02-25T20:20:36.25462Z","iopub.status.idle":"2022-02-25T20:20:54.230373Z","shell.execute_reply.started":"2022-02-25T20:20:36.254595Z","shell.execute_reply":"2022-02-25T20:20:54.229583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Fused annotaions","metadata":{}},{"cell_type":"code","source":"df_train = pd.merge(train_annotations, train_images, how=\"left\", right_index=True, left_on=\"image_id\")\ndf_train = pd.merge(df_train, train_categories, how=\"left\", right_index=True, left_on=\"category_id\")\ndf_train = pd.merge(df_train, train_institutions, how=\"left\", right_index=True, left_on=\"institution_id\")\n# df_train = pd.merge(df_train, train_genera, how=\"left\", right_index=True, left_on=\"genus_id\")\n\ndisplay(df_train.head())\nprint(f\"training images: {len(df_train)}\")","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:20:54.232943Z","iopub.execute_input":"2022-02-25T20:20:54.233825Z","iopub.status.idle":"2022-02-25T20:20:55.138578Z","shell.execute_reply.started":"2022-02-25T20:20:54.233785Z","shell.execute_reply":"2022-02-25T20:20:55.137707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sample images ","metadata":{}},{"cell_type":"code","source":"# shuffle\ndf_train.sample(frac=1)\n\nfig, axarr = plt.subplots(nrows=2, ncols=5, figsize=(12, 6))\nfor i, (_, row) in enumerate(df_train[:10].iterrows()):\n    img_path = os.path.join(PATH_DATASET, \"train_images\", row[\"file_name\"])\n    img = plt.imread(img_path)\n    axarr[i // 5, i % 5].imshow(img)\n#     print(row)\nfig.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:20:55.140215Z","iopub.execute_input":"2022-02-25T20:20:55.140493Z","iopub.status.idle":"2022-02-25T20:20:58.567031Z","shell.execute_reply.started":"2022-02-25T20:20:55.140455Z","shell.execute_reply":"2022-02-25T20:20:58.565623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport numpy as np\nfrom tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\ndef _color_means(img_path):\n    img = plt.imread(img_path)\n    means = {i: np.mean(img[..., i]) / 255.0 for i in range(3)}\n    std = {i: np.std(img[..., i]) / 255.0 for i in range(3)}\n    return means, std\n\nimages = glob.glob(os.path.join(PATH_DATASET, \"train_images\", \"*\", \"*\", \"*.jpg\"))\n# images += glob.glob(os.path.join(PATH_DATASET, \"test_images\", \"*\", \"*.jpg\"))\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(images[:15000]))","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:20:58.568313Z","iopub.execute_input":"2022-02-25T20:20:58.569039Z","iopub.status.idle":"2022-02-25T20:25:57.986177Z","shell.execute_reply.started":"2022-02-25T20:20:58.569001Z","shell.execute_reply":"2022-02-25T20:25:57.985386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_color_mean = pd.DataFrame([c[0] for c in clr_mean_std]).describe()\ndisplay(img_color_mean)\nimg_color_std = pd.DataFrame([c[1] for c in clr_mean_std]).describe()\ndisplay(img_color_std)\n\nimg_color_mean = list(img_color_mean.T[\"mean\"])\nimg_color_std = list(img_color_std.T[\"mean\"])\nprint(img_color_mean, img_color_std)","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:25:57.987833Z","iopub.execute_input":"2022-02-25T20:25:57.988099Z","iopub.status.idle":"2022-02-25T20:25:58.046079Z","shell.execute_reply.started":"2022-02-25T20:25:57.988065Z","shell.execute_reply":"2022-02-25T20:25:58.04538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training with Lightning⚡Flash\n\n**Follow the example:** https://lightning-flash.readthedocs.io/en/stable/reference/image_classification.html\n\n\n**Later you would need to adjust the image size to used model:**\n\n| **Base model** | resolution |\n|----------------|------------|\n| EfficientNetB0 | 224        |\n| EfficientNetB1 | 240        |\n| EfficientNetB2 | 260        |\n| EfficientNetB3 | 300        |\n| EfficientNetB4 | 380        |\n| EfficientNetB5 | 456        |\n| EfficientNetB6 | 528        |\n| EfficientNetB7 | 600        |","metadata":{}},{"cell_type":"code","source":"!pip install -q effdet \"icevision[all]\" 'lightning-flash[image]'\n# !pip install -q \"pytorch-lightning==1.4.*\"\n!pip uninstall -y wandb","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-25T20:25:58.047471Z","iopub.execute_input":"2022-02-25T20:25:58.047741Z","iopub.status.idle":"2022-02-25T20:27:18.184854Z","shell.execute_reply.started":"2022-02-25T20:25:58.047704Z","shell.execute_reply":"2022-02-25T20:27:18.183852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip download -q effdet \"icevision[all]\" 'lightning-flash[image]' --dest frozen_packages --prefer-binary\n!rm frozen_packages/torch-*\n!ls -l frozen_packages","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-25T20:27:18.190261Z","iopub.execute_input":"2022-02-25T20:27:18.192451Z","iopub.status.idle":"2022-02-25T20:29:14.304456Z","shell.execute_reply.started":"2022-02-25T20:27:18.192406Z","shell.execute_reply":"2022-02-25T20:29:14.303594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nimport flash\nfrom flash.core.data.utils import download_data\nfrom flash.image import ImageClassificationData, ImageClassifier","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:29:14.306477Z","iopub.execute_input":"2022-02-25T20:29:14.306807Z","iopub.status.idle":"2022-02-25T20:29:26.505673Z","shell.execute_reply.started":"2022-02-25T20:29:14.306763Z","shell.execute_reply":"2022-02-25T20:29:26.504844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Create the DataModule 🗄️","metadata":{}},{"cell_type":"code","source":"from dataclasses import dataclass\nfrom torchvision import transforms as T\nfrom typing import Tuple, Callable\nfrom flash.core.data.io.input_transform import InputTransform\n\n@dataclass\nclass ImageClassificationInputTransform(InputTransform):\n\n    image_size: Tuple[int, int] = (224, 224)\n\n    def input_per_sample_transform(self):\n        return T.Compose([\n            T.ToTensor(),\n            T.Resize(self.image_size),\n            # T.Normalize([0.778, 0.756, 0.709], [0.246, 0.250, 0.253]),\n            T.Normalize(img_color_mean, img_color_std),\n        ])\n\n    def train_input_per_sample_transform(self):\n        return T.Compose([\n            T.ToTensor(),\n            T.Resize(self.image_size),\n            # T.Normalize([0.778, 0.756, 0.709], [0.246, 0.250, 0.253]),\n            T.Normalize(img_color_mean, img_color_std),\n            T.RandomHorizontalFlip(),\n            T.RandomAffine(degrees=10, scale=(0.9, 1.1), translate=(0.1, 0.1)),\n            # T.ColorJitter(),\n            # T.RandomAutocontrast(),\n            # T.RandomPerspective(distortion_scale=0.1),\n        ])\n\n    def target_per_sample_transform(self) -> Callable:\n        return torch.as_tensor","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:29:26.507066Z","iopub.execute_input":"2022-02-25T20:29:26.507355Z","iopub.status.idle":"2022-02-25T20:29:26.519025Z","shell.execute_reply.started":"2022-02-25T20:29:26.507319Z","shell.execute_reply":"2022-02-25T20:29:26.517617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"datamodule = ImageClassificationData.from_data_frame(\n    input_field=\"file_name\",\n    target_fields=\"category_id\",\n    # for simplicity take just half of the data\n    train_data_frame=df_train[:len(df_train) // 2],\n    train_images_root=os.path.join(PATH_DATASET, \"train_images\"),\n    train_transform=ImageClassificationInputTransform,\n    batch_size=128,\n    transform_kwargs={\"image_size\": (224, 224)},\n    num_workers=3,\n)","metadata":{"execution":{"iopub.status.busy":"2022-02-25T20:29:26.521114Z","iopub.execute_input":"2022-02-25T20:29:26.521406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Build the task ⚙️","metadata":{}},{"cell_type":"code","source":"model = ImageClassifier(\n    backbone=\"efficientnet_b0\",\n    num_classes=datamodule.num_classes,\n    pretrained=True,\n    optimizer=\"AdamW\",\n    learning_rate=0.001,\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Finetune the model 🛠️","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.loggers import CSVLogger\n# from pytorch_lightning.callbacks import StochasticWeightAveraging\n\n# Trainer Args\nGPUS = int(torch.cuda.is_available())  # Set to 1 if GPU is enabled for notebook\n\n# swa = StochasticWeightAveraging(swa_epoch_start=0.6)\nlogger = CSVLogger(save_dir='logs/')\n\ntrainer = flash.Trainer(\n    max_epochs=3,\n    # gradient_clip_val=0.01,\n    gpus=GPUS,\n    precision=16 if GPUS else 32,\n    logger=logger,\n    accumulate_grad_batches=32,\n)","metadata":{"_kg_hide-output":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.finetune(model, datamodule=datamodule, strategy=\"freeze\")\n\ntrainer.save_checkpoint(\"image_classification_model.pt\")","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics = 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 = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(15, 5)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference 🎉","metadata":{}},{"cell_type":"code","source":"test_images = pd.DataFrame(test_data).set_index(\"image_id\")\ndisplay(test_images.head())\nprint(f\"inference for {len(test_images)} images\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"datamodule = ImageClassificationData.from_data_frame(\n    input_field=\"file_name\",\n    # target_fields=\"category_id\",\n    predict_data_frame=test_images,\n    # for simplicity take just fraction of the data\n    # predict_data_frame=test_images[:len(test_images) // 100],\n    predict_images_root=os.path.join(PATH_DATASET, \"test_images\"),\n    batch_size=16,\n    transform_kwargs={\"image_size\": (224, 224)},\n    num_workers=2,\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nfor lbs in trainer.predict(model, datamodule=datamodule, output=\"labels\"):\n    # lbs = [torch.argmax(p[\"preds\"].float()).item() for p in preds]\n    predictions += lbs","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame({\"id\": test_images.index, \"Predicted\": predictions}).set_index(\"id\")\nsubmission.to_csv(\"submission.csv\")\n\n! head submission.csv","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}