{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"#### This notebook was inspired from [here](https://www.kaggle.com/code/ammarnassanalhajali/layout-parser-model-training)","metadata":{"_uuid":"2a9e103a-9be5-4118-ace4-e429b1c32450","_cell_guid":"b8a4d973-2613-48d4-83fd-2e9d65cc5e91","trusted":true}},{"cell_type":"markdown","source":"# Detectron2\n\n[Detectron2](https://detectron2.readthedocs.io/en/latest/index.html) is a popular open-source software library developed by Facebook AI Research (FAIR) for building computer vision models. It serves as a powerful framework for object detection, instance segmentation, and keypoint detection tasks. Detectron2 is built on top of PyTorch, geared towards a more convenient way to build modular, flexible pipelines for specific Computer Vision Tasks such as object detection, instance segmentation. \n\nDetectron2 has a collection of trained models for these tasks in their [model zoo](https://github.com/facebookresearch/detectron2/blob/main/MODEL_ZOO.md). We can also use detectron2 to train pre-implemented state-of-the-art models from scratch for new datasets, as we do in this notebook. \n\nRead the [documentation](https://detectron2.readthedocs.io/en/latest/index.html).","metadata":{"_uuid":"02db164f-14fe-473f-9360-3528dc711121","_cell_guid":"acaa7248-40b8-4c14-9e01-ccea0d785f9b","trusted":true}},{"cell_type":"markdown","source":"# 1 Install detectron2","metadata":{"_uuid":"dcbab65a-2deb-44b8-9110-956e73ae7f31","_cell_guid":"7f55d342-6fbb-4797-82e6-b679e324be3c","trusted":true}},{"cell_type":"markdown","source":"## 1.1 Recommended Way (is not working on kaggle)","metadata":{"_uuid":"38f68376-6839-49d5-8be7-f0bf92f9e747","_cell_guid":"bb56b025-c2e8-4555-8e32-52d13c05128e","trusted":true}},{"cell_type":"code","source":"# !python -m pip install 'git+https://github.com/facebookresearch/detectron2.git'","metadata":{"_uuid":"7d1fb578-8ef5-4e24-bd21-c65ecb29d5de","_cell_guid":"be83a0e9-a75a-476d-b7fa-c93d19c3f838","collapsed":false,"_kg_hide-input":false,"_kg_hide-output":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:512\"","metadata":{"_uuid":"ccf71dcd-74a9-4f25-81c6-3ab00698c859","_cell_guid":"eb606811-bf2e-45e4-a167-08a90af1cab4","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1.2 Fast Way\nIgnore the warnings.","metadata":{"_uuid":"ea41b1ca-563d-482f-871c-ecb44dfd1411","_cell_guid":"9cb7c8f7-2854-475b-a217-53d75602d543","trusted":true}},{"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":{"_uuid":"6cbb4532-5e4a-4472-93ee-668b7b59173e","_cell_guid":"fcc9bb73-4ea3-4a1f-9bfd-60d6238ee3cf","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2 Notebook Config","metadata":{"_uuid":"a4259cfb-08a6-4cc2-8f3f-c88307bc957e","_cell_guid":"3773d64b-e163-44f4-a67b-d1aa735ea789","trusted":true}},{"cell_type":"code","source":"\nimport os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:512\"\n\n\nfrom 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 = True\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\nimport random\nimport os\nimport numpy as np\nimport torch\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(SEED)\n\n\"\"\"## 2.2 Paths\"\"\"\n\nfrom pathlib import Path\n\nTRAIN_IMG_DIR = Path(\"/kaggle/input/binarizedbadlad/binTrain\")\n\nTRAIN_COCO_PATH = Path(\"/kaggle/input/dlsprint2/badlad/labels/coco_format/train/badlad-train-coco.json\")\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(\"\")\n\n\"\"\"## 2.3 imports\"\"\"\n\n# detectron2\nfrom 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\n\n# torch\nimport torch\nimport os\n\nimport gc\n\nimport warnings\n# Ignore \"future\" warnings and Data-Frame-Slicing warnings.\nwarnings.filterwarnings('ignore')\n\nsetup_logger()\n\nimport os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:512\"\n\n\"\"\"# 3 COCO Annotations Data\n\n## 3.1 Load\n\"\"\"\n\nwith TEST_METADATA_PATH.open() as f:\n    test_dict = json.load(f)\n\n\nprint(\"#### LABELS AND METADATA LOADED ####\")\n\n\"\"\"## 3.2 Observe\"\"\"\n\ndef organize_coco_data(data_dict: 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    print(thing_classes)\n\n    # thing_classes = ['paragraph', 'text_box', 'image', 'table']\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\n\n\nthing_classes_test, images_metadata_test, _ = organize_coco_data(\n    test_dict\n)\n\nthing_classes = thing_classes_test\nprint(\"THINGS CLASSES\")\nprint(thing_classes)\ntest_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)\n\ndef 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\n\n\n\"\"\"## 4.3 Registering and Loading Data for `detectron2`\"\"\"\nDATA_REGISTER_TEST     = \"badlad_test\"\n\n# 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\n# dataset_dicts_test = DatasetCatalog.get(DATA_REGISTER_TEST)\nmetadata_dicts_test = MetadataCatalog.get(DATA_REGISTER_TEST)","metadata":{"_uuid":"a2b86369-4dc4-4c48-8a0a-24f1dbed9b4c","_cell_guid":"f5920213-4dd4-4d65-8d7d-ee0a7b3d9343","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"_uuid":"0348b282-e449-4dc4-8bda-5aff605352c5","_cell_guid":"c2b93beb-b981-471f-8976-19d051462ad8","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! git clone https://github.com/microsoft/unilm.git --depth=1 --quiet\n! sed -i 's/from collections import Iterable/from collections.abc import Iterable/' unilm/dit/object_detection/ditod/table_evaluation/data_structure.py","metadata":{"_uuid":"bbfa1340-e1f6-4b4a-9448-c7082cdbfa96","_cell_guid":"6f0c5632-81a2-4ee1-836c-f23fd97c13ef","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"unilm\")\n\nimport cv2\n\nfrom unilm.dit.object_detection.ditod import add_vit_config","metadata":{"_uuid":"e615baeb-945f-4063-a93c-963a5be33a72","_cell_guid":"da69737a-2f69-465c-a681-6066e6be6ba2","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile cascade_dit_base.yaml\n_BASE_: \"/kaggle/input/dit-publay-finetuned/Base-RCNN-FPN.yaml\"\nMODEL:\n  PIXEL_MEAN: [ 127.5, 127.5, 127.5 ]\n  PIXEL_STD: [ 127.5, 127.5, 127.5 ]\n  WEIGHTS: \"https://layoutlm.blob.core.windows.net/dit/dit-pts/dit-base-224-p16-500k-62d53a.pth\"\n  VIT:\n    NAME: \"dit_base_patch16\"\n  ROI_HEADS:\n    NAME: CascadeROIHeads\n  ROI_BOX_HEAD:\n    CLS_AGNOSTIC_BBOX_REG: True\n  RPN:\n    POST_NMS_TOPK_TRAIN: 2000\nSOLVER:\n  WARMUP_ITERS: 1000\n  IMS_PER_BATCH: 16\n  MAX_ITER: 60000\n  CHECKPOINT_PERIOD: 2000\nTEST:\n  EVAL_PERIOD: 2000","metadata":{"_uuid":"19abaefc-d466-49e4-b9ca-8be2e5873d65","_cell_guid":"d364d0a9-de62-4227-a2fb-a0a346ad49a6","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"_uuid":"e3fce8d5-dd99-460f-936e-7c5ad6fee623","_cell_guid":"50a2d031-2d3d-4ca0-bf92-4ccc174d3f19","collapsed":false,"jupyter":{"outputs_hidden":false},"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","metadata":{"_uuid":"44c4ad7d-f5c6-4eb4-8493-16fbc977d4aa","_cell_guid":"b4799c1b-a888-4f8d-80fb-1cccd59d0377","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inf_cfg = get_cfg()\nadd_vit_config(inf_cfg)\ninf_cfg.merge_from_file(\"/kaggle/working/cascade_dit_base.yaml\")\ninf_cfg.SOLVER.IMS_PER_BATCH = 64\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.25\ninf_cfg.MODEL.DEVICE = \"cuda\"\ninf_cfg.DATALOADER.NUM_WORKERS = 2  # lower this if CUDA overflow occurs\ninf_cfg.MODEL.WEIGHTS = str(\"/kaggle/input/dit-publay-finetuned/dit-pub-50000.pth\")\ninf_cfg.OUTPUT_DIR = str(OUTPUT_DIR)\nprint(\"creating cfg.OUTPUT_DIR -> \", inf_cfg.OUTPUT_DIR)\nOUTPUT_DIR.mkdir(exist_ok=True)\nmodel = rebuild_model()\nmodel.eval()","metadata":{"_uuid":"5d28c6a4-a958-4594-bd87-655b95b6f11a","_cell_guid":"b15ce3c2-639c-4bb7-ad55-144f10baa645","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def binarize(image):\n    # Convert image to grayscale if it has more than one channel\n    if len(image.shape) > 2 and image.shape[2] in [3, 4]:\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n\n    # Binarize the grayscale image\n    _, binary_image = cv2.threshold(image, 128, 255, cv2.THRESH_BINARY)\n\n    # Convert the binary image back to 3-channel format (optional, but could be useful for consistency)\n    binary_image = cv2.cvtColor(binary_image, cv2.COLOR_GRAY2RGB)\n\n    return binary_image\n","metadata":{"_uuid":"535a6cad-011c-4809-bc82-dc25b01193f7","_cell_guid":"48f12f79-977e-4184-93f3-afe2ce2ea1a5","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BinarizeDatasetMapper(DatasetMapper):\n    def __init__(self, cfg, is_train=True):\n        super().__init__(cfg, is_train)\n\n    def __call__(self, dataset_dict):\n        dataset_dict = super().__call__(dataset_dict)\n        image = dataset_dict[\"image\"].permute(1, 2, 0).cpu().numpy()\n        binarized_image = binarize(image)\n        dataset_dict[\"image\"] = torch.tensor(binarized_image.transpose(2, 0, 1)).to(dataset_dict[\"image\"].device)\n        return dataset_dict","metadata":{"_uuid":"3ba6e993-cd55-47b4-bc92-91c055b453de","_cell_guid":"47a1e74a-1092-414b-9b19-2f03f48cad5d","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(batch):\n    images = [data[\"image\"].cpu().numpy() for data in batch]\n    images = [binarize(image) for image in images]\n    images = torch.stack([torch.tensor(image.transpose(2, 0, 1)) for image in images])\n    batch[0][\"image\"] = images\n    return batch\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH = 2  # lower this if CUDA overflow occurs\n# mapper = BinarizeDatasetMapper(inf_cfg, is_train=False)\ntest_loader = build_detection_test_loader(inf_cfg, DATA_REGISTER_TEST,batch_size=BATCH)","metadata":{"_uuid":"b9b60d2d-3772-4857-8b74-83e0330aef2a","_cell_guid":"b74cfd88-33f9-48e9-884e-cf70be8f869c","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ACCEPTANCE_THRESHOLD = 0.6","metadata":{"_uuid":"3a04ea15-7fd9-4efd-a66a-3101601de529","_cell_guid":"9de4d55c-c58d-45ed-b069-be8c81164f8d","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:512\"\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode(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":{"_uuid":"da2c44bc-1e96-45d2-923e-c27bad9995e2","_cell_guid":"302679e7-c6e3-4221-a127-e91f8b9c2f8a","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ACCEPTANCE_THRESHOLDS = {\n    \"paragraph\": 0.5,\n    \"text_box\": 0.3,\n    \"image\": 0.5,\n    \"table\": 0.55,\n}\n\n# @retry_if_cuda_oom\n# def get_masks(prediction):\n#     # get masks for each category\n#     pred_masks = (prediction.pred_masks != 0)\n#     pred_classes = prediction.pred_classes\n\n#     rles = []\n#     for cat in range(len(thing_classes_test)):\n#         pred_mask = pred_masks[pred_classes == cat]\n#         pred_mask = torch.any(pred_mask, dim=0)\n        \n#         threshold = ACCEPTANCE_THRESHOLDS[thing_classes[cat]]\n#         take = prediction.scores >= threshold\n#         pred_mask = pred_mask & take\n        \n#         rles.append(rle_encode(pred_mask.short().to(\"cpu\").numpy()))\n\n#     return rles\n\ndef get_masks(prediction):\n    # get masks for each category\n    rles = []\n    for cat in range(len(thing_classes)):\n        threshold = ACCEPTANCE_THRESHOLDS.get(thing_classes[cat], 0.4)  # Get threshold or set to 0.4 if not present\n        if threshold==0.4:\n            print(\"thresh : 0.4\")\n        take = prediction.scores >= threshold\n        pred_masks = (prediction.pred_masks[take] != 0)\n        pred_classes = prediction.pred_classes[take]\n        \n        pred_mask = torch.any(pred_masks[pred_classes == cat], dim=0)\n        rles.append(rle_encode(pred_mask.short().to(\"cpu\").numpy()))\n\n    return rles\n\n\n# def 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)):\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":{"_uuid":"32fe4ae1-f0a5-4ef5-8318-4c54c04ebc7e","_cell_guid":"f0521c26-8932-4e03-87c6-3f128065a1ef","collapsed":false,"jupyter":{"outputs_hidden":false},"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))\n            ]\n\n            results.extend(result)\n\n        del outputs, output\n\n    return results","metadata":{"_uuid":"0a716fc2-7c5d-45a3-a30a-735bb994cf8c","_cell_guid":"dea5f5f6-e9b9-46a0-aa5f-2da6b6518a59","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_file = open(\"submission.csv\", \"w\")\nsubmission_file.write(\"Id,Predicted\\n\")\n\nresults: list[str] = []\n\nfor 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\nsubmission_file.writelines(results)\nsubmission_file.close()","metadata":{"_uuid":"aa945c38-e2bd-4cc8-a2af-9fb270ac9c9b","_cell_guid":"d3a579d0-d2b3-4acf-9eb6-5a9ac199ee3b","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if Path(\"submission.csv\").exists:\n    display(FileLink(\"submission.csv\"))","metadata":{"_uuid":"517a1805-46b6-4fb6-9372-31f7661c97be","_cell_guid":"1185d525-562c-475d-b98b-f58c96d82ea2","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r detectron2/","metadata":{"_uuid":"26cdb869-5bee-404f-8581-0f6899a22223","_cell_guid":"0ed19be4-de35-495b-8ffd-aa4c950752f6","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}