{"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":"markdown","source":"**Note:**  Comment của em sẽ được viết bằng tiếng Việt\n---\nNotebook gốc mất đến 7 tiếng để chạy nên em xin điều chỉnh một vài số để notebook này chạy nhanh hơn","metadata":{}},{"cell_type":"code","source":"# xem trong thư mục có những file gì\n! ls -l /kaggle/input/herbarium-2022-fgvc9","metadata":{"execution":{"iopub.status.busy":"2022-08-11T06:24:21.781175Z","iopub.execute_input":"2022-08-11T06:24:21.781497Z","iopub.status.idle":"2022-08-11T06:24:22.78269Z","shell.execute_reply.started":"2022-08-11T06:24:21.781396Z","shell.execute_reply":"2022-08-11T06:24:22.781833Z"},"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\n# giống hàm set_theme(), đặt theme mặc định cho seaborn\nsn.set() \n\nPATH_DATASET = \"/kaggle/input/herbarium-2022-fgvc9\"\n\n# đọc training data và testing data từ file\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\n# in tên các cột (features) trong training data\npprint(train_data.keys())\n# số lượng data point trong testing data\npprint(len(test_data))\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T06:26:12.101253Z","iopub.execute_input":"2022-08-11T06:26:12.101738Z","iopub.status.idle":"2022-08-11T06:26:20.673208Z","shell.execute_reply.started":"2022-08-11T06:26:12.101699Z","shell.execute_reply":"2022-08-11T06:26:20.672469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Brief visualisations","metadata":{}},{"cell_type":"code","source":"# kiểm tra các subfeatures của annotations\ntrain_annotations = pd.DataFrame(train_data['annotations'])\ndisplay(train_annotations.head(3))\n\n# thống kê số lượng của mỗi genus, institution, category\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-08-11T06:26:24.161298Z","iopub.execute_input":"2022-08-11T06:26:24.161782Z","iopub.status.idle":"2022-08-11T06:26:27.152318Z","shell.execute_reply.started":"2022-08-11T06:26:24.161747Z","shell.execute_reply":"2022-08-11T06:26:27.151637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# chi tiết của mỗi category_id\ntrain_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-08-11T06:26:31.305511Z","iopub.execute_input":"2022-08-11T06:26:31.305778Z","iopub.status.idle":"2022-08-11T06:26:31.337384Z","shell.execute_reply.started":"2022-08-11T06:26:31.305748Z","shell.execute_reply":"2022-08-11T06:26:31.336322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# genus_id ứng với genus gì\ntrain_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-08-11T06:26:33.957919Z","iopub.execute_input":"2022-08-11T06:26:33.958644Z","iopub.status.idle":"2022-08-11T06:26:33.970311Z","shell.execute_reply.started":"2022-08-11T06:26:33.958608Z","shell.execute_reply":"2022-08-11T06:26:33.969529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# institution_id nào ứng với collection code nào\ntrain_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-08-11T06:26:36.793428Z","iopub.execute_input":"2022-08-11T06:26:36.793902Z","iopub.status.idle":"2022-08-11T06:26:36.803741Z","shell.execute_reply.started":"2022-08-11T06:26:36.793864Z","shell.execute_reply":"2022-08-11T06:26:36.802902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# image_id nào ứng với file ảnh nào\ntrain_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-08-11T06:26:38.733372Z","iopub.execute_input":"2022-08-11T06:26:38.733848Z","iopub.status.idle":"2022-08-11T06:26:39.492477Z","shell.execute_reply.started":"2022-08-11T06:26:38.733812Z","shell.execute_reply":"2022-08-11T06:26:39.49177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot khoảng cách Phylogenetic Distances giữa các genus cây (giống cây)\ntrain_distances = pd.DataFrame(train_data['distances'])\ndisplay(train_distances.head())\n\nfig = plt.figure(figsize=(18, 18))\n# reshape lại data theo genus_id để plot heatmap\nheat = train_distances.pivot(index=\"genus_id_y\", columns=\"genus_id_x\", values=\"distance\")\n# _ chỉ là throwaway variable\n_= sn.heatmap(heat, ax=fig.gca())","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2022-08-11T06:26:41.117494Z","iopub.execute_input":"2022-08-11T06:26:41.118311Z","iopub.status.idle":"2022-08-11T06:26:58.748918Z","shell.execute_reply.started":"2022-08-11T06:26:41.118257Z","shell.execute_reply":"2022-08-11T06:26:58.747907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Fused annotaions","metadata":{}},{"cell_type":"code","source":"# nối các subcategory của image_id category_id institution_id vào annotations (không nối genus_id vì đã có tên genus trong category_id)\ndf_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-08-11T06:27:17.852394Z","iopub.execute_input":"2022-08-11T06:27:17.853158Z","iopub.status.idle":"2022-08-11T06:27:18.753565Z","shell.execute_reply.started":"2022-08-11T06:27:17.853122Z","shell.execute_reply":"2022-08-11T06:27:18.752846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sample images ","metadata":{}},{"cell_type":"code","source":"# in một số hình để xem cách đọc file ảnh đã đúng chưa\n# 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    # path đến file ảnh\n    img_path = os.path.join(PATH_DATASET, \"train_images\", row[\"file_name\"])\n    # để ảnh vào từng subplot trong plot để hiện thị\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-08-11T07:52:08.369391Z","iopub.execute_input":"2022-08-11T07:52:08.36992Z","iopub.status.idle":"2022-08-11T07:52:11.256381Z","shell.execute_reply.started":"2022-08-11T07:52:08.369881Z","shell.execute_reply":"2022-08-11T07:52:11.254792Z"},"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\n# hàm để tính mean và std (standard deviation) của màu trong file hình (theo RGB)\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\n# glob để lấy các ảnh có tên có dạng PATH_DATASET/train_images/*/*/*.jpg\nimages = glob.glob(os.path.join(PATH_DATASET, \"train_images\", \"*\", \"*\", \"*.jpg\"))\n# images += glob.glob(os.path.join(PATH_DATASET, \"test_images\", \"*\", \"*.jpg\"))\n# sửa thành 300 hình vì tính toán sẽ khá lâu (ban đầu là 15000)\n# tqdm để hiện thanh progress meter như ở dưới\n# Parallel trong thư viện joblib để thực hiện công việc nhanh hơn bằng việc pipeline song song các task \n# (https://joblib.readthedocs.io/en/latest/generated/joblib.Parallel.html)\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(images[:300]))","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:46:54.207823Z","iopub.execute_input":"2022-08-11T07:46:54.2091Z","iopub.status.idle":"2022-08-11T07:47:11.582284Z","shell.execute_reply.started":"2022-08-11T07:46:54.209013Z","shell.execute_reply":"2022-08-11T07:47:11.581449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# một số thông tin thống kê về các mean và std của các hình (theo các màu RGB là 0 1 2)\n# c[0] là mean, c[1] là std\nimg_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\n# chọn mean của 2 bảng này làm img_color_mean và img_color_std để sử dụng ở dưới\n# ý nghĩa: img_color_mean[0] là mean của mean của màu 0 trong tất cả các hình\n#          img_color_std[0] là mean của std của màu 0 trong tất cả các hình\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-08-11T07:54:32.794752Z","iopub.execute_input":"2022-08-11T07:54:32.795048Z","iopub.status.idle":"2022-08-11T07:54:32.83554Z","shell.execute_reply.started":"2022-08-11T07:54:32.795014Z","shell.execute_reply":"2022-08-11T07:54:32.834791Z"},"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":"# tải EfficientDet để finetune model Flash bên dưới\n!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-08-11T08:02:52.276848Z","iopub.execute_input":"2022-08-11T08:02:52.277421Z","iopub.status.idle":"2022-08-11T08:03:54.872202Z","shell.execute_reply.started":"2022-08-11T08:02:52.277382Z","shell.execute_reply":"2022-08-11T08:03:54.871177Z"},"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-08-11T08:04:11.481855Z","iopub.execute_input":"2022-08-11T08:04:11.482658Z","iopub.status.idle":"2022-08-11T08:05:48.458574Z","shell.execute_reply.started":"2022-08-11T08:04:11.482618Z","shell.execute_reply":"2022-08-11T08:05:48.457506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch và các module cần thiết\nimport torch\n\nimport flash\nfrom flash.core.data.utils import download_data\nfrom flash.image import ImageClassificationData, ImageClassifier","metadata":{"execution":{"iopub.status.busy":"2022-08-11T08:07:48.987513Z","iopub.execute_input":"2022-08-11T08:07:48.98781Z","iopub.status.idle":"2022-08-11T08:08:00.876064Z","shell.execute_reply.started":"2022-08-11T08:07:48.987774Z","shell.execute_reply":"2022-08-11T08:08:00.875291Z"},"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# tác giả follow cách sử dụng thư viện Flash trong\n# https://lightning-flash.readthedocs.io/en/stable/reference/image_classification.html\n# class này chứa các cách transform ảnh để model có thể predict kể cả khi hình đã bị biến đổi \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            # sử dụng mean và std tìm được ở trên\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-08-11T08:11:27.406304Z","iopub.execute_input":"2022-08-11T08:11:27.406972Z","iopub.status.idle":"2022-08-11T08:11:27.417793Z","shell.execute_reply.started":"2022-08-11T08:11:27.406917Z","shell.execute_reply":"2022-08-11T08:11:27.416992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# datamodule này cũng follow hướng dẫn của flash trong \n# https://lightning-flash.readthedocs.io/en/latest/api/generated/flash.image.classification.data.ImageClassificationData.html\n# data module là data của các hình để train\ndatamodule = ImageClassificationData.from_data_frame(\n    input_field=\"file_name\",\n    target_fields=\"category_id\",\n    # for simplicity take just half of the data\n    # em lấy 1 phần 100 data vì một nửa chạy và train sẽ khá lâu\n    train_data_frame=df_train[:len(df_train) // 100],\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-08-11T08:26:19.886641Z","iopub.execute_input":"2022-08-11T08:26:19.887436Z","iopub.status.idle":"2022-08-11T08:26:23.980641Z","shell.execute_reply.started":"2022-08-11T08:26:19.887393Z","shell.execute_reply":"2022-08-11T08:26:23.979885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Build the task ⚙️","metadata":{}},{"cell_type":"code","source":"# tạo model với backbone là efficientnet_b0 để finetune ở dưới\nmodel = ImageClassifier(\n    backbone=\"efficientnet_b0\",\n    num_classes=datamodule.num_classes,\n    pretrained=True,\n    optimizer=\"AdamW\",\n    learning_rate=0.001,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T08:26:40.812771Z","iopub.execute_input":"2022-08-11T08:26:40.813071Z","iopub.status.idle":"2022-08-11T08:26:41.092258Z","shell.execute_reply.started":"2022-08-11T08:26:40.813036Z","shell.execute_reply":"2022-08-11T08:26:41.09149Z"},"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\n# lấy GPU để train model\nGPUS = int(torch.cuda.is_available())  # Set to 1 if GPU is enabled for notebook\n\n# swa = StochasticWeightAveraging(swa_epoch_start=0.6)\n# logger ghi epoch train ra log để quan sát\nlogger = CSVLogger(save_dir='logs/')\n\n# tạo trainer cho model\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,"execution":{"iopub.status.busy":"2022-08-11T08:26:43.195793Z","iopub.execute_input":"2022-08-11T08:26:43.19642Z","iopub.status.idle":"2022-08-11T08:26:43.204587Z","shell.execute_reply.started":"2022-08-11T08:26:43.196373Z","shell.execute_reply":"2022-08-11T08:26:43.203352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train và finetune model theo backbone\ntrainer.finetune(model, datamodule=datamodule, strategy=\"freeze\")\n\n# lưu model đã được train\ntrainer.save_checkpoint(\"image_classification_model.pt\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-11T08:26:46.087561Z","iopub.execute_input":"2022-08-11T08:26:46.088239Z","iopub.status.idle":"2022-08-11T08:34:37.82062Z","shell.execute_reply.started":"2022-08-11T08:26:46.0882Z","shell.execute_reply":"2022-08-11T08:34:37.819652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot accuracy và cross entropy loss để xem hiệu quả của model\n# bởi vì em giảm số data cho vào để thời gian chạy được nhanh nên graph khác với graph gốc của tác giả\n# tuy nhiên vẫn thấy được accuracy tăng và loss giảm theo số lần train\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 = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(15, 5)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T08:34:57.240876Z","iopub.execute_input":"2022-08-11T08:34:57.241425Z","iopub.status.idle":"2022-08-11T08:34:57.71769Z","shell.execute_reply.started":"2022-08-11T08:34:57.241384Z","shell.execute_reply":"2022-08-11T08:34:57.71684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference 🎉","metadata":{}},{"cell_type":"code","source":"# lấy test_image để thực hiện predict với model đã được train\ntest_images = pd.DataFrame(test_data).set_index(\"image_id\")\ndisplay(test_images.head())\nprint(f\"inference for {len(test_images)} images\")","metadata":{"execution":{"iopub.status.busy":"2022-08-11T08:37:21.525829Z","iopub.execute_input":"2022-08-11T08:37:21.526604Z","iopub.status.idle":"2022-08-11T08:37:21.741751Z","shell.execute_reply.started":"2022-08-11T08:37:21.526565Z","shell.execute_reply":"2022-08-11T08:37:21.740938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# thực hiện lại các bước như cho training data\ndatamodule = ImageClassificationData.from_data_frame(\n    input_field=\"file_name\",\n    # target_fields=\"category_id\",\n    #predict_data_frame=test_images,\n    # em comment dòng trên và lấy 2 dòng dưới của tác giả để predict data nhanh hơn\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":{"execution":{"iopub.status.busy":"2022-08-11T08:39:05.313435Z","iopub.execute_input":"2022-08-11T08:39:05.313968Z","iopub.status.idle":"2022-08-11T08:39:06.171442Z","shell.execute_reply.started":"2022-08-11T08:39:05.313914Z","shell.execute_reply":"2022-08-11T08:39:06.170694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# thực hiện predict và lưu vào predictions\npredictions = []\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":{"execution":{"iopub.status.busy":"2022-08-11T08:39:10.848939Z","iopub.execute_input":"2022-08-11T08:39:10.849667Z","iopub.status.idle":"2022-08-11T08:39:56.342814Z","shell.execute_reply.started":"2022-08-11T08:39:10.849631Z","shell.execute_reply":"2022-08-11T08:39:56.342022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# tạo dataframe cho đúng format với yêu cầu của cuộc thi và nộp file submission\nsubmission = 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":[]}]}