{"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":"## Install Detectron2\n","metadata":{}},{"cell_type":"code","source":"%%capture\nimport sys, os, distutils.core\n# Note: This is a faster way to install detectron2 in Colab, but it does not include all functionalities (e.g. compiled operators).\n# See https://detectron2.readthedocs.io/tutorials/install.html for full installation instructions\n!git clone 'https://github.com/facebookresearch/detectron2'\ndist = distutils.core.run_setup(\"./detectron2/setup.py\")\n!python -m pip install {' '.join([f\"'{x}'\" for x in dist.install_requires])}\nsys.path.insert(0, os.path.abspath('./detectron2'))","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:14:45.106104Z","iopub.execute_input":"2023-08-06T15:14:45.106458Z","iopub.status.idle":"2023-08-06T15:15:36.064263Z","shell.execute_reply.started":"2023-08-06T15:14:45.106426Z","shell.execute_reply":"2023-08-06T15:15:36.062716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Import Libraries\n","metadata":{}},{"cell_type":"code","source":"from detectron2.utils.memory import retry_if_cuda_oom\nfrom detectron2.utils.logger import setup_logger\nfrom detectron2.checkpoint import DetectionCheckpointer\nfrom detectron2.modeling import build_model\nfrom detectron2.evaluation import COCOEvaluator, inference_on_dataset\nimport detectron2.data.transforms as T\nfrom detectron2.data import detection_utils as utils\nfrom detectron2.data import DatasetCatalog, MetadataCatalog, build_detection_test_loader, build_detection_train_loader, DatasetMapper\nfrom detectron2.utils.visualizer import Visualizer\nfrom detectron2.structures import BoxMode\nfrom detectron2.engine import DefaultPredictor, DefaultTrainer\nfrom detectron2.config import get_cfg\nfrom detectron2 import model_zoo\n\nimport pandas as pd\nimport numpy as np\nfrom tqdm.notebook import tqdm  # progress bar\nimport matplotlib.pyplot as plt\nimport json\nimport cv2\nimport copy\nfrom typing import Optional\n\nfrom IPython.display import FileLink\nimport sys\n# torch\nimport torch\n\nimport gc\n\nimport warnings\n# Ignore \"future\" warnings and Data-Frame-Slicing warnings.\nwarnings.filterwarnings('ignore')\n\nsetup_logger()","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:15:36.067456Z","iopub.execute_input":"2023-08-06T15:15:36.069654Z","iopub.status.idle":"2023-08-06T15:15:37.025495Z","shell.execute_reply.started":"2023-08-06T15:15:36.069597Z","shell.execute_reply":"2023-08-06T15:15:37.024513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Downloading unilm\n","metadata":{}},{"cell_type":"code","source":"#better to use gdown \n!pip install gdown\n!gdown 1KQTZ6mXstpckzAix3k3XPtY-iEdqyeKD","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:17:56.009068Z","iopub.execute_input":"2023-08-06T15:17:56.009484Z","iopub.status.idle":"2023-08-06T15:18:10.859625Z","shell.execute_reply.started":"2023-08-06T15:17:56.009449Z","shell.execute_reply":"2023-08-06T15:18:10.858352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Unzipping unilm\n","metadata":{}},{"cell_type":"code","source":"# Replace '/kaggle/working/unilm.zip' with the actual path to your 'unilm.zip' file\nzip_file_path = '/kaggle/working/unilm.zip'\n\n# Replace 'unilm' with the name of the folder where you want to unzip the contents\noutput_folder = 'unilm'\n\n\n# Unzip the file\n!unzip $zip_file_path -d $output_folder","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:18:31.597910Z","iopub.execute_input":"2023-08-06T15:18:31.598359Z","iopub.status.idle":"2023-08-06T15:18:31.604396Z","shell.execute_reply.started":"2023-08-06T15:18:31.598319Z","shell.execute_reply":"2023-08-06T15:18:31.603430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setting Path\n","metadata":{}},{"cell_type":"code","source":"sys.path.insert(1, \"/kaggle/working/unilm/layoutlmv3\")\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:18:53.882884Z","iopub.execute_input":"2023-08-06T15:18:53.883255Z","iopub.status.idle":"2023-08-06T15:18:53.887894Z","shell.execute_reply.started":"2023-08-06T15:18:53.883225Z","shell.execute_reply":"2023-08-06T15:18:53.886961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! sed -i 's/from collections import Iterable/from collections.abc import Iterable/' /kaggle/working/unilm/layoutlmv3/examples/object_detection/ditod/table_evaluation/data_structure.py\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:04.927253Z","iopub.execute_input":"2023-08-06T15:19:04.927644Z","iopub.status.idle":"2023-08-06T15:19:05.900332Z","shell.execute_reply.started":"2023-08-06T15:19:04.927594Z","shell.execute_reply":"2023-08-06T15:19:05.898922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Importing vit config\n","metadata":{}},{"cell_type":"code","source":"from examples.object_detection.ditod import add_vit_config\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:11.473510Z","iopub.execute_input":"2023-08-06T15:19:11.474471Z","iopub.status.idle":"2023-08-06T15:19:20.442243Z","shell.execute_reply.started":"2023-08-06T15:19:11.474432Z","shell.execute_reply":"2023-08-06T15:19:20.441294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cfg = get_cfg()\n# # Add PointRend-specific config\n# add_vit_config(cfg)\n# # Load a config from file\n# cfg.merge_from_file(\"unilm/layoutlmv3/examples/object_detection/cascade_layoutlmv3.yaml\")\n# print(cfg)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:29.596857Z","iopub.execute_input":"2023-08-06T15:19:29.597220Z","iopub.status.idle":"2023-08-06T15:19:29.601915Z","shell.execute_reply.started":"2023-08-06T15:19:29.597191Z","shell.execute_reply":"2023-08-06T15:19:29.600898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setting Condition\n","metadata":{}},{"cell_type":"code","source":"from datetime import datetime\n\n# if False, model is set to `PRETRAINED_PATH` model\nis_train = True\n\n# if True, evaluate on validation dataset\nis_evaluate = False\n\n# if True, run inference on test dataset\nis_inference = True\n\n# if True and `is_train` == True, `PRETRAINED_PATH` model is trained further\nis_resume_training = False\n\n# Perform augmentation\nis_augment = False\n\nSEED = 42","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:43.891644Z","iopub.execute_input":"2023-08-06T15:19:43.892020Z","iopub.status.idle":"2023-08-06T15:19:43.898261Z","shell.execute_reply.started":"2023-08-06T15:19:43.891989Z","shell.execute_reply":"2023-08-06T15:19:43.896445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining Path\n","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\n\nTEST_IMG_DIR = Path(\"/kaggle/input/dlsprint2/badlad/images/test\")\n\nTEST_METADATA_PATH = Path(\"/kaggle/input/dlsprint2/badlad/badlad-test-metadata.json\")\n\n# Training output directory\nOUTPUT_DIR = Path(\"./output\")\nOUTPUT_MODEL = OUTPUT_DIR/\"model_final.pth\"\n\n# Path to your pretrained model weights\nPRETRAINED_PATH = Path(\"\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:47.209950Z","iopub.execute_input":"2023-08-06T15:19:47.210653Z","iopub.status.idle":"2023-08-06T15:19:47.216564Z","shell.execute_reply.started":"2023-08-06T15:19:47.210592Z","shell.execute_reply":"2023-08-06T15:19:47.215423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## JSON Load\n","metadata":{}},{"cell_type":"code","source":"with TEST_METADATA_PATH.open() as f:\n    test_dict = json.load(f)\n\nprint(\"#### LABELS AND METADATA LOADED ####\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:51.139716Z","iopub.execute_input":"2023-08-06T15:19:51.140660Z","iopub.status.idle":"2023-08-06T15:19:51.208735Z","shell.execute_reply.started":"2023-08-06T15:19:51.140600Z","shell.execute_reply":"2023-08-06T15:19:51.207755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Organizing COCO\n","metadata":{}},{"cell_type":"code","source":"def organize_coco_data(data_dict: dict) -> tuple[list[str], list[dict], list[dict]]:\n    thing_classes: list[str] = []\n\n    # Map Category Names to IDs\n    for cat in data_dict['categories']:\n        thing_classes.append(cat['name'])\n\n    # Images\n    images_metadata: list[dict] = data_dict['images']\n\n    # Convert COCO annotations to detectron2 annotations format\n    data_annotations = []\n    for ann in data_dict['annotations']:\n        # coco format -> detectron2 format\n        annot_obj = {\n            # Annotation ID\n            \"id\": ann['id'],\n\n            # Segmentation Polygon (x, y) coords\n            \"gt_masks\": ann['segmentation'],\n\n            # Image ID for this annotation (Which image does this annotation belong to?)\n            \"image_id\": ann['image_id'],\n\n            # Category Label (0: paragraph, 1: text box, 2: image, 3: table)\n            \"category_id\": ann['category_id'],\n\n            \"x_min\": ann['bbox'][0],  # left\n            \"y_min\": ann['bbox'][1],  # top\n            \"x_max\": ann['bbox'][0] + ann['bbox'][2],  # left+width\n            \"y_max\": ann['bbox'][1] + ann['bbox'][3]  # top+height\n        }\n        data_annotations.append(annot_obj)\n\n    return thing_classes, images_metadata, data_annotations","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:52.827283Z","iopub.execute_input":"2023-08-06T15:19:52.827693Z","iopub.status.idle":"2023-08-06T15:19:52.836290Z","shell.execute_reply.started":"2023-08-06T15:19:52.827661Z","shell.execute_reply":"2023-08-06T15:19:52.835086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"thing_classes_test, images_metadata_test, _ = organize_coco_data(\n    test_dict\n)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:54.266557Z","iopub.execute_input":"2023-08-06T15:19:54.267331Z","iopub.status.idle":"2023-08-06T15:19:54.272309Z","shell.execute_reply.started":"2023-08-06T15:19:54.267292Z","shell.execute_reply":"2023-08-06T15:19:54.271165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_metadata = pd.DataFrame(images_metadata_test)\ntest_metadata = test_metadata[['id', 'file_name', 'width', 'height']]\ntest_metadata = test_metadata.rename(columns={\"id\": \"image_id\"})\nprint(\"test_metadata size=\", len(test_metadata))\ntest_metadata.head(5)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:54.522196Z","iopub.execute_input":"2023-08-06T15:19:54.522564Z","iopub.status.idle":"2023-08-06T15:19:54.587440Z","shell.execute_reply.started":"2023-08-06T15:19:54.522534Z","shell.execute_reply":"2023-08-06T15:19:54.586413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Registering Data\n","metadata":{}},{"cell_type":"code","source":"DATA_REGISTER_TEST     = \"badlad_test\"\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:55.576747Z","iopub.execute_input":"2023-08-06T15:19:55.577100Z","iopub.status.idle":"2023-08-06T15:19:55.581784Z","shell.execute_reply.started":"2023-08-06T15:19:55.577073Z","shell.execute_reply":"2023-08-06T15:19:55.580671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Detectron2 Format","metadata":{}},{"cell_type":"code","source":"def convert_coco_to_detectron2_format(\n    imgdir: Path,\n    metadata_df: pd.DataFrame,\n    annot_df: Optional[pd.DataFrame] = None,\n    target_indices: Optional[np.ndarray] = None,\n):\n\n    dataset_dicts = []\n    for _, train_meta_row in tqdm(metadata_df.iterrows(), total=len(metadata_df)):\n        # Iterate over each image\n        image_id, filename, width, height = train_meta_row.values\n\n        annotations = []\n\n        # If train/validation data, then there will be annotations\n        if annot_df is not None:\n            for _, ann in annot_df.query(\"image_id == @image_id\").iterrows():\n                # Get annotations of current iteration's image\n                class_id = ann[\"category_id\"]\n                gt_masks = ann[\"gt_masks\"]\n                bbox_resized = [\n                    float(ann[\"x_min\"]),\n                    float(ann[\"y_min\"]),\n                    float(ann[\"x_max\"]),\n                    float(ann[\"y_max\"]),\n                ]\n\n                annotation = {\n                    \"bbox\": bbox_resized,\n                    \"bbox_mode\": BoxMode.XYXY_ABS,\n                    \"segmentation\": gt_masks,\n                    \"category_id\": class_id,\n                }\n\n                annotations.append(annotation)\n\n        # coco format -> detectron2 format dict\n        record = {\n            \"file_name\": str(imgdir/filename),\n            \"image_id\": image_id,\n            \"width\": width,\n            \"height\": height,\n            \"annotations\": annotations\n        }\n\n        dataset_dicts.append(record)\n\n    if target_indices is not None:\n        dataset_dicts = [dataset_dicts[i] for i in target_indices]\n\n    return dataset_dicts","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:56.255242Z","iopub.execute_input":"2023-08-06T15:19:56.255946Z","iopub.status.idle":"2023-08-06T15:19:56.267172Z","shell.execute_reply.started":"2023-08-06T15:19:56.255910Z","shell.execute_reply":"2023-08-06T15:19:56.265983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Register Test Inference data\nDatasetCatalog.register(\n    DATA_REGISTER_TEST,\n    lambda: convert_coco_to_detectron2_format(\n        TEST_IMG_DIR,\n        test_metadata,\n    )\n)\n\n# Set Test data categories\nMetadataCatalog.get(DATA_REGISTER_TEST).set(\n    thing_classes=thing_classes_test\n)\n\ndataset_dicts_test = DatasetCatalog.get(DATA_REGISTER_TEST)\nmetadata_dicts_test = MetadataCatalog.get(DATA_REGISTER_TEST)\n\nprint(\"dicts test size=\", len(dataset_dicts_test))\nprint(\"################\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:56.803247Z","iopub.execute_input":"2023-08-06T15:19:56.803635Z","iopub.status.idle":"2023-08-06T15:19:57.603258Z","shell.execute_reply.started":"2023-08-06T15:19:56.803584Z","shell.execute_reply":"2023-08-06T15:19:57.602183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Downloading Model Weight and Configs\n","metadata":{}},{"cell_type":"code","source":"#better to use gdown to fetch from drive\n!gdown 1OkOEy7ZoF7Hmd24wAlvzVb7cEszkcrvk  #model weight\n\n!gdown 1CwIgwAFY4s7Nz-ST7Al2KGL1qtrlIhFx  #config of layoutlmv3","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:19:59.537390Z","iopub.execute_input":"2023-08-06T15:19:59.537772Z","iopub.status.idle":"2023-08-06T15:20:13.424117Z","shell.execute_reply.started":"2023-08-06T15:19:59.537741Z","shell.execute_reply":"2023-08-06T15:20:13.422901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setting Model Path\n","metadata":{}},{"cell_type":"code","source":"MODEL_PATH=Path(\"/kaggle/working/final_train_layoutlmmv3.pth\")\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:13.426783Z","iopub.execute_input":"2023-08-06T15:20:13.427450Z","iopub.status.idle":"2023-08-06T15:20:13.432578Z","shell.execute_reply.started":"2023-08-06T15:20:13.427413Z","shell.execute_reply":"2023-08-06T15:20:13.431648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setting Test Hyperparameters\n","metadata":{}},{"cell_type":"code","source":"inf_cfg = get_cfg()\n\nadd_vit_config(inf_cfg)\n# Load a config from file\ninf_cfg.merge_from_file(\"/kaggle/working/unilm/layoutlmv3/examples/object_detection/cascade_layoutlmv3.yaml\")\ninf_cfg.MODEL.CONFIG_PATH=\"/kaggle/working/config.json\"\ninf_cfg.MODEL.ROI_HEADS.BATCH_SIZE_PER_IMAGE = 128\ninf_cfg.MODEL.ROI_HEADS.NUM_CLASSES = 4\ninf_cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = 0.5\ninf_cfg.MODEL.DEVICE = \"cuda\"\n\ninf_cfg.DATALOADER.NUM_WORKERS = 1  # lower this if CUDA overflow occurs\ninf_cfg.MODEL.WEIGHTS = str(MODEL_PATH)\nBATCH = 1 # lower this if CUDA overflow occurs\ntest_loader = build_detection_test_loader(inf_cfg, DATA_REGISTER_TEST, batch_size=BATCH)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:13.434509Z","iopub.execute_input":"2023-08-06T15:20:13.435425Z","iopub.status.idle":"2023-08-06T15:20:14.645005Z","shell.execute_reply.started":"2023-08-06T15:20:13.435392Z","shell.execute_reply":"2023-08-06T15:20:14.644081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#set acceptance threshold to 0.5\nACCEPTANCE_THRESHOLD = 0.5  # for all categories","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:14.647733Z","iopub.execute_input":"2023-08-06T15:20:14.648155Z","iopub.status.idle":"2023-08-06T15:20:14.652690Z","shell.execute_reply.started":"2023-08-06T15:20:14.648120Z","shell.execute_reply":"2023-08-06T15:20:14.651487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"#### MODEL: {inf_cfg.MODEL.WEIGHTS} FOR INFERENCE ####\")\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:14.654100Z","iopub.execute_input":"2023-08-06T15:20:14.654737Z","iopub.status.idle":"2023-08-06T15:20:14.664065Z","shell.execute_reply.started":"2023-08-06T15:20:14.654703Z","shell.execute_reply":"2023-08-06T15:20:14.662973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rebuild_model():\n    model = build_model(inf_cfg)\n    _ = DetectionCheckpointer(model).load(inf_cfg.MODEL.WEIGHTS)\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:14.665589Z","iopub.execute_input":"2023-08-06T15:20:14.666166Z","iopub.status.idle":"2023-08-06T15:20:14.674145Z","shell.execute_reply.started":"2023-08-06T15:20:14.666134Z","shell.execute_reply":"2023-08-06T15:20:14.673209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = rebuild_model()\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:14.675726Z","iopub.execute_input":"2023-08-06T15:20:14.676058Z","iopub.status.idle":"2023-08-06T15:20:22.277151Z","shell.execute_reply.started":"2023-08-06T15:20:14.676025Z","shell.execute_reply":"2023-08-06T15:20:22.276109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CUDA Problems\n","metadata":{}},{"cell_type":"code","source":"!export LRU_CACHE_CAPACITY=1\n!export 'PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:512'","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:22.279428Z","iopub.execute_input":"2023-08-06T15:20:22.280024Z","iopub.status.idle":"2023-08-06T15:20:24.262099Z","shell.execute_reply.started":"2023-08-06T15:20:22.279987Z","shell.execute_reply":"2023-08-06T15:20:24.260786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vars_to_del = [\"trainer\", \"predictor\", \"outputs\"]\n\nfor v in vars_to_del:\n    if v in globals():\n        print(f\"Deleting {v}\")\n        del globals()[v]\n    elif v in locals():\n        print(f\"Deleting {v}\")\n        del locals()[v]","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:24.264920Z","iopub.execute_input":"2023-08-06T15:20:24.265747Z","iopub.status.idle":"2023-08-06T15:20:24.272634Z","shell.execute_reply.started":"2023-08-06T15:20:24.265705Z","shell.execute_reply":"2023-08-06T15:20:24.271480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference Utils\n","metadata":{}},{"cell_type":"code","source":"def rle_encode(mask):\n#     print(mask)\n    pixels = mask.T.flatten()\n    use_padding = False\n    if pixels[0] or pixels[-1]:\n        use_padding = True\n        pixel_padded = np.zeros([len(pixels) + 2], dtype=pixels.dtype)\n        pixel_padded[1:-1] = pixels\n        pixels = pixel_padded\n    rle = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    if use_padding:\n        rle = rle - 1\n    rle[1::2] = rle[1::2] - rle[:-1:2]\n    return ' '.join(str(x) for x in rle)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:24.276797Z","iopub.execute_input":"2023-08-06T15:20:24.277103Z","iopub.status.idle":"2023-08-06T15:20:24.287135Z","shell.execute_reply.started":"2023-08-06T15:20:24.277078Z","shell.execute_reply":"2023-08-06T15:20:24.286149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@retry_if_cuda_oom\ndef get_masks(prediction):\n    # get masks for each category\n    take = prediction.scores >= ACCEPTANCE_THRESHOLD\n    pred_masks = (prediction.pred_masks[take] != 0)\n    pred_classes = prediction.pred_classes[take]\n  \n    rles = []\n    for cat in range(len(thing_classes_test)):\n        pred_mask = pred_masks[pred_classes == cat]\n        \n        pred_mask = retry_if_cuda_oom(torch.any)(pred_mask, dim=0)\n#         pred_mask = torch.any(pred_mask, dim=0)\n        rles.append(rle_encode(pred_mask.short().to(\"cpu\").numpy()))\n\n    return rles","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:24.288773Z","iopub.execute_input":"2023-08-06T15:20:24.289257Z","iopub.status.idle":"2023-08-06T15:20:24.299505Z","shell.execute_reply.started":"2023-08-06T15:20:24.289223Z","shell.execute_reply":"2023-08-06T15:20:24.298553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_inference(data):\n    results = []\n    with torch.no_grad():\n        outputs = model(data)\n        if torch.cuda.is_available():\n            torch.cuda.synchronize()\n\n        for idx, output in enumerate(outputs):\n            output = output[\"instances\"]\n\n            rles = get_masks(output)\n\n            result = [\n                f\"{data[idx]['image_id']}_{cat},{rles[cat]}\\n\"\n                for cat in range(len(thing_classes_test))\n            ]\n\n            results.extend(result)\n\n        del outputs, output\n\n    return results","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:24.300876Z","iopub.execute_input":"2023-08-06T15:20:24.301602Z","iopub.status.idle":"2023-08-06T15:20:24.310904Z","shell.execute_reply.started":"2023-08-06T15:20:24.301567Z","shell.execute_reply":"2023-08-06T15:20:24.309873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Running Inference on Test Data and Creating Submission File\n","metadata":{}},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:27.860101Z","iopub.execute_input":"2023-08-06T15:20:27.860466Z","iopub.status.idle":"2023-08-06T15:20:28.185497Z","shell.execute_reply.started":"2023-08-06T15:20:27.860436Z","shell.execute_reply":"2023-08-06T15:20:28.184449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if is_inference:\n    model.eval()\n    submission_file = open(\"submission.csv\", \"w\")\n    submission_file.write(\"Id,Predicted\\n\")\n\n    results: list[str] = []\n    \n    for i, data in enumerate(tqdm(test_loader)):\n        res = run_inference(data)\n        results.extend(res)\n        \n        if i % (500 // BATCH) == 0:\n            print(f\"Inference on batch {i}/{len(test_loader)} done\")\n            submission_file.writelines(results)\n            results = []\n\n    submission_file.writelines(results)\n    submission_file.close()","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:20:29.999456Z","iopub.execute_input":"2023-08-06T15:20:30.000168Z","iopub.status.idle":"2023-08-06T15:21:04.835378Z","shell.execute_reply.started":"2023-08-06T15:20:30.000132Z","shell.execute_reply":"2023-08-06T15:21:04.833963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if Path(\"submission.csv\").exists:\n    display(FileLink(\"submission.csv\"))","metadata":{},"execution_count":null,"outputs":[]}]}