{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":33679,"databundleVersionId":3212216,"sourceType":"competition"},{"sourceId":3627198,"sourceType":"datasetVersion","datasetId":1942232},{"sourceId":88750059,"sourceType":"kernelVersion"}],"dockerImageVersionId":30163,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! ls -l /kaggle/input/","metadata":{"execution":{"iopub.status.busy":"2025-03-02T18:21:44.671302Z","iopub.execute_input":"2025-03-02T18:21:44.671626Z","iopub.status.idle":"2025-03-02T18:21:45.744836Z","shell.execute_reply.started":"2025-03-02T18:21:44.671532Z","shell.execute_reply":"2025-03-02T18:21:45.743985Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Browse test images ","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nPATH_DATASET = \"/kaggle/input/herbarium-2022-fgvc9\"\n\nwith open(os.path.join(PATH_DATASET, \"test_metadata.json\")) as fp:\n    test_data = json.load(fp)\n\nprint(len(test_data))\ndf_test = pd.DataFrame(test_data).set_index(\"image_id\")\ndisplay(df_test.head())","metadata":{"execution":{"iopub.status.busy":"2025-03-02T18:21:45.747336Z","iopub.execute_input":"2025-03-02T18:21:45.748037Z","iopub.status.idle":"2025-03-02T18:21:46.436065Z","shell.execute_reply.started":"2025-03-02T18:21:45.747982Z","shell.execute_reply":"2025-03-02T18:21:46.435331Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axarr = plt.subplots(nrows=2, ncols=5, figsize=(12, 6))\nfor i, (_, row) in enumerate(df_test[:10].iterrows()):\n    img_path = os.path.join(PATH_DATASET, \"test_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":"2025-03-02T18:21:46.437053Z","iopub.execute_input":"2025-03-02T18:21:46.437256Z","iopub.status.idle":"2025-03-02T18:21:48.638371Z","shell.execute_reply.started":"2025-03-02T18:21:46.437231Z","shell.execute_reply":"2025-03-02T18:21:48.637668Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference with Lightning⚡Flash\n","metadata":{}},{"cell_type":"code","source":"!pip install -q 'lightning-flash[image]' \"torchmetrics==0.7.*\" --find-links /kaggle/input/herbarium-eda-baseline-flash-efficientnet/frozen_packages/ --no-index\n!pip install -q timm -U --find-links /kaggle/input/herbarium-submissions/packages/ --no-index\n!pip uninstall -y wandb","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-03-02T18:21:48.639386Z","iopub.execute_input":"2025-03-02T18:21:48.639591Z","iopub.status.idle":"2025-03-02T18:22:20.508389Z","shell.execute_reply.started":"2025-03-02T18:21:48.639565Z","shell.execute_reply":"2025-03-02T18:22:20.507661Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport flash\nfrom flash.image import ImageClassificationData, ImageClassifier","metadata":{"execution":{"iopub.status.busy":"2025-03-02T18:22:20.509752Z","iopub.execute_input":"2025-03-02T18:22:20.510013Z","iopub.status.idle":"2025-03-02T18:22:32.961742Z","shell.execute_reply.started":"2025-03-02T18:22:20.509981Z","shell.execute_reply":"2025-03-02T18:22:32.961018Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1. Load the task ⚙️","metadata":{}},{"cell_type":"code","source":"model = ImageClassifier.load_from_checkpoint(\n#     \"/kaggle/input/herbarium-eda-baseline-flash-efficientnet/image_classification_model.pt\"\n    \"/kaggle/input/herbarium-submissions/herbarium-classif-2nwcf7mv_convnext_base_384_in22ft1k-384px.pt\"\n)","metadata":{"execution":{"iopub.status.busy":"2025-03-02T18:22:32.9629Z","iopub.execute_input":"2025-03-02T18:22:32.963169Z","iopub.status.idle":"2025-03-02T18:23:16.48791Z","shell.execute_reply.started":"2025-03-02T18:22:32.963135Z","shell.execute_reply":"2025-03-02T18:23:16.487197Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Trainer Args\nGPUS = int(torch.cuda.is_available())  # Set to 1 if GPU is enabled for notebook\ntrainer = flash.Trainer(gpus=GPUS)","metadata":{"_kg_hide-output":false,"execution":{"iopub.status.busy":"2025-03-02T18:23:16.4903Z","iopub.execute_input":"2025-03-02T18:23:16.490842Z","iopub.status.idle":"2025-03-02T18:23:16.545042Z","shell.execute_reply.started":"2025-03-02T18:23:16.490807Z","shell.execute_reply":"2025-03-02T18:23:16.54431Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2. Run predictions 🎉","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    color_mean = (0.778, 0.756, 0.709)\n    color_std = (0.246, 0.250, 0.253)\n\n    def input_per_sample_transform(self):\n        return T.Compose([\n            T.ToTensor(),\n            T.Resize(self.image_size),\n            T.Normalize(self.color_mean, self.color_std),\n        ])\n\n    def target_per_sample_transform(self) -> Callable:\n        return torch.as_tensor","metadata":{"execution":{"iopub.status.busy":"2025-03-02T18:23:16.546247Z","iopub.execute_input":"2025-03-02T18:23:16.546494Z","iopub.status.idle":"2025-03-02T18:23:16.557292Z","shell.execute_reply.started":"2025-03-02T18:23:16.54646Z","shell.execute_reply":"2025-03-02T18:23:16.55671Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"datamodule = ImageClassificationData.from_data_frame(\n    input_field=\"file_name\",\n    predict_data_frame=df_test,\n    # for simplicity take just fraction of the data\n    # predict_data_frame=df_test[:len(df_test) // 1000],\n    predict_images_root=os.path.join(PATH_DATASET, \"test_images\"),\n    predict_transform=ImageClassificationInputTransform,\n    batch_size=3,\n    transform_kwargs={\"image_size\": (384, 384)},\n    num_workers=3,\n)","metadata":{"execution":{"iopub.status.busy":"2025-03-02T18:23:16.558261Z","iopub.execute_input":"2025-03-02T18:23:16.558503Z","iopub.status.idle":"2025-03-02T18:30:09.509328Z","shell.execute_reply.started":"2025-03-02T18:23:16.55847Z","shell.execute_reply":"2025-03-02T18:30:09.50818Z"},"trusted":true},"outputs":[],"execution_count":null},{"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.DataFrame({\"id\": df_test.index, \"Predicted\": predictions}).set_index(\"id\")\nsubmission.to_csv(\"submission.csv\")\n\n! head submission.csv","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}