{"cells":[{"metadata":{},"cell_type":"markdown","source":"Please consider upvoting the dataset: https://www.kaggle.com/raddar/vinbigdata-competition-jpg-data-2x-downsampled"},{"metadata":{},"cell_type":"markdown","source":"#### Detectron2 library can be used to quickly train object detection models. \nIn this kernel, I will be refactoring code from detectron2's official github for this. Most of my code here is from: https://github.com/facebookresearch/detectron2/blob/master/tools/plain_train_net.py\n\nHere, I am trying to make custom training loop without using detectrons's built in trainer. This would give me more control over the training.\n\nAt around 15,000 epochs, this model can give 0.2 lb score"},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport glob\n%matplotlib inline","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"#### Object detection models require very large training iterations. Kaggle kernels have limited time allowed for continuous running of gpus. So I will be stopping the training when execution time approaches 3 hours"},{"metadata":{"trusted":true},"cell_type":"code","source":"# set start time\nimport time\nkernel_start_time = time.time()","execution_count":null,"outputs":[]},{"metadata":{"_kg_hide-output":true,"trusted":true},"cell_type":"code","source":"!pip install detectron2 -f \\\n  https://dl.fbaipublicfiles.com/detectron2/wheels/cu102/torch1.7/index.html","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### **** **Import everything**"},{"metadata":{"trusted":true},"cell_type":"code","source":"import detectron2\nimport cv2\n\n\nimport logging\nfrom collections import OrderedDict\nimport torch\nfrom torch.nn.parallel import DistributedDataParallel\n\nimport detectron2.utils.comm as comm\nfrom detectron2.checkpoint import DetectionCheckpointer, PeriodicCheckpointer\nfrom detectron2.config import get_cfg\nfrom detectron2.data import (\n    MetadataCatalog,\n    build_detection_test_loader,\n    build_detection_train_loader,\n)\nfrom detectron2.engine import default_argument_parser, default_setup, launch\nfrom detectron2.evaluation import (\n    COCOEvaluator,\n    inference_on_dataset,\n    print_csv_format,\n)\nfrom detectron2.modeling import build_model\nfrom detectron2.solver import build_lr_scheduler, build_optimizer\nfrom detectron2.utils.events import (\n    CommonMetricPrinter,\n    EventStorage,\n    JSONWriter,\n    TensorboardXWriter,\n)\n\nfrom detectron2.data import DatasetCatalog, MetadataCatalog\nfrom tqdm import tqdm\nimport numpy as np\nfrom detectron2 import model_zoo\nfrom detectron2.structures import BoxMode\nimport detectron2.utils.comm as comm\nfrom detectron2.engine import DefaultPredictor","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Configs"},{"metadata":{"trusted":true},"cell_type":"code","source":"# configs\nDEBUG=False\nTRAIN_CSV_PATH = \"../input/vinbigdata-competition-jpg-data-2x-downsampled/train_downsampled.csv\"\nTRAIN_FOLDER = \"../input/vinbigdata-competition-jpg-data-2x-downsampled/train/train\"\nEPOCHS =30000000 # set to high value. Training stops after around 2.5 hours\nRUNTIME = 2.5*60*60 # SECONDS\n# RUNTIME = 1*60*60\nEVALUATE_EVERY_N_STEPS = 15\nOUTPUT_DIR=\"outputs\"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"traindf = pd.read_csv(TRAIN_CSV_PATH)\ntraindf.sample(2)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Validation split"},{"metadata":{"trusted":true},"cell_type":"code","source":"from sklearn import model_selection\ngkf = model_selection.GroupKFold(5)\nfor train_index, test_index in gkf.split(X=traindf,  groups=traindf[\"image_id\"]):\n    dftrain_val = traindf.loc[test_index,:]\n    dftrain_tr = traindf.loc[train_index,:]\n    break\n\ndftrain_tr.shape,dftrain_val.shape","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Format data\nDetectron2 data prep: create list of dicts  \n![image.png](attachment:image.png)","attachments":{"image.png":{"image/png":"iVBORw0KGgoAAAANSUhEUgAAAW0AAAFuCAYAAABKlUmMAAAgAElEQVR4Ae2dPY/cxrKGz3+RA2n/gbODm1xFBziBU8MwBGOhWHByAjtRIhxIuQMHSi4EGBsoVqjAmzhxKEWbKFc0GS+62VVTbDY/ZoYzU00+wYLcIdldXfX20zU9ZPMfj588avjDB2gADaCBOjTwDwJVR6CIE3FCA2ggaKAH7du7z81u9zX9fW7ubr8hE+fbCBpAA2jAiQbK0H5419zeAGsyGzIbNIAGvGkAaDsZPb0JA3uAFRrwqQGgDbT52osG0EBFGgDaFQWLzMdn5kNciMslNQC0gTZZFhpAAxVpAGhXFKxLjubURfaIBnxqAGgDbbIsNIAGKtIA0K4oWGQ+PjMf4kJcLqkBoA20ybLQABqoSANAu6JgXXI0py6yRzTgUwNAG2iTZaEBNFCRBoB2RcEi8/GZ+RAX4nJJDZShzYJRZB4MZmgADbjUQA/alxwxqIsMBQ2gATRwmAaANtmEy2yCjnxYR8Zf2/EX0AbaQBsNoIGKNAC0KwoW2dR2siliTayHNFCG9s2z5u6hfXvNw90zRmHAjgbQABpwooEitOMrx3h7DSJ1ItKhjIPPyUa3qIEitF/ff23IsOkQW+wQtBnde9dAH9ppagRoI17v4sU+NLpFDQBtpgCYBkIDaKAiDRSg/aq5331u7m55G/sWR3HaTPaKBnxrYA9tvWPkY/P6BmAjXN/CJT7EZ6sa2ENbvh7ckGlvVQy0GxCiAf8aKEC7vUebHyL9B48ORozQwPY0ALTlGwZbfoxCA2igAg0A7QqCRDa1vWyKmBPzIQ30of3kUcPDNQhmSDB8jjbQwHU1UIQ2j7FfNyh0CvyPBtDAkAaK0H6st//xOPuQ4/icToUG0MA1NFCGNvO8/CCDBtAAGnCpAaCNMF0K8xoZDHWSOdegAaANtIE2GkADFWlAof2/T/+HwFUUuBoyAmwkc0UDy2tAoY1zl3cuPsWnaAANLK0BhbabTPuXH5ovn140H375pnls9/Ms+Obb5u2fL+K54fwvf37X/GgXuhq6Nr/u/dN1fsOQdka/JF/lPsp9yv/r1AJxXVVcFdpLjwZHlxdh+7x5+71AO+1nwnv5vgBqe84QtM05sYy1Qzu0TwC+1raamB6tO8pYFdjWrAOFtptM+/vvmr8/JVDbfdupEoT+/u3bk4S2amg/edTY9tn9NQuatjEdsXYNKLSraugYtCPo99MmcZrFAt/sD4JMMtMw7VKaejFllPz242/Pmy/vnzZxK2XkWe7N0+aDHIvbH5qXaXqnvf47nf758Iuc2/3W0SlfBroJ20r2nvpZWPZgx4szTkogTo0B129nsFJou8m0x6CTATkCNYGvB+cExd7npvwitAXYBrLxvAPmgxWmUkay29ry8v0e0vn0hVwfzo91RyC389Ly7aI9x5Rhp5VMGy/RmYH2doBxCT1Rx7ieFNpVOWos0xZgHQvtCFgDw1Be/Kyb5Y75KwLVQn6GvfYau28HlrAfoZ3Ks4OAgF+gPmYfx8Y7Bf7BP541oNCuItNWIHezzqKDj4V2+gHTZvHt/sLQLtWTQD8NbZku2U8Dib1AG+AU+4P0HbbVT2MptKsK9IzM9fGx0C5l2gcK3UI3+jW3tzBdYq+x+7Mz7QNtrCretK160KC35ZIJhTaZdnJqAmz4IfFYoVnoFqGdzz8niMu95vb6IrSfPEo/cmbTOFeC29Sc9k/Pm+aPX5vm9+9vyz7991/x+B8/v2/+Ze+1l/b8833z+6+hjL+an0rHdVXKz83dLS+lPla3XLccWM/pS4X2OStZvOw8c5XOrTAbmTYQQHbu3Mju+RZw23PsHLWpr9Q2C914vGBvhLGW/0PzMoB87vRIqj/Wo2WENs+fwinZfexnV4f2k0dNXAN+97XZ3b8qDwwTMTu27VxXB+jWFCeFdlWZNh0QMOUauH3XPOxY/31NcKIt5QFRoY2Dyg7CL3X4ZSrbJ451xJE4TcdJoU2mPe0sua1O7tQobbl7Y4Yf8yz5lP9Thr3bfWxel+a7Tymba/lG51ADCm1GuAvDxqEY0AAaQAP+NaDQJtP2Hyw6FDFCA2hAoY0YEAMaQANowL8GFNpk2v6DRYciRmgADSi0EQNiQANoAA3414BC202mndbkiIsh2f3sh7v2wZLrPEyymLDTo/Zznr7sP7CT1h8JT24eUM5itmfxoFz/nZ0YrSNGCm03AbWPeNv9DBJAu4V2u+qf2c/85Cau2MXtc2hgEQ0otN1k2vEx85RB2/2NB7yfaZuVDguPyQPrdWRVxJE45hpQaOcH3P6fpkyKy6XGaYLnzdvf2pcDh7U8Xoa3yIT1OewCUDKdoOt2ZAsv9Y63a5nYB2faTF/WODlwmiZf28TaJoNTfk6w9YD1Ty4ZP55GBCyX1NvW61Jou8m0BVpT21IWLrANcJOFoQIQI+j3YB57a4w+9SggLWSxLbD35bXlHwju1D67ip+KUYAtNshCWEB7ka+X6ucpjXEcfzvUgEK7OiGPQDv+iGmPx30D2SwQnakHydbD2+DTefG4ADQBdam3xhShXbC3Y2Nmv9jJlowXDaxfAwrtNWXak9DuTLGkKQ7JYvMsN2XvOjUi2bxOrcgUSXoV2IFALUI7+2YQOiLQXn9nBLjEeI4GFNpzTnZ1js2kBZQJqKPQjte9aGym3AGiQNtCWbLsUE8p05b6j9gCbTqqq351hIax/7IaVmhvJtOOWayZf04Q1x/5ClMTuSgj5D8NT7fk54/9X4S2HXwky674h0jeXHPZTj2mN47VHwuFdi3BjJCzWXDaj5mzhZ3NxDMQd8vovjUm+KF7PJs+SZlIC+791MhBb42RgSJvh0zRhDrsFE74PPxvjzvKiKbuHjk7tHlzjf7+Uks/xs7jBw+FdnWZ9rmgVZhPlicOdV77XHVT7vHw4c01x/sO3VXlO4U2I1878hWnPlJmbOfB8dfxmcI5fDeV7Z+jTsr0pYGtxEOhTaa9F2BpemQusEvXdt5wY3/UJMM5PcPhzTWn+xAdVuVDhfZWRinauR+c8AW+QAP1aUChTaZdX/DocMQMDWxPAwptgr+94BNzYo4G6tOAQptMu77g0eGIGRrYngYU2gR/e8En5sQcDdSnAYW2m0w7PVQS79aw++kX7nh3xhJ3YMgaIkeUtYQN7a2F7ZOZdp9OVF8nImbE7JIaUGhfstLRuiKo02Pmdr82aNsnMgu3VNn7we3+qG8K5XA+wEAD29KAQttNpm1hZ/eXhvYJAJyVaRds73SuOCCl9Uvs/gl2dcqnnKruvSV22wLvKfFWaJ9SyCWvFWDGbVq7o/d4eYTgfl2QzoMx+Sp+pekRmTrJ1gaRekZtiLDe120frJHrL+mvS9TF04gA5xI6o45WZwptN5n2RIYosFYA5lMo+f8Joh1wpzoEvh0xCNQF5ul/rc8sKKWf5XWG8qcy7Yl2dmxyfi7QBto16bV2WxXatTSkB1q7sl8CqsJ0DM4CX4GzgDGWZ5ZulaVRzXlTNkRfbgjatWgHOxlc1qABhXZVmbYBqKzAFzNpyZKzaY3ei33HYC5lSB1pULADAdCm86+h89OGOnWs0K4lgFPADMctYMfa1SsrwFygbcEvAB+CfZbtxzrJtPkhUL69sUULC2pAob2KTDs4pjS/POCwIrQjbMffStO7rgTtQoY+NoDUfGxqTvvsL0G4edbcPXxtdrvPzd3t/oXMNfsU2+vMgi8RN4X2JSpboo5ZwIzg7t7BoT9ERih3j8XpE/NWmFiHzbTDfn7cZt8laOsAsq9r7jeAJfx0yTKuDm3eXEMmO5CYXbIfXKouhXYtmfbZHROBn2XaG8qaz+7fc3Uu3lwDuM+lLWflKrSr7awLO7T4dGLKzjVbX7hOfH/6V+GpbB8fn+5jfOjDhwptMu19QErTIwB77x9XnZc315BhbyyJUmi76ogbCwK+dzogoEMGBIcaUGiTaQMOBg80gAb8a0ChTbD8B4sYESM0gAYU2mTaiAEgoAE04F8DCm2C5T9YxIgYoQE0oNB2k2mnB2Pi3Rp2/9AfBA64tvfAzkRdk+fLo/DxgZxvm7d/dh/OoePR8dAAGjhWAwrtYwtY/LoI2+E318yu70Roj4F57Fi0T6Adnpq0+xODwey2UQ53NaCBzWpAoe0m07YLLdn9M4q0BOHSZwLVsWOlc+acL9exJQNDA2hgTAMK7bGTvByLTyuaNUDUrgj39Oh5enpR3hhTfCgmPZYu59ilWyNg83VH4v/7R9sFwvbcta4roj4eGTR5GhHIzNEJ5yyjE4W2m0x7BA5xBb8StMNUSP750CJOhXVEBMJWVKXP5LjAWkFtp3TG7F/pMaC9TGcUfbHFn2MaUGiPneTmmM2oLajtvoBxCNqFc0uALn0mfugdG6pLbGG72flH0QxbQLyUBhTaVWTaEY7tj5Qv3z9v/v6z3Y/TJnap1ADJAZCWplh6EB56FVmCb+/8gbqWChLl0OHRABoQDSi05QPX23gnRgD10+ZDmA755Yf4lpoAUZ2qkKx2AKRAG/G71rjoly3fzgY0oNCuItNOGfCH9y2sYzb953fN2/cvmt4PjgPQzt+SHrPm8ENjlqlHuH/a//hoOzqZdhf8U3PavLmm6y+rJfbxzaEaUGgfeuG1zm9hmu7jlmmMT/v/2+P7t8XIHSI2E++c8/5pE//PoK33V+udJHuAA+1uR7s6tHlzDVnpQFZ6LU6ds16Fdi2Z9jmdQdldGFflD95cA7g3Am6FdlUddCPBISbzB5GpbB9fzvclvvLtK4U2mbbvQNGRBuLDm2vIsDeWxCm0gcIAFDYmCHSADtCAbw0otMm0fQeKjkR80AAaCBpQaCMIBIEG0AAa8K8BhTaZtv9g0aGIERpAAwptxIAYrqKBm2fN3cPXZrdLfw/vmtubb/hxjd9S0MCABhTabjLtuGJeesLR7g804CqguZQt8gKFfAXDS9Vv/W/3l6w/Qfvh7hmddEm/UtZq9aTQdgO/CIf0hKPd9ybCaNv+KcmD/ZeAbJ/U7JXhAtpnjgXQXi1cenr21ocrtUeh7SbTtm+rsfveHHwJaF+7zdb/dn9Ju4A20F5STxsoS6Fd06jYWTvkU3eFP1lHpHOOXVckwueH5mXcyholWcacFpuSdUu+2IWjIqzluu62s2hVfp7YINmzrmliyjDTILqQVTjPfG7j1DnnU3fRrEk/BHH32rlfw8XWM7V/0tOIQBtobwC0U33okOMKbTeZ9lQAv/+u+fDbt3uhRzjuYaOwFkgmOCtQFdYJ1AJROT+BzE5btGVmYB/LtEMZUl5oT25DBGb7lnZbTylwse4CtCOw7edZHZN+mDM9MxWLdBxo8yN2Sbt8dh5dKLSrdXCCrEC5B7kcThFue8iHdneuKcE4qyP6qnTeEORyG06FdskeWfEwDRadNpXqk8HKgn/I/nN+Tqa9T0DO6WfKXo2fFdrVZNoCm2x6YSlo92AXxF6C7gS086mLMNXSyapLZRY6VtGewsCTDz6960r15b603w4KtpxlYAfaq4HJWfRxKR1WVI9CuxaH96YFsqxzElYF4HWuKcE4qyP6qnReCnwsz86DjwCzA/KCcDq2yfGSPYdm2lKWbFOZU/YsrhOgDbRFg2xnaUGhXUWmLZmhyQglo10q05Yf5yy8egNFEFc2h2xhlp/fQjzLtAWyE9MTRWiXro2DyH7ap3ddaeCwnWTquD0322dO+zxzl1ZT7ONj0YBCWz5wv02wlDs7/v7tafP2z/2dE5Owmsq0DZCljqG7NwTGcp4MHAJ++Ty8yiyA3A4E0c8pu9XzBOD55zoVZH4MlQFMj+2BHcqe9EOpDjMYHqIDoA1QDtEL556mF4V2FZl2luER/NOC78J/TI/M+krsIlb0PxexUmgjihUAsMZOBbRdgID+X0//V2iTadcTtFV1sARtFoxCf6vS9RkTKIU2DqPToAE0gAb8a0ChTabtP1h0KGKEBtCAQhsxIAY0gAbQgH8NKLQ9ZNoebEC0/kVLjIjRljWg0N6yEwbbfvOquU9vVLl/w9tUBv10xh9dqBNAo4GuBhTaHrJcDzZYgcSHRu5fcUsWUEYDaMCNBhTaFlbsh0Wi2ncXkmF3R3m0gT/QwHU1oND2kOV6sEEFmaZGgPZ1BarxINNzk+kRk+v2CYU2gcgCAbSBBAMFGnCoAYW2hyzXgw06eN2+ax52H5vXN/wAqT5xKGBsy5INYrT6gUahjfiT+OWOkYd3zS3AXn0HQPdAvzYNKLQ9ZLkebNAAkmkDbLJWNOBQAwpthZVDI69iG3PadFj6AhpwqAGFtocs14MNOkAAbTqsww6r+sS2zepToY0Ysrk9oL3ZTkFfyPoCA4SrvqDQ9pDlerBBOywP17gSqsYFgBCXjWtAoU2n6GcXPMbe9wk6wSdo4LoaUGh7yHI92NARpNz+t/va8GTkdYXaicvGMy18sW0tKrQRwraFQPyJPxqoQwMKbQ9ZrgcbEG4dwiVOxGmrGlBob9UBtJvOjwbQQE0aUGh7yHI92FBT8LAV2KCB7WlAoU3wtxd8Yp5inm7v3KW3FO1Yd+Z6txW++dhoHLgBoBgHhbaHLNeDDRFkv/zQfPn0ovnwyzfNY7uf7lp4+f5F8+X906JDw/U//va8+fLnd82P51xw6uZp8+FTssPuc2fFYFwGB6kE7Ye7Z4dfm/x9e/fZwOZzc3e78OqQlxpYBJpnfGOT+GrU3zzcNqhFhfagoLcIgQjq583b7wXaad8htP/+7dvmcYJ23N9ivE5t86nQjqDbL+PbQmn//xJ9q/PMgAB8YbC2dbxr7h6+NruFy44+UD+/i+9eBdrHfbtXaHvIcj3YEMX1/XfN358SqO2+K2h/27z980XTQtvsnwqwLV6vMDki0y5dm7LEUSgd4ue44mSWvceBIvvskDLzc998bKK9ZxoQQr8Kg1l83mGOf8i0ybSXyHakDJkeidswRSHTFKkjyPTIyzBNIsfz6ZI4GKRrwzn2eJqS6WTOhWkasYftoyZmibsjIVYCr0AtTRd0AJw+iwAqAFW+/s/NVqPt+Ty6KTeWZ48LWOfO+er5JvuXz2y5oc3y+aGZdmkQSWV1fBfreEWmLfo6YkumfYTTBNYK1QTgOAcuc9oW5Pn0RXZ+gG4sswdum+2nOfYj7N0C1M8G7eBvC6QIU/OEbPw/wdAArwfasbgZQEus2qmKVzHbimUJRLX+9sXTPSAO1WNsC3UUB4pwbXae2DO5LQE6+s0MFGIbmfZgFj3p5yePGoX2nJM5p52Dkkzb+iN8JhCXTNv+EBk/Sz9elo4/LkzDtOf90HyQaRARPduTRG/jFvdLwMl83EL4Y3P/8LWdRpDjAm2FafsDZAe0cu7QNq8/mxqQsuJWXoGXrjloeQUp9z7coVGAabAvlTv3W4L1Zeuj/Ruf7MBjz3sMtE/Sr0Lbw3yyBxs64hroZEPQljtKSlC2n5Wulx8TJVtvYdLOVXemTgZsmmM35wz88JNDs+RjgVk+nZBgnUNwEFilssNnISuVsu1+mguOt8FJth3OFwC/OfAulZj9mm8KuT3STltXfs7Q/9GmNEVl9/PzgfYy0KZDD3ToXHAylWFv+bvp/hBoAS1+taAuHS9l2nKNbKUstvNjNctXM6AtEJatllu6Nn12eBbcAi/U0bk2DgzZfL1k+IfcVirlvAkvrc7KE52fAu007RKnbMLgMAR+oL0MtD1kuR5s0M4oIi5sexC1twjKnLadn87nsPP/E/QlUw82RLB/+qF5GTplPidesGmO3Ws+J8J0CERT/iqB11zTnZbo/4jWHt9DsHO+Kae1MQOyOd5OL3xuHiTjlmMCUvlc/s+hmGCYZ/0x7hkoh2ycnB6RbxZii9go2zQwPDzs/dHTXWZL73goK53TGbw6dXzdfzORzzey1emRouM24oRD294C1dz5IXAVf6U7PfTOEbl9UI6HbX6Oydyl/M5USTpf5s0PtXnt558L2i3cMtCmKQb7I6Cc1z7NNzBfnK7TaRCrh7CfgGjL1bgJqOWpzRzYqSwZGDplCMwtaLW8ZKvAWMqXrb0m1KHXZT6Rtsjx/LpwXNovZcu21JYxaEsdc++eEdtWslVoe8hyPdignWQlAaY9M6ZSEgQ6oDtH/AWeJUiF+lKWetLTlAmMxQx1oTa1A9RAJr2UL8egrfP8AzYs1E6vfUeh7dVA7JoBnZWL9OwaWAo0E3GQbLwI1CVsSGUUp0cmbJvtYxl4Spm0wHTg2Ow6gq1j0J6w4aB6lvLLBctRaHvIcj3YsPaA077CIKiw+9quH7IEdGwnFsiU5tz12ImPjsvUw9K2m3bIoFP6gVGmZU4eMKQdaeokH+DGbNiKthXaW2kw7SxAy3RM/IN/0IBvDSi0PWS5HmxAsL4FS3yIz9Y1oNDeuiNoPzBAA2igBg0otD1kuR5sqCFo2Ahc0MB2NaDQRgTbFQGxJ/ZooB4NKLQ9ZLkebIjiTQ+yuH5zzRV+PCw+fj9mhzzpGZ8OdbqOSumhkiXvwJA7U2yZcseI/WzMj4Vj8S6KE67vQFrsye4fb+/UCPdCt0+BDt818rH5v/jmnuy+aVPuvqxsrZSxB4oK7Y52S7n5olfia3loJ24zm0KZ2R0qJ9/xMmTnmT5XaHeCeKbKqqkjQtv5m2uuEKOjoR2e9hSAmyc/Xeghf6BFOn4GsKNtlfLs7X4yUJwA3UWhHbSUbNJb7PL/I+gyAGa+i7f9mTZ1/hc/ZH6N52SfTfo62HLfvmFH7Q1tSHXYh6TawWL/hOrg4HGF/jTZzgGbFNoeslwPNkRH2mVS7X5yYm/tkcy5B8Mtu/7YYJ77umPaZX1l989t6+zyM/CE62LHPhQkQzGMWeHn5v7+sy7pGsq/vwuLNu1hMtveVE+00QDy0OtL51vI2n05N/+sB1wLersvvsk/K/he6hrbhnoDrHtxKkC7HYz2fu7ZLLZVtFVojzmJY935LoFP3MqbaUwGKXC72ptr4gJTz5u3v7UvKA5Lu6otxs52ZUGzhopd5CqIWLJjaWP+hh1d2ErK6L5L85K6iZ3RZrNzO2EOjvTVu5vBpekB/dq9h4A8uddZT0S+vgfwx/2PzetQT/w/vLwg/W+hnYCmbyLPgSyZqtrQXzCp9UF6QMiWPdcX0va43naWVYcyko3RN6XMWwa8h3bd8cHplNi29iUOpXNGdRNtTLZFG2ws+i+GyCF9tE7m+vAC5ym0PWS5HmwYFUwKiMBaF2/KVu2L0L7mm2vSqoBxHe5kW1xBME77pJUDM5tDu2O7BNwCbAN5GYzk5Q5tO1N5wTd2WukC4rWxOroz5rAMULRZdgLZ2FduAXc8R+AqZQi04+cG3rHeBBwLw+S32B4Bd16mwjF74YCcL8ePAXcEYeYDE8s2wy+8DELOEVuHBtDOwGCAK9dPbYN90k4pS9YU17pl4ArbrI7OOYWBaap+B8cV2rYDsN/NrHN/SKZtPw+fCcRzuIXz4mcJgKXjpfW02/OOeHNNgnb8IdVO78T9FrKTNphzpZ2daxLUOysRps/ED3Kd620Epum80qkFunk2FzptDovwWSonTIMoVPTzFhwBeGHJ0pipGmgXpzqsXeZc8WXnmmiPacOQjeHzib928BuGtrS9B0MpN9oaYJnZI8d1QMneAGSOj9kY7LMDaLRXYpViZ49HP+XgDnVJnMMgLYPATBvG7LvEMYW2hyzXgw1znD4EbVkPuwO3JAT7Wel6WTO7BMGD31wzA9qTNtisvNAGsXe//KxMkewHrzm+vPo5Fo7SaQ0kO3CU4wU4hHa0gMhgZcrqzK+aejvg0TraKZn9VEQ3Y+zYFcuy2eV+vzPNI2UPbVM57Xz7wNKr8n5JAWVWVmzLw7vm7n4EhqbtB8W/MFh2/FCKS/ps0A+pTAv6g2zK2n+JaxXal6hsLXX0gJdlmBbQ0mZ7Tel4KdOWa2QrZU1uZ0B70oZJaLe38HUGmSsIeNIXUzaVAGKza7svZRXg0d5G9rF5nZ8fy+8CN9pssuMOeKQOa1depgwQkiGaso72h2SeqcwI31KGOgbtaGeCfbS/mxWrbbZt0t4521S+zvvr/H4aKIvQNoPfQB3FQXPgXG3DFY8rtD1kuR5smBOUHkSzudweEPP54/z/ofljeblCgvDsaYcZ0JYfIRW6uQ22DJneyX6IjO0UG68o4hCzFjJZljvHphwgAi/JJAuZmGSTt/KqrwQoyeY6x+OxcWh3fuALNg/YIOW3GX03k+3UOdLu1k99mLZlGv+ldnfm91O5sQzxj9RVOL9Xppyb+1w+n9iW6zU/Phag3dpQ8L/UlcVO+r/42GMGrtAWY9lOz/u1sNpPB3zJwRUhbo8X7qrIz8l/8Pv0olGgBoGl82eB2wI3DhCp/nyeesSGqAN7PPxAGf6XHyqT6Pu+KLRVOsgZty2MDHTm1pU6rc3eeh01P0cy3FCHHLMQE4CF8wpZcvRtPMfYm2eRtrxQjz0u5Vo7dODaT42U5p2L0E5t6LU71SmDhbChBM+23AyOMvhkdrY+M22fE6tUVm5LsCkCNtaR7kjRDLw7sOlgaI8Pvf1G4jrwbUN8cY2tQttDluvBhmsEgTqnB0p8hI8uqoGhAWfOAHPmcxTaF3XImRtFW+jgaAANnKSB9C2j9+3DAbsU2h6yXA82nBToSwRU5p7tAy/Z/qwplEvYSh2Tt9i519uFYihzyHaaqrOfT7Gcyy7JsIemTc5V7wHlKrQRD5kJGkADaMC/BhTaHrJcDzYgWv+iJUbEaMsaUGhv2Qm0HQigATRQiwYU2h6yXA821BI47AQyaGCbGlBoI4BtCoC4E3c0UJcGFNoeslwPNkQBp4dK4sMtdv+AX3jpCKYjHODD8tOm6UGjA8pZ0v/yQEq8m+FSdzEcqbXiQy6lstIDQKWHVZb0HWWZflCKwxGfKbRxrnFuhEN6ss/uH+HgQb9m65UMnrdkndcq6wDYlqF95sBYJ8QAABAKSURBVFjM9Mv+ybvslVkzr79EjIG26ceO4rJk7BXaHrJcDzZE5/Ye/T7Do9lbgvYBnacH7UvEYqZ9NUB7NhzItKu9h16hPTvYMwW++vLS+h77pUkzsKfsUo/L2iJjD8fYdT3y8+yxFIMIuOzBGrs2SL4uSOehm1R+mAKy5ch6J8XFoOyaJnN0EIG7X4NFyu5oo+fHF40scds5b6K+NsM8cD0LU2bv4Y587Y/OGhf9TDu/vvcknV07pLSGs6xXoutiHNaWbv3ZGiDSztwGxw+QHBL7rZ2r0PaQ5XqwYZYAEmg6EJSOEbbhuEA6/J/g1YHWWKYtwDZlRLAacOdQjf+PHJf1r9VmqSNAX+qJA016E00B0Hkds3wl/sgXwDKfq03y9hyxJ5wz8+8UaLfAGwCdqT+eV5rTvn3X3N8929sa4WigGxcfMv+bMmP70lN4PdDn5835P9ZdaEtuE5n2Pl5z/OroHIX23M7BeektNAaQkz4pAbr0mQgjQt68xit8bqcJEtws7NrjAtx2revOcVleVewWaFtAFupQoI/ZK3YPbQsDQPRZGCTEnnRtb3pkqMylPj8AXoPQzm3JyxxaRU+uk0enSwOCnDN3OwDtMKh1BoXcxrnlc97VYa/Q9pDlerBhEsAzs0E75SBTJB2IjkEwn1rRKZD9FEyeaXcy8TFIyjKyY/VLx7SDh92X43O3A/aUMveLQ3sqCzZtHIS2QFenNtrlUTt3ZuhSn4VjoY68jML0zBxt6ssYZK1vUzbQnv/NbZavjTYueb5C+5KV1l5XCTa2TTlQ5a3ms6E9A5BtHfv54s6a3gNA7tg9cI5th9gdpnXCtR37DxGsZ2gfkHEOQTtOzdgsearMfKoi92W6vgPZ/Jyh/0uZdmn6ZcrGofL5nEzbQqKWTFvmqIcg1sl6ZVriU//difl56osEVJ2ayDtKOt6ZI8/O6Q0ccSAwL1aYA+1QZpzCeN78LRl6Vo/aPPb5ALRLUz7xW4mdshkr1xw7ZU67vbYwD2zKD+0sQlsyZJMZt+UNv2Nx8iUAJchmtgz6vQTt3PYE7HDfeefbwNw6OO+q4CbTPlaACYIy9fHl037qQn7002Pvn8Y7NHqQTzDT8+z8roBbp0ZedOd/i1MoxgaZxjHXdyA/F9pi4xEg7X8baL8ZWD90znn/tPPW+kEwFWJ2CrRDPQJaXQ5UIGwAp8fiNIiBfDb18XD3qrl7MECMILVvlDHHQltKdUj9hbb2/CIDRzY9s7NvRM/OuX/TvuUFaNc3ZaLQ9pDlerCh1yHmdJpLnzOQuQ5m7qfYN1BXFX46pd1ce9VsEn0NDyYKbZw07CR3vsmnOgJgJDM/IiMea9/FfxgElsASDYxqQKHtIcv1YMMYwFwdK02PLAVsGQDC1IqdspHOZI+b6Red5inM37vynbSjhm029dKdoglTLiP3f9fQPmwcBXSp3yi0Swf5rKLsG/EfLH70jb5r1IBC20OW68GGGoOIzcAHDWxHAwptgr6doBNrYo0G6tWAQttDluvBBsRcr5iJHbHbggYU2lto7Nw2xoco8ntxmTNmzhgNoAEHGlBoe8hyPdhgwR4fuDjkIQcHAbX2s0/miQbWpwGFNsEtBDc8yWbXlADKZFpoAA1cWQMKbQ9ZrgcbOoMX0KaDXrmDdvSILejxyaNGoY04yLTRQEEDgBJQOtOAQttDluvBhg64DlhruXOdsyBjGzBGA+vRgEKboA4EVVZHY26bjIvBGA040IBC20OW68GGzuBFpk0nddBJO5rEns1rUqGNMAqZNj9Ebr6D0C8K/YKB46r9QqHtIcv1YEOnk45BW6ZNWGXtqgLuxAuYEIsNaEChjfgLGcUYtOUVTuHJSR7AARYbgAWMKDDiCnFXaHvIcj3Y0BHmBLTbd/19bY56AesVgt1pG/Uz0KCBKjWg0KZD90fRqcfYT30vIT7v+xyf4BM0MK4BhbaHLNeDDUEwkwtG6dtEzMtdyVqqzFoAxDgg8I8//yi0CY6/4BATYoIG0ECuAYW2hyzXgw25g/ifToMG0IAnDSi0PRmFLXQSNIAG0EBZAwptD1muBxsQSlko+AW/oAEfGlBoExAfASEOxAENoIExDSi0PWS5HmwYcxbH6ExoAA1cWwMK7WsbQv10BjSABtDAtAYU2h6yXA82IJpp0eAjfIQGrqcBhTZBuF4Q8D2+RwNoYK4GFNoeslwPNsx1HOfRydAAGriGBhTa16icOhE9GkADaOAwDSi0PWS5HmxAQIcJCH/hLzRwWQ0otHH8ZR2Pv/E3GkADx2hAoe0hy/Vggzrxn++b339tmj9+/av56eab/gp2N7fNf34Ox5vmv/8uHH/yqPnpeXv89+9v+9eHVQHDet3hJQpDLw1mNcGy31hREb9sWAMKbYXVhp3R8QHQBgz0BTTgUAMKbQ9ZrgcbOuB2GDDs4ys1Gti2BhTaCGHbQiD+xB8N1KEBhbaHLNeDDQi3DuESJ+K0VQ0otLfqANpN50cDaKAmDSi0yXIRbk3CxVb0ulUNKLS36gDaTedHA2igJg0otMm0EW5NwsVW9LpVDSi03Tjg5lVzHx44kb/7V9wryq2HaAANoIGkAYW2x0z79f3XZge0ESvAQgNoQDWg0HaTaZvgAG2+AnvUJTahy2tqQKFNpo0QrylE6kZ/aGCeBhTaHh1Gpj0viB5jh03EDg2cRwMKbTLt8zgY4eJXNIAGltSAQnvJQpcqi0wbsS+lJcpBS2vRgEKbTBtRr0XUtAMtr1kDCm2PjSTTpvN51CU2octrakChTaaNEK8pROpGf2hgngYU2h4dRqY9L4geY4dNxA4NnEcDCm0y7fM4GOHiVzSABpbUgEJ7yUKXKotMG7EvpSXKQUtr0YBC202mzYJRusbAWkRGOwAmGlhOAwptnLqcU/ElvkQDaOBcGgDaZoGqczmZcunAaAANLKUBoA20mY5BA2igIg0A7YqCtdRITTlkfWigXg2UoX3zrLl7aN8e83D37KKj8O3d5/TWmo/N65tvLlo3Qq5XyMSO2G1FA0VoR3A+vGturwhNbvejE26lE9JOtH6IBorQDsC8dIbdM/rNx2Z35YGjZxNTKXzzQQNo4Moa6EM7TY0AbUZ/Bi00gAb8aQBoX3nUpFP46xTEhJh41kAB2q+a+93n5u72yj8C3r5rHjzYAdT5OowG0IAjDeyhrXeMOLprQ2xibptO46jTeM7CsG393xL20JZOEdf+INNG/OsXPzEmxjVqoADt9h5tfohE0DUKGpvR7do1cDi0b26b//zcNH/82jT//Xd53vun5+3x37+/LX+t//df8fo/fn7f/GvoXvCxW/5k2oQ577J/5VsTW/yDBlangTqh/eRRo09O3r9aXVDWninQPrJhNHC8BvrQfvKoqeLhmnh3iYOHgMhkGDTRABq4oAaK0K7hMfb4mDvTI3SWC3YWssPjs0N8t5zvitB+rHPGl89kddpjN3DrYcqwd0PH6cSAHA2ggRVroAztFTeYEX+5ER9f4ks0cHkNAG0GKLIyNIAGKtIA0K4oWGQ1l89q8Dk+96YBoA20ybLQABqoSANlaHv+IbIi53obobGHrBEN1K+BIrRruOUP8dUvPmJIDNHA4RooQruKh2vIuPlKiwbQwAY10Id2mhphwajDR0CyBnyGBtDAuTUAtDc4Up9bVJQPuNDA+TRQgDZvrkFw5xMcvsW3aOA0DeyhrXeMDDw+fo2MVGzizTXMXV5Df9SJ7hxqYA9tMY431yBU0QJbtIAG3GmgAG3eXMPXt9O+vuE//IcGzqeBw6Ht4c01jP7uRn866fk6Kb7Ft1YDQJsBgAEADaCBijTQh3Ytb66pyMl2lGSfrAkNoIFTNFCENo+xI6pTRMW16AcNnE8DRWi7fnMNGTZfZdEAGtiwBsrQ3rBDyBDOlyHgW3yLBk7XANBmgCJrQwNooCINAO2KgkWWcnqWgg/xYe0aANpAmywLDaCBijQAtCsKVu0ZAvaT5aKB0zUAtIE2WRYaQAMVaQBoVxQsspTTsxR8iA9r1wDQBtpkWWgADVSkAaBdUbBqzxCwnywXDZyuAaANtMmy0AAaqEgDQLuiYJGlnJ6l4EN8WLsGgDbQJstCA2igIg0A7YqCVXuGgP1kuWjgdA30oR3fEfm12e0+N3e33zACA3U0gAbQgCMNAG1HwSALOT0LwYf4cO0a6EMbiJFVoAE0gAbcagBoI0634lx7xkT7+FZwjAaANtAG2mgADVSkAaBdUbCOGZW5hmwODaxLA0AbaJNloQE0UJEGgHZFwSJjWlfGRDyJ5zEaANpAmywLDaCBijQAtCsK1jGjMteQzaGBdWkAaANtsiw0gAYq0gDQrihYZEzrypiIJ/E8RgNAG2iTZaEBNFCRBoB2RcE6ZlTmGrI5NLAuDQBtoE2WhQbQQEUaANoVBYuMaV0ZE/EknsdoAGgDbbIsNIAGKtIA0K4oWMeMylxDNocG1qWBPrR5cw1ZBwMZGkADbjUAtBGnW3GSIa4rQySey8SzD20gBsTQABpAA241ALQRp1txkpktk5nhx3X5EWgDbaCNBtBARRoA2hUFi4xpXRkT8SSex2gAaANtsiw0gAYq0gDQrihYx4zKXEM2hwbWpQGgDbTJstAAGqhIA0C7omCRMa0rYyKexPMYDQBtoE2WhQbQQEUaANoVBeuYUZlryObQwLo0ALSBNlkWGkADFWkAaFcULDKmdWVMxJN4HqMBoA20ybLQABqoSANAu6JgHTMqcw3ZHBpYlwaANtAmy0IDaKAiDQDtioJFxrSujIl4Es9jNNCHNm+uIetgIEMDaMCtBoA24nQrzmOyEK4he127BvrQBmJADA2gATTgVgNAG3G6FefaMybax7eCYzQAtIE20EYDaKAiDQDtioJ1zKjMNWRzaGBdGgDaQJssCw2ggYo0ALQrChYZ07oyJuJJPI/RANAG2mRZaAANVKSBLrTffGx2u6/Nbve5ubv9phjI1/fhuPx9bF7flM87ZgThGjIPNIAG0MC4BrrQTqPN7d3nZvfwrrkdA3IEPNBGYOMCwz/4Bw0sq4EitB/fvmsedhNABtrFbyIIdFmB4k/8iQa6GgDaFc1lId6uePEH/tiiBoA20OYbAxpAAxVpoAzttNLf/ZuRHxmZHkHoFQl9ixkZbV7nN5EytFNnbO8UGZjbBtpAG2ijATRwcQ2UoU2mffFAkBWtMysirsR1aQ2Uoc3dI0CbDAoNoAGXGgDaCNOlMJfOTiiPjHctGgDaQBtoowE0UJEGgHZFwVpLpkA7yHrRwPEaKEKbx9iPdyhixHdoAA2cUwNdaMfb+Fgw6pwOp2w6NBpAA6dooAttpgqY20MDaAANuNYA0EagrgV6SkbCtWS0a9QA0AbaQBsNoIGKNAC0KwrWGrMG2kQ2jAYO0wDQBtpkWWgADVSkAaBdUbDISA7LSPAX/lqjBoA20CbLQgNooCINAO2KgrXGrIE2kQ2jgcM08P/ZrBfUvlR+mwAAAABJRU5ErkJggg=="}}},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_train_dicts(df):\n    dataset_dicts = []\n    for idx, img_id in enumerate(df[\"image_id\"].unique()):\n        record ={}\n        filename = os.path.join(TRAIN_FOLDER,img_id+\".jpg\")\n        record[\"file_name\"] = filename\n        # print(filename)\n        h,w = cv2.imread(filename).shape[:2]\n        h_ratio = 0.5\n        w_ratio = 0.5\n        record[\"image_id\"] = img_id\n        record[\"height\"] = h\n        record[\"width\"] = w\n        objects = []\n        try:\n            for _,row in df[df[\"image_id\"]==img_id].iterrows():\n                object_ = {\n                    \"bbox\" : [int(row[\"x_min\"])*w_ratio,int(row[\"y_min\"])*h_ratio,int(row[\"x_max\"])*w_ratio,int(row[\"y_max\"])*h_ratio],\n                    \"bbox_mode\" : BoxMode.XYXY_ABS,\n                    \"category_id\" : row[\"class_id\"]\n                }\n                objects.append(object_)\n        except:\n            objects = []\n            pass\n        record[\"annotations\"] = objects\n        dataset_dicts.append(record)\n    return dataset_dicts\n# get_train_dicts(dftrain_val_sampled)\nclasses = dftrain_tr[[\"class_id\",\"class_name\"]].drop_duplicates().sort_values(\"class_id\")\nprint(classes[classes[\"class_name\"]!=\"No finding\"][\"class_id\"].tolist())\nclasses = classes[classes[\"class_name\"]!=\"No finding\"][\"class_name\"].tolist()\nclasses","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Register Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"if DEBUG == True:\n    nsample = 200\n    train_sampled = []\n    for class_ in classes+[\"No finding\"]:\n        sampledimids = np.random.choice(dftrain_tr[dftrain_tr[\"class_name\"]==class_][\"image_id\"].unique(),nsample)\n        train_sampled.extend(sampledimids)\n    dftrain_tr_sampled = dftrain_tr[dftrain_tr[\"image_id\"].isin(train_sampled)]\n\n\n    nsample = 100\n    val_sampled = []\n    for class_ in classes+[\"No finding\"]:\n        sampledimids = np.random.choice(dftrain_val[dftrain_val[\"class_name\"]==class_][\"image_id\"].unique(),nsample)\n        val_sampled.extend(sampledimids)\n    dftrain_val_sampled = dftrain_val[dftrain_val[\"image_id\"].isin(val_sampled)]\n    print(dftrain_tr_sampled.shape,dftrain_val_sampled.shape)\n\n    d= \"train\"\n    DatasetCatalog.register(\"vin_\" + d, lambda d=d: get_train_dicts(dftrain_tr_sampled))\n    MetadataCatalog.get(\"vin_\" + d).set(thing_classes=classes)\n    d= \"val\"\n    DatasetCatalog.register(\"vin_\" + d, lambda d=d: get_train_dicts(dftrain_val_sampled))\n    MetadataCatalog.get(\"vin_\" + d).set(thing_classes=classes,evaluator_type=\"coco\")\n    \nelse:\n    d= \"train\"\n    DatasetCatalog.register(\"vin_\" + d, lambda d=d: get_train_dicts(dftrain_tr))\n    MetadataCatalog.get(\"vin_\" + d).set(thing_classes=classes)\n    d= \"val\"\n    DatasetCatalog.register(\"vin_\" + d, lambda d=d: get_train_dicts(dftrain_val))\n    MetadataCatalog.get(\"vin_\" + d).set(thing_classes=classes,evaluator_type=\"coco\")\n    ","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Prepare detectron2 config\nPotential for hyperparameter tuning"},{"metadata":{"_kg_hide-output":true,"trusted":true},"cell_type":"code","source":"os.makedirs(OUTPUT_DIR,exist_ok=True) \ncfg = get_cfg()\ncfg.merge_from_file(model_zoo.get_config_file(\"COCO-Detection/faster_rcnn_R_50_FPN_3x.yaml\")) # Pretrained model\ncfg.DATASETS.TRAIN = (\"vin_train\",)\ncfg.DATASETS.TEST = (\"vin_val\",)\ncfg.DATALOADER.FILTER_EMPTY_ANNOTATIONS=False #negative samples also included\ncfg.DATALOADER.NUM_WORKERS = 4\ncfg.MODEL.WEIGHTS = model_zoo.get_checkpoint_url(\"COCO-Detection/faster_rcnn_R_50_FPN_3x.yaml\")  \ncfg.SOLVER.IMS_PER_BATCH = 8\ncfg.SOLVER.BASE_LR = 0.00025  \ncfg.MODEL.ROI_HEADS.BATCH_SIZE_PER_IMAGE = 512  \ncfg.MODEL.ROI_HEADS.NUM_CLASSES = 14\ncfg.OUTPUT_DIR=OUTPUT_DIR\nprint(cfg)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Build model, optimizer, scheduler, checkpointer(for saving model), writer(logging), data loader"},{"metadata":{"trusted":true},"cell_type":"code","source":"# set up model,optimizers,scheduler\nmodel = build_model(cfg)\nmodel.train()\noptimizer = build_optimizer(cfg, model)\nscheduler = build_lr_scheduler(cfg, optimizer)\n# checkpointer helps in saving --> checkpointer.save('filename.pth')\ncheckpointer = DetectionCheckpointer(\n        model, cfg.OUTPUT_DIR, optimizer=optimizer, scheduler=scheduler\n    )\n# writer object saves loss metrics\nwriter = JSONWriter(os.path.join( \"test_metrics2.json\"))\ndata_loader = build_detection_train_loader(cfg)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Evaluation function for calculating coco metrics: APmean, AP50..."},{"metadata":{"trusted":true},"cell_type":"code","source":"def do_evaluate(cfg, model):\n    results = OrderedDict()\n    for dataset_name in cfg.DATASETS.TEST:\n        print(dataset_name)\n        data_loader = build_detection_test_loader(cfg, dataset_name)\n        evaluator = COCOEvaluator(dataset_name, cfg, True, \"inference\")\n        results_i = inference_on_dataset(model, data_loader, evaluator)\n        results[dataset_name] = results_i\n        print(results)\n        if comm.is_main_process():\n            logger.info(\"Evaluation results for {} in csv format:\".format(dataset_name))\n            print_csv_format(results_i)\n    if len(results) == 1:\n        results = list(results.values())[0]\n    return results","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Training Loop\nPoints to note:\n* Training loop needs to be included inside an EventStorage block.\n* storage to be 'steped' after every epoch\n* storage.put_scalars() + writer.write() records loss in json file\n* for loss calculation :refer https://github.com/facebookresearch/detectron2/blob/master/tools/plain_train_net.py\n* checkpointer.save(filename) saves file to cfg.OUTPUT_DIR\n"},{"metadata":{"trusted":true},"cell_type":"code","source":"loss_list = []\neval_list = []\ntotal_loss_list = []\nlogger = logging.getLogger(\"detectron2\")\nos.makedirs(OUTPUT_DIR,exist_ok=True)\nwith EventStorage(0) as storage:\n    for data, iteration in zip(data_loader, range(0, EPOCHS)):\n        storage.step() #--.needed\n#         print(iteration)\n        \n        loss_dict = model(data)\n#         print(loss_dict)\n        losses = sum(loss for loss in loss_dict.values())\n#         print(losses)\n        loss_dict_reduced = {k: v.item() for k, v in comm.reduce_dict(loss_dict).items()}\n#         print(loss_dict_reduced)\n        losses_reduced = sum(loss for loss in loss_dict_reduced.values())\n#         print(losses_reduced)\n        loss_list.append(loss_dict_reduced)\n        total_loss_list.append(losses_reduced)\n        # break\n\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n        scheduler.step()\n        storage.put_scalars(total_loss=losses_reduced, **loss_dict_reduced)\n        writer.write()\n        if iteration % EVALUATE_EVERY_N_STEPS==0:\n            results = do_evaluate(cfg, model)\n    #         print(results)\n            eval_list.append(results[\"bbox\"])\n            checkpointer.save(\"model_final\") \n        if (time.time() - kernel_start_time) >RUNTIME:\n            print(\"time uppp\")\n            break","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_loss = pd.DataFrame(loss_list)\ndf_eval = pd.DataFrame(eval_list)\ndf_totloss = pd.DataFrame({\"tot_loss\":total_loss_list})\ndf_totloss.plot()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_eval[\"AP50\"].plot()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Evaluate"},{"metadata":{},"cell_type":"markdown","source":"Evaluation config"},{"metadata":{"trusted":true},"cell_type":"code","source":"\ncfg.MODEL.WEIGHTS = os.path.join(cfg.OUTPUT_DIR, \"model_final.pth\")  # path to the model we just trained\ncfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = 0\ncfg.MODEL.ROI_HEADS.NMS_THRESH_TEST = 0.4\npredictor = DefaultPredictor(cfg)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from detectron2.utils.visualizer import Visualizer\nfrom detectron2.utils.visualizer import ColorMode\nimport random\nmodel.eval()\nval_dict = DatasetCatalog.get(\"vin_val\")\nval_metadata = MetadataCatalog.get(\"vin_val\")\nfor d in random.sample(val_dict, 10):\n    img = cv2.imread(d[\"file_name\"])\n    visualizer = Visualizer(img[:, :, ::-1], metadata=val_metadata, scale=1)\n    out = visualizer.draw_dataset_dict(d)\n    plt.imshow(out.get_image())\n    plt.show()\n    outputs = predictor(img)  # format is documented at https://detectron2.readthedocs.io/tutorials/models.html#model-output-format\n    v = Visualizer(img[:, :, ::-1],\n                   metadata=val_metadata, \n                   scale=1, \n                   )\n    out = v.draw_instance_predictions(outputs[\"instances\"].to(\"cpu\"))\n    plt.imshow(out.get_image())\n    plt.show()\n    print(\"-------------------------\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_loss.to_csv(\"loss_list.csv\")\ndf_eval.to_csv(\"eval_list.csv\")\ndf_totloss.to_csv(\"total_loss_list.csv\")","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}