{"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":"code","source":"# Machine Learning and Data Science Imports\nimport tensorflow_probability as tfp\nimport tensorflow_datasets as tfds\nimport tensorflow_addons as tfa\nimport tensorflow_hub as hub\nfrom skimage import exposure\nimport pandas as pd; pd.options.mode.chained_assignment = None\nimport numpy as np\nimport scipy\n\n# Built In Imports\nfrom datetime import datetime\nfrom glob import glob\nimport warnings\nimport IPython\nimport urllib\nimport zipfile\nimport pickle\nimport shutil\nimport string\nimport math\nimport tqdm\nimport time\nimport os\nimport gc\nimport re\n\n# Visualization Imports\nfrom matplotlib.colors import ListedColormap\nimport matplotlib.patches as patches\nimport plotly.graph_objects as go\nimport matplotlib.pyplot as plt\nimport plotly.express as px\nimport seaborn as sns\nfrom PIL import Image\nimport matplotlib\nimport plotly\nimport PIL\nimport cv2\n\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\nimport torch\nimport torchvision\n\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection import FasterRCNN\nfrom torchvision.models.detection.rpn import AnchorGenerator\n\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.utils.data.sampler import SequentialSampler\n\nfrom matplotlib import pyplot as plt\n\n\nwarnings.filterwarnings(\"ignore\")\n\n# PRESETS\nFIG_FONT = dict(family=\"Helvetica, Arial\", size=14, color=\"#7f7f7f\")\nLABEL_COLORS = [px.colors.label_rgb(px.colors.convert_to_RGB_255(x)) for x in sns.color_palette(\"Spectral\", 15)]\nLABEL_COLORS_WOUT_NO_FINDING = LABEL_COLORS[:8]+LABEL_COLORS[9:]\n\n# Other Imports\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nfrom tqdm.notebook import tqdm\nimport pydicom\n\nprint(\"\\n... IMPORTS COMPLETE ...\\n\")","metadata":{"execution":{"iopub.status.busy":"2023-09-14T03:36:38.406295Z","iopub.execute_input":"2023-09-14T03:36:38.406705Z","iopub.status.idle":"2023-09-14T03:36:43.598195Z","shell.execute_reply.started":"2023-09-14T03:36:38.406674Z","shell.execute_reply":"2023-09-14T03:36:43.597176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a style=\"text-align: font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\" id=\"background_information\">1&nbsp;&nbsp;BACKGROUND INFORMATION</a>","metadata":{}},{"cell_type":"markdown","source":"<h3 style=\"text-align: font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\">1.1  THE DATA</h3>\n\n---\n\n<b style=\"text-decoration: underline; font-family: Verdana;\">BACKGROUND INFORMATION</b>\n\nIn this competition, we are classifying common thoracic lung diseases and localizing critical findings. <br>**This is an object detection and classification problem.**\n\nFor each test image, you will be predicting a bounding box and class for all findings. If you predict that there are no findings, you should create a prediction of **`14 1 0 0 1 1`** *(14 is the class ID for no finding, and this provides a one-pixel bounding box with a confidence of 1.0)*\n\nNote that the images are in **DICOM** format, which means they contain additional data that might be useful for visualizing and classifying.\n\n![Example Radiographs](https://i.imgur.com/QWmbhXx.png)\n\n<br>\n\n<b style=\"text-decoration: underline; font-family: Verdana;\">DATASET INFORMATION</b>\n\nThe dataset comprises **`18,000`** postero-anterior (PA) CXR scans in DICOM format, which were de-identified to protect patient privacy. \n\nAll images were labeled by a panel of experienced radiologists for the presence of **14** critical radiographic findings as listed below:\n\n> **`0`** - Aortic enlargement <br>\n**`1`** - Atelectasis <br>\n**`2`** - Calcification <br>\n**`3`** - Cardiomegaly <br>\n**`4`** - Consolidation <br>\n**`5`** - ILD <br>\n**`6`** - Infiltration <br>\n**`7`** - Lung Opacity <br>\n**`8`** - Nodule/Mass <br>\n**`9`** - Other lesion <br>\n**`10`** - Pleural effusion <br>\n**`11`** - Pleural thickening <br>\n**`12`** - Pneumothorax <br>\n**`13`** - Pulmonary fibrosis <br>\n**`14`** - \"No finding\" observation was intended to capture the absence of all findings above\n\nNote that a key part of this competition is working with ground truth from multiple radiologists. That means that the same image will have multiple ground-truth labels as annotated by different radiologists.\n\n<br>\n\n<b style=\"text-decoration: underline; font-family: Verdana;\">DATA FILES</b>\n> **`train.csv`** - the train set metadata, with one row for each object, including a class and a bounding box (multiple rows per image possible)<br>\n**`sample_submission.csv`** - a sample submission file in the correct format\n\n<br>\n\n<b style=\"text-decoration: underline; font-family: Verdana;\">TRAIN COLUMNS</b>\n> **`image_id`** - unique image identifier<br>\n**`class_name`** - the name of the class of detected object (or \"No finding\")<br>\n**`class_id`** - the ID of the class of detected object<br>\n**`rad_id`** - the ID of the radiologist that made the observation<br>\n**`x_min`** - minimum X coordinate of the object's bounding box<br>\n**`y_min`** - minimum Y coordinate of the object's bounding box<br>\n**`x_max`** - maximum X coordinate of the object's bounding box<br>\n**`y_max`** - maximum Y coordinate of the object's bounding box","metadata":{}},{"cell_type":"markdown","source":"<h3 style=\"text-align: font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\">1.2  THE GOAL</h3>\n\n---\n\nIn this competition, you’ll automatically localize and classify **`14`** types of thoracic abnormalities from chest radiographs. You'll work with a dataset consisting of **`18,000`** scans that have been annotated by experienced radiologists. You can train your model with **`15,000`** independently-labeled images and will be evaluated on a test set of **`3,000`** images. These annotations were collected via VinBigData's web-based platform, VinLab. Details on building the dataset can be found in our recent paper “VinDr-CXR: An open dataset of chest X-rays with radiologist's annotations”.","metadata":{}},{"cell_type":"markdown","source":"<a style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\" id=\"setup\">2&nbsp;&nbsp;NOTEBOOK SETUP</a>","metadata":{}},{"cell_type":"code","source":"# Define the root data directory\nDATA_DIR = \"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection\"\n\n# Define the paths to the training and testing dicom folders respectively\nTRAIN_DIR = os.path.join(DATA_DIR, \"train\")\nTEST_DIR = os.path.join(DATA_DIR, \"test\")\n\n# Capture all the relevant full train/test paths\nTRAIN_DICOM_PATHS = [os.path.join(TRAIN_DIR, f_name) for f_name in os.listdir(TRAIN_DIR)]\nTEST_DICOM_PATHS = [os.path.join(TEST_DIR, f_name) for f_name in os.listdir(TEST_DIR)]\nprint(f\"\\n... The number of training files is {len(TRAIN_DICOM_PATHS)} ...\")\nprint(f\"... The number of testing files is {len(TEST_DICOM_PATHS)} ...\")\n\n# Define paths to the relevant csv files\nTRAIN_CSV = os.path.join(DATA_DIR, \"train.csv\")\nSS_CSV = os.path.join(DATA_DIR, \"sample_submission.csv\")\n\n# Create the relevant dataframe objects\ntrain_df = pd.read_csv(TRAIN_CSV)\nss_df = pd.read_csv(SS_CSV)\n\ntrain_df.fillna(0, inplace=True)\ntrain_df.loc[train_df[\"class_id\"] == 14, ['x_max', 'y_max']] = 1.0\n\nprint(\"\\n\\nTRAIN DATAFRAME\\n\\n\")\ndisplay(train_df.head(3))\n\nprint(\"\\n\\nSAMPLE SUBMISSION DATAFRAME\\n\\n\")\ndisplay(ss_df.head(3))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:01.820738Z","iopub.execute_input":"2023-09-14T04:12:01.821118Z","iopub.status.idle":"2023-09-14T04:12:02.017193Z","shell.execute_reply.started":"2023-09-14T04:12:01.821087Z","shell.execute_reply":"2023-09-14T04:12:02.016178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: navy; background-color: #ffffff;\" id=\"setup\">3&nbsp;&nbsp;HELPER FUNCTIONS</a>","metadata":{}},{"cell_type":"code","source":"def dicom2array(path, voi_lut=True, fix_monochrome=True):\n    \"\"\" Convert dicom file to numpy array \n    \n    Args:\n        path (str): Path to the dicom file to be converted\n        voi_lut (bool): Whether or not VOI LUT is available\n        fix_monochrome (bool): Whether or not to apply monochrome fix\n        \n    Returns:\n        Numpy array of the respective dicom file \n        \n    \"\"\"\n    # Use the pydicom library to read the dicom file\n    dicom = pydicom.read_file(path)\n    \n    # VOI LUT (if available by DICOM device) is used to \n    # transform raw DICOM data to \"human-friendly\" view\n    if voi_lut:\n        data = apply_voi_lut(dicom.pixel_array, dicom)\n    else:\n        data = dicom.pixel_array\n        \n    # The XRAY may look inverted\n    #   - If we want to fix this we can\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(data) - data\n    \n    # Normalize the image array and return\n    data = data - np.min(data)\n    data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data\n\ndef plot_image(img, title=\"\", figsize=(8,8), cmap=None):\n    \"\"\" Function to plot an image to save a bit of time \"\"\"\n    plt.figure(figsize=figsize)\n    \n    if cmap:\n        plt.imshow(img, cmap=cmap)\n    else:\n        img\n        plt.imshow(img)\n        \n    plt.title(title, fontweight=\"bold\")\n    plt.axis(False)\n    plt.show()\n    \ndef get_image_id(path):\n    \"\"\" Function to return the image-id from a path \"\"\"\n    return path.rsplit(\"/\", 1)[1].rsplit(\".\", 1)[0]\n\ndef create_fractional_bbox_coordinates(row):\n    \"\"\" Function to return bbox coordiantes as fractions from DF row \"\"\"\n    frac_x_min = row[\"x_min\"]/row[\"img_width\"]\n    frac_x_max = row[\"x_max\"]/row[\"img_width\"]\n    frac_y_min = row[\"y_min\"]/row[\"img_height\"]\n    frac_y_max = row[\"y_max\"]/row[\"img_height\"]\n    return frac_x_min, frac_x_max, frac_y_min, frac_y_max\n\ndef draw_bboxes(img, tl, br, rgb, label=\"\", label_location=\"tl\", opacity=0.1, line_thickness=0):\n    \"\"\" TBD \n    \n    Args:\n        TBD\n        \n    Returns:\n        TBD \n    \"\"\"\n    rect = np.uint8(np.ones((br[1]-tl[1], br[0]-tl[0], 3))*rgb)\n    sub_combo = cv2.addWeighted(img[tl[1]:br[1],tl[0]:br[0],:], 1-opacity, rect, opacity, 1.0)    \n    img[tl[1]:br[1],tl[0]:br[0],:] = sub_combo\n\n    if line_thickness>0:\n        img = cv2.rectangle(img, tuple(tl), tuple(br), rgb, line_thickness)\n        \n    if label:\n        # DEFAULTS\n        FONT = cv2.FONT_HERSHEY_SIMPLEX\n        FONT_SCALE = 1.666\n        FONT_THICKNESS = 3\n        FONT_LINE_TYPE = cv2.LINE_AA\n        \n        if type(label)==str:\n            LABEL = label.upper().replace(\" \", \"_\")\n        else:\n            LABEL = f\"CLASS_{label:02}\"\n        \n        text_width, text_height = cv2.getTextSize(LABEL, FONT, FONT_SCALE, FONT_THICKNESS)[0]\n        \n        label_origin = {\"tl\":tl, \"br\":br, \"tr\":(br[0],tl[1]), \"bl\":(tl[0],br[1])}[label_location]\n        label_offset = {\n            \"tl\":np.array([0, -10]), \"br\":np.array([-text_width, text_height+10]), \n            \"tr\":np.array([-text_width, -10]), \"bl\":np.array([0, text_height+10])\n        }[label_location]\n        img = cv2.putText(img, LABEL, tuple(label_origin+label_offset), \n                          FONT, FONT_SCALE, rgb, FONT_THICKNESS, FONT_LINE_TYPE)\n    \n    return img","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:08.359631Z","iopub.execute_input":"2023-09-14T04:12:08.359997Z","iopub.status.idle":"2023-09-14T04:12:08.38083Z","shell.execute_reply.started":"2023-09-14T04:12:08.359965Z","shell.execute_reply":"2023-09-14T04:12:08.379795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = px.histogram(train_df.image_id.value_counts(), \n                   log_y=True, color_discrete_sequence=['indianred'], opacity=0.7,\n                   labels={\"value\":\"Number of Annotations Per Image\"},\n                   title=\"<b>DISTRIBUTION OF # OF ANNOTATIONS PER PATIENT   \" \\\n                         \"<i><sub>(Log Scale for Y-Axis)</sub></i></b>\",\n                   )\nfig.update_layout(showlegend=False,\n                  xaxis_title=\"<b>Number of Unique Images</b>\",\n                  yaxis_title=\"<b>Count of All Object Annotations</b>\",\n                  font=FIG_FONT,)\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-07T18:00:51.39991Z","iopub.execute_input":"2023-09-07T18:00:51.400387Z","iopub.status.idle":"2023-09-07T18:00:53.277232Z","shell.execute_reply.started":"2023-09-07T18:00:51.400353Z","shell.execute_reply":"2023-09-07T18:00:53.276033Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data preparation","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_ids = train_df['image_id'].unique()\nvalid_ids = image_ids[-3000:]\ntrain_ids = image_ids[:-3000]","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:13.314044Z","iopub.execute_input":"2023-09-14T04:12:13.314421Z","iopub.status.idle":"2023-09-14T04:12:13.329637Z","shell.execute_reply.started":"2023-09-14T04:12:13.314381Z","shell.execute_reply":"2023-09-14T04:12:13.32862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df = train_df[train_df['image_id'].isin(valid_ids)]\ntrain_df = train_df[train_df['image_id'].isin(train_ids)]\nvalid_df.shape, train_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:14.968064Z","iopub.execute_input":"2023-09-14T04:12:14.968555Z","iopub.status.idle":"2023-09-14T04:12:15.031866Z","shell.execute_reply.started":"2023-09-14T04:12:14.968516Z","shell.execute_reply":"2023-09-14T04:12:15.02987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VinBigDataset(Dataset):\n    \n    def __init__(self, dataframe, image_dir, transforms=None):\n        super().__init__()\n        \n        self.image_ids = dataframe[\"image_id\"].unique()\n        self.df = dataframe\n        self.image_dir = image_dir\n        self.transforms = transforms\n        \n    def __getitem__(self, index):\n        \n        image_id = self.image_ids[index]\n        records = self.df[(self.df['image_id'] == image_id)]\n        records = records.reset_index(drop=True)\n\n        dicom = pydicom.dcmread(f\"{self.image_dir}/{image_id}.dicom\")\n        \n        image = dicom.pixel_array\n        \n        if \"PhotometricInterpretation\" in dicom:\n            if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n                image = np.amax(image) - image\n        \n        intercept = dicom.RescaleIntercept if \"RescaleIntercept\" in dicom else 0.0\n        slope = dicom.RescaleSlope if \"RescaleSlope\" in dicom else 1.0\n        \n        if slope != 1:\n            image = slope * image.astype(np.float64)\n            image = image.astype(np.int16)\n            \n        image += np.int16(intercept)        \n        \n        image = np.stack([image, image, image])\n        image = image.astype('float32')\n        image = image - image.min()\n        image = image / image.max()\n        image = image * 255.0\n        image = image.transpose(1,2,0)\n       \n        if records.loc[0, \"class_id\"] == 0:\n            records = records.loc[[0], :]\n        \n        boxes = records[['x_min', 'y_min', 'x_max', 'y_max']].values\n        area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])\n        area = torch.as_tensor(area, dtype=torch.float32)\n        labels = torch.tensor(records[\"class_id\"].values, dtype=torch.int64)\n\n        # suppose all instances are not crowd\n        iscrowd = torch.zeros((records.shape[0],), dtype=torch.int64)\n\n        target = {}\n        target['boxes'] = boxes\n        target['labels'] = labels\n        target['image_id'] = torch.tensor([index])\n        target['area'] = area\n        target['iscrowd'] = iscrowd\n\n        if self.transforms:\n            sample = {\n                'image': image,\n                'bboxes': target['boxes'],\n                'labels': labels\n            }\n            sample = self.transforms(**sample)\n            image = sample['image']\n            \n            target['boxes'] = torch.tensor(sample['bboxes'])\n\n        if target[\"boxes\"].shape[0] == 0:\n            # Albumentation cuts the target (class 14, 1x1px in the corner)\n            target[\"boxes\"] = torch.from_numpy(np.array([[0.0, 0.0, 1.0, 1.0]]))\n            target[\"area\"] = torch.tensor([1.0], dtype=torch.float32)\n            target[\"labels\"] = torch.tensor([0], dtype=torch.int64)\n            \n        return image, target\n    \n    def __len__(self):\n        return self.image_ids.shape[0]","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:17.066924Z","iopub.execute_input":"2023-09-14T04:12:17.067323Z","iopub.status.idle":"2023-09-14T04:12:17.086242Z","shell.execute_reply.started":"2023-09-14T04:12:17.067273Z","shell.execute_reply":"2023-09-14T04:12:17.085157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Albumentations\ndef get_train_transform():\n    return A.Compose([\n        A.Flip(0.5),\n        A.ShiftScaleRotate(scale_limit=0.1, rotate_limit=45, p=0.25),\n        A.LongestMaxSize(max_size=800, p=1.0),\n\n        # FasterRCNN will normalize.\n        A.Normalize(mean=(0, 0, 0), std=(1, 1, 1), max_pixel_value=255.0, p=1.0),\n        ToTensorV2(p=1.0)\n    ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})\n\ndef get_valid_transform():\n    return A.Compose([\n        A.Normalize(mean=(0, 0, 0), std=(1, 1, 1), max_pixel_value=255.0, p=1.0),\n        ToTensorV2(p=1.0)\n    ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:22.109978Z","iopub.execute_input":"2023-09-14T04:12:22.110337Z","iopub.status.idle":"2023-09-14T04:12:22.118457Z","shell.execute_reply.started":"2023-09-14T04:12:22.110307Z","shell.execute_reply":"2023-09-14T04:12:22.117357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load a model; pre-trained on COCO\nmodel = torchvision.models.detection.fasterrcnn_resnet50_fpn(pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T03:37:37.800604Z","iopub.execute_input":"2023-09-14T03:37:37.800963Z","iopub.status.idle":"2023-09-14T03:37:39.521571Z","shell.execute_reply.started":"2023-09-14T03:37:37.800934Z","shell.execute_reply":"2023-09-14T03:37:39.520545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 15\n\n# get number of input features for the classifier\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n\n# replace the pre-trained head with a new one\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:24.435714Z","iopub.execute_input":"2023-09-14T04:12:24.436174Z","iopub.status.idle":"2023-09-14T04:12:24.444809Z","shell.execute_reply.started":"2023-09-14T04:12:24.436134Z","shell.execute_reply":"2023-09-14T04:12:24.443727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(batch):\n    return tuple(zip(*batch))\n\ntrain_dataset = VinBigDataset(train_df, TRAIN_DIR, get_train_transform())\nvalid_dataset = VinBigDataset(valid_df, TRAIN_DIR, get_valid_transform())\n\n\n# split the dataset in train and test set\nindices = torch.randperm(len(train_dataset)).tolist()\n\ntrain_data_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=4,\n    collate_fn=collate_fn\n)\n\nvalid_data_loader = DataLoader(\n    valid_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=4,\n    collate_fn=collate_fn\n)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:26.22372Z","iopub.execute_input":"2023-09-14T04:12:26.224102Z","iopub.status.idle":"2023-09-14T04:12:26.245087Z","shell.execute_reply.started":"2023-09-14T04:12:26.224072Z","shell.execute_reply":"2023-09-14T04:12:26.244045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:28.768247Z","iopub.execute_input":"2023-09-14T04:12:28.769074Z","iopub.status.idle":"2023-09-14T04:12:28.774474Z","shell.execute_reply.started":"2023-09-14T04:12:28.769033Z","shell.execute_reply":"2023-09-14T04:12:28.773397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, targets = next(iter(train_data_loader))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:12:30.745936Z","iopub.execute_input":"2023-09-14T04:12:30.746312Z","iopub.status.idle":"2023-09-14T04:13:01.061942Z","shell.execute_reply.started":"2023-09-14T04:12:30.746281Z","shell.execute_reply":"2023-09-14T04:13:01.060654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = list(image.to(device) for image in images)\ntargets = [{k: v.to(device) for k, v in t.items()} for t in targets]","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:13:11.716015Z","iopub.execute_input":"2023-09-14T04:13:11.716455Z","iopub.status.idle":"2023-09-14T04:13:17.592814Z","shell.execute_reply.started":"2023-09-14T04:13:11.71641Z","shell.execute_reply":"2023-09-14T04:13:17.591672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"boxes = targets[2]['boxes'].cpu().numpy().astype(np.int32)\nsample = images[2].permute(1,2,0).cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:13:26.848836Z","iopub.execute_input":"2023-09-14T04:13:26.849213Z","iopub.status.idle":"2023-09-14T04:13:26.864353Z","shell.execute_reply.started":"2023-09-14T04:13:26.849183Z","shell.execute_reply":"2023-09-14T04:13:26.862826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\nfor box in boxes:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n    \nax.set_axis_off()\nax.imshow(sample)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:13:39.723398Z","iopub.execute_input":"2023-09-14T04:13:39.723778Z","iopub.status.idle":"2023-09-14T04:13:40.212303Z","shell.execute_reply.started":"2023-09-14T04:13:39.723748Z","shell.execute_reply":"2023-09-14T04:13:40.211311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"class Averager:\n    def __init__(self):\n        self.current_total = 0.0\n        self.iterations = 0.0\n\n    def send(self, value):\n        self.current_total += value\n        self.iterations += 1\n\n    @property\n    def value(self):\n        if self.iterations == 0:\n            return 0\n        else:\n            return 1.0 * self.current_total / self.iterations\n\n    def reset(self):\n        self.current_total = 0.0\n        self.iterations = 0.0","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:14:07.10279Z","iopub.execute_input":"2023-09-14T04:14:07.103159Z","iopub.status.idle":"2023-09-14T04:14:07.110861Z","shell.execute_reply.started":"2023-09-14T04:14:07.10313Z","shell.execute_reply":"2023-09-14T04:14:07.109748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(device)\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.SGD(params, lr=0.005, momentum=0.9, weight_decay=0.0005)\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=4, gamma=0.1)\n\nnum_epochs = 12","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:14:14.762531Z","iopub.execute_input":"2023-09-14T04:14:14.762952Z","iopub.status.idle":"2023-09-14T04:14:14.836016Z","shell.execute_reply.started":"2023-09-14T04:14:14.762922Z","shell.execute_reply":"2023-09-14T04:14:14.835062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_hist = Averager()\nitr = 1\n\nfor epoch in range(num_epochs):\n    loss_hist.reset()\n    \n    for images, targets in train_data_loader:\n        \n        images = list(image.to(device) for image in images)\n        targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n\n        loss_dict = model(images, targets)\n\n        losses = sum(loss for loss in loss_dict.values())\n        loss_value = losses.item()\n\n        loss_hist.send(loss_value)\n\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n\n        if itr % 100 == 0:\n            print(f\"Iteration #{itr} loss: {loss_hist.value}\")\n\n        itr += 1\n        \n        # !!!REMOVE THIS!!!\n        break\n    \n    # update the learning rate\n    if lr_scheduler is not None:\n        lr_scheduler.step()\n\n    print(f\"Epoch #{epoch} loss: {loss_hist.value}\")   \n    # print(\"Saving epoch's state...\")\n    # torch.save(model.state_dict(), f\"model_state_epoch_{epoch}.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:14:25.172107Z","iopub.execute_input":"2023-09-14T04:14:25.17268Z","iopub.status.idle":"2023-09-14T04:19:53.598074Z","shell.execute_reply.started":"2023-09-14T04:14:25.172633Z","shell.execute_reply":"2023-09-14T04:19:53.596545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, targets = next(iter(valid_data_loader))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:19:53.601671Z","iopub.execute_input":"2023-09-14T04:19:53.604827Z","iopub.status.idle":"2023-09-14T04:20:24.781878Z","shell.execute_reply.started":"2023-09-14T04:19:53.604768Z","shell.execute_reply":"2023-09-14T04:20:24.77986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = list(img.to(device) for img in images)\ntargets = [{k: v.to(device) for k, v in t.items()} for t in targets]","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:20:24.784674Z","iopub.execute_input":"2023-09-14T04:20:24.785139Z","iopub.status.idle":"2023-09-14T04:20:25.051071Z","shell.execute_reply.started":"2023-09-14T04:20:24.785092Z","shell.execute_reply":"2023-09-14T04:20:25.04788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"boxes = targets[1]['boxes'].cpu().numpy().astype(np.int32)\nsample = images[1].permute(1,2,0).cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:20:25.052761Z","iopub.execute_input":"2023-09-14T04:20:25.05316Z","iopub.status.idle":"2023-09-14T04:20:25.12394Z","shell.execute_reply.started":"2023-09-14T04:20:25.053123Z","shell.execute_reply":"2023-09-14T04:20:25.122432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\ncpu_device = torch.device(\"cpu\")\n\noutputs = model(images)\noutputs = [{k: v.to(cpu_device) for k, v in t.items()} for t in outputs]","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:20:49.348861Z","iopub.execute_input":"2023-09-14T04:20:49.349249Z","iopub.status.idle":"2023-09-14T04:20:50.195181Z","shell.execute_reply.started":"2023-09-14T04:20:49.349219Z","shell.execute_reply":"2023-09-14T04:20:50.194081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\nfor box in boxes:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n    \nax.set_axis_off()\nax.imshow(sample)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T04:20:51.568932Z","iopub.execute_input":"2023-09-14T04:20:51.569823Z","iopub.status.idle":"2023-09-14T04:20:53.191223Z","shell.execute_reply.started":"2023-09-14T04:20:51.569779Z","shell.execute_reply":"2023-09-14T04:20:53.190273Z"},"trusted":true},"execution_count":null,"outputs":[]}]}