{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! pip install monai\n! pip install lightning\nimport cv2\nimport numpy as np\nimport torch\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.tensorboard import SummaryWriter\n\nimport monai\nfrom monai.apps.detection.metrics.coco import COCOMetric\nfrom monai.apps.detection.metrics.matching import matching_batch\nfrom monai.apps.detection.networks.retinanet_detector import RetinaNetDetector\nfrom monai.apps.detection.networks.retinanet_network import (\n    RetinaNet,\n    resnet_fpn_feature_extractor,\n)\nfrom monai.apps.detection.utils.anchor_utils import AnchorGeneratorWithAnchorShape\nfrom monai.data import DataLoader, Dataset, box_utils, load_decathlon_datalist\nfrom monai.data.utils import no_collation\nfrom monai.networks.nets import resnet\nfrom monai.transforms import ScaleIntensityRanged\nfrom monai.utils import set_determinism\nfrom monai.data.utils import decollate_batch\n\nimport pandas as pd\nimport os\nfrom tqdm.notebook import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:09:20.038850Z","iopub.execute_input":"2025-05-02T19:09:20.039137Z","iopub.status.idle":"2025-05-02T19:11:35.792995Z","shell.execute_reply.started":"2025-05-02T19:09:20.039117Z","shell.execute_reply":"2025-05-02T19:11:35.791622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025'\nBOX_SIZE = 24   # Bounding box size (in pixels)\nTRAIN_SPLIT = 0.95  \n\nGT_BOX_MODE = \"cccwhd\"\nLR= 1e-3\nSPACING = [0.703125, 0.703125, 1.25]\nPATCH_SIZE = [128,128,64]\nVAL_PATCH_SIZE = [192,192,64]\nFG_LABELS = [0]\nN_INPUT_CHANNELS = 1\nSPARTIAL_DIMS = 3\nSCORE_THRESH = 0.02\nNMS_THRESH = 0.22\nRETURNED_LAYERS = [1,2]\nCONVL_T_STRIDE = [2,2,1]\nBASE_ANCHOR_SHAPES = [[6,8,4],[8,6,5],[10,10,6]]\nBALANCED_SAMPLER_POS_FRACTION = 0.3\nVERBOSE = False\nAMP = True\nBATCH_SIZE= 2\nACCUMULATE = 4\nPREPROCESSED_DATSET_DIR = \"dataset_v1\"\nEPOCHS = 100","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:35.794954Z","iopub.execute_input":"2025-05-02T19:11:35.795974Z","iopub.status.idle":"2025-05-02T19:11:35.803590Z","shell.execute_reply.started":"2025-05-02T19:11:35.795940Z","shell.execute_reply":"2025-05-02T19:11:35.802566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the preprocessing function to extract slices, normalize, and generate YOLO annotations.\ndef prepare_dataset(dataset_dir, labels_df, data_path):\n    \"\"\"\n    Extract slices containing motors and save images with corresponding YOLO annotations.\n\n    Steps:\n    - Load the motor labels.\n    - Perform a train/validation split by tomogram.\n    - For each motor, extract slices in a range (± trust parameter).\n    - Normalize each slice and save it.\n    - Generate YOLO format bounding box annotations with a fixed box size.\n    - Create a YAML configuration file for YOLO training.\n\n    Returns:\n        dict: A summary containing dataset statistics and file paths.\n    \"\"\"\n    # Define YOLO dataset structure and parameters\n\n    os.makedirs(dataset_dir, exist_ok=True)\n    # Load the labels CSV\n\n    labels_df = labels_df[labels_df['Number of motors'] > 0]\n\n    for tomo_id in tqdm(labels_df['tomo_id'].unique(), desc=\"Processing tomograms\"):\n        tomo_path = os.path.join(DATA_DIR, \"train\", f\"{tomo_id}\")\n        slice_files = sorted([f for f in os.listdir(tomo_path) if f.endswith('.jpg')])\n        \n        if not slice_files:\n            print(f\"Skipping missing tomogram {tomo_id}\")\n            continue\n\n        \n        indices = np.linspace(0, len(slice_files)-1, 128).astype(int)\n\n        tomo_array = np.array([cv2.resize(cv2.imread(os.path.join(tomo_path, slice_files[i]), cv2.IMREAD_GRAYSCALE), (512,512))  for i in indices])\n        np.save(f\"{dataset_dir}/{tomo_id}.npy\", tomo_array)\n        \ndf = pd.read_csv(os.path.join(DATA_DIR, \"train_labels.csv\"))[:10] #Remove it\nprepare_dataset(\"dataset_v1\", df, DATA_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:35.804632Z","iopub.execute_input":"2025-05-02T19:11:35.805044Z","iopub.status.idle":"2025-05-02T19:11:45.479686Z","shell.execute_reply.started":"2025-05-02T19:11:35.805013Z","shell.execute_reply":"2025-05-02T19:11:45.478453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = df[df['Number of motors'] > 0].copy() #[:20]\n\n#Uncomment it\n#train_df, val_df = train_test_split(df, train_size = TRAIN_SPLIT) \n\n# for test\ntrain_df, val_df = df[:16],df[:16] #Remove it\n\nprint(train_df.head())\nprint(val_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:45.480886Z","iopub.execute_input":"2025-05-02T19:11:45.481201Z","iopub.status.idle":"2025-05-02T19:11:45.505734Z","shell.execute_reply.started":"2025-05-02T19:11:45.481176Z","shell.execute_reply":"2025-05-02T19:11:45.504648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.optim.lr_scheduler import _LRScheduler\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n\nclass GradualWarmupScheduler(_LRScheduler):\n    \"\"\"Gradually warm-up(increasing) learning rate in optimizer.\n    Proposed in 'Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour'.\n\n    Args:\n        optimizer (Optimizer): Wrapped optimizer.\n        multiplier: target learning rate = base lr * multiplier if multiplier > 1.0. if multiplier = 1.0, lr starts from 0 and ends up with the base_lr.\n        total_epoch: target learning rate is reached at total_epoch, gradually\n        after_scheduler: after target_epoch, use this scheduler(eg. ReduceLROnPlateau)\n    \"\"\"\n\n    def __init__(self, optimizer, multiplier, total_epoch, after_scheduler=None):\n        self.multiplier = multiplier\n        if self.multiplier < 1.0:\n            raise ValueError(\"multiplier should be greater thant or equal to 1.\")\n        self.total_epoch = total_epoch\n        self.after_scheduler = after_scheduler\n        self.finished = False\n        super(GradualWarmupScheduler, self).__init__(optimizer)\n\n    def get_lr(self):\n        if self.last_epoch > self.total_epoch:\n            if self.after_scheduler:\n                if not self.finished:\n                    self.after_scheduler.base_lrs = [base_lr * self.multiplier for base_lr in self.base_lrs]\n                    self.finished = True\n                return self.after_scheduler.get_last_lr()\n            return [base_lr * self.multiplier for base_lr in self.base_lrs]\n\n        if self.multiplier == 1.0:\n            return [base_lr * (float(self.last_epoch) / self.total_epoch) for base_lr in self.base_lrs]\n        else:\n            return [\n                base_lr * ((self.multiplier - 1.0) * self.last_epoch / self.total_epoch + 1.0)\n                for base_lr in self.base_lrs\n            ]\n\n    def step_ReduceLROnPlateau(self, metrics, epoch=None):\n        if epoch is None:\n            epoch = self.last_epoch + 1\n        self.last_epoch = (\n            epoch if epoch != 0 else 1\n        )  # ReduceLROnPlateau is called at the end of epoch, whereas others are called at beginning\n        if self.last_epoch <= self.total_epoch:\n            warmup_lr = [\n                base_lr * ((self.multiplier - 1.0) * self.last_epoch / self.total_epoch + 1.0)\n                for base_lr in self.base_lrs\n            ]\n            for param_group, lr in zip(self.optimizer.param_groups, warmup_lr):\n                param_group[\"lr\"] = lr\n        else:\n            if epoch is None:\n                self.after_scheduler.step(metrics, None)\n            else:\n                self.after_scheduler.step(metrics, epoch - self.total_epoch)\n\n    def step(self, epoch=None, metrics=None):\n        if type(self.after_scheduler) != ReduceLROnPlateau:\n            if self.finished and self.after_scheduler:\n                if epoch is None:\n                    self.after_scheduler.step(None)\n                else:\n                    self.after_scheduler.step(epoch - self.total_epoch)\n                self._last_lr = self.after_scheduler.get_last_lr()\n            else:\n                return super(GradualWarmupScheduler, self).step(epoch)\n        else:\n            self.step_ReduceLROnPlateau(metrics, epoch)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:45.508109Z","iopub.execute_input":"2025-05-02T19:11:45.508449Z","iopub.status.idle":"2025-05-02T19:11:45.522556Z","shell.execute_reply.started":"2025-05-02T19:11:45.508424Z","shell.execute_reply":"2025-05-02T19:11:45.521598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DetectionDataset3d(torch.utils.data.Dataset):\n    train_dir = os.path.join(DATA_DIR, \"train\")\n\n    def __init__(self, df):\n        self.df = df\n\n    def __len__(self):\n        return self.df.tomo_id.nunique()\n\n    def __getitem__(self, idx):\n        tomo_id = self.df.tomo_id.unique()[idx]\n\n        tomo = np.load(f\"{PREPROCESSED_DATSET_DIR}/{tomo_id}.npy\")\n\n        boxes = []\n        labels = []\n\n        tomo_motors = self.df[self.df['tomo_id'] == tomo_id]\n\n        for _, motor in tomo_motors.iterrows():\n            if pd.isna(motor['Motor axis 0']):\n                continue\n            \n            d,h,w = tomo.shape\n            \n            od,oh,ow = (int(motor['Array shape (axis 0)']),\n                    int(motor['Array shape (axis 1)']),\n                    int(motor['Array shape (axis 2)']))\n            \n            z,y,x = (int(motor['Motor axis 0']) / od * d,\n                int(motor['Motor axis 1']) / oh * h,\n                int(motor['Motor axis 2']) / ow * w)\n            \n            spacing = float(motor['Voxel spacing']) \n\n            size = BOX_SIZE\n\n            boxes.append([y, x, z, size, size, size])\n            labels.append(0)\n\n        return {\n            \"image\": torch.from_numpy(tomo.transpose(1,2,0)).unsqueeze(0),  # Shape: (C, H, W, D)\n            \"box\": torch.tensor(boxes, dtype=torch.float32),   # (N, 6)\n            \"label\": torch.tensor(labels, dtype=torch.long),  # (N,)\n        }\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:45.523870Z","iopub.execute_input":"2025-05-02T19:11:45.524656Z","iopub.status.idle":"2025-05-02T19:11:45.549267Z","shell.execute_reply.started":"2025-05-02T19:11:45.524609Z","shell.execute_reply":"2025-05-02T19:11:45.547990Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport numpy as np\n\n\ndef normalize_image_to_uint8(image):\n    \"\"\"\n    Normalize image to uint8\n    Args:\n        image: numpy array\n    \"\"\"\n    draw_img = image\n    if np.amin(draw_img) < 0:\n        draw_img -= np.amin(draw_img)\n    if np.amax(draw_img) > 1:\n        draw_img /= np.amax(draw_img)\n    draw_img = (255 * draw_img).astype(np.uint8)\n    return draw_img\n\ndef visualize_one_xy_slice_in_3d_image(gt_boxes, image, pred_boxes, gt_box_index=0):\n    \"\"\"\n    Prepare a 2D xy-plane image slice from a 3D image for visualization.\n    It draws the (gt_box_index)-th GT box and predicted boxes on the same slice.\n    The GT box will be green rect overlayed on the image.\n    The predicted boxes will be red boxes overlayed on the image.\n\n    Args:\n        gt_boxes: numpy sized (M, 6)\n        image: image numpy array, sized (H, W, D)\n        pred_boxes: numpy array sized (N, 6)\n    \"\"\"\n    draw_box = gt_boxes[gt_box_index, :]\n    draw_box_center = [round((draw_box[axis] + draw_box[axis + 3] - 1) / 2.0) for axis in range(3)]\n    draw_box = np.round(draw_box).astype(int).tolist()\n    draw_box_z = draw_box_center[2]  # the z-slice we will visualize\n\n    # draw image\n    draw_img = normalize_image_to_uint8(image[:, :, draw_box_z])\n    draw_img = cv2.cvtColor(draw_img, cv2.COLOR_GRAY2BGR)\n\n    # draw GT box, notice that cv2 uses Cartesian indexing instead of Matrix indexing.\n    # so the xy position needs to be transposed.\n    cv2.rectangle(\n        draw_img,\n        pt1=(draw_box[1], draw_box[0]),\n        pt2=(draw_box[4], draw_box[3]),\n        color=(0, 255, 0),  # green for GT\n        thickness=1,\n    )\n    # draw predicted boxes\n    for bbox in pred_boxes:\n        bbox = np.round(bbox).astype(int).tolist()\n        if bbox[5] < draw_box[2] or bbox[2] > draw_box[5]:\n            continue\n        cv2.rectangle(\n            draw_img,\n            pt1=(bbox[1], bbox[0]),\n            pt2=(bbox[4], bbox[3]),\n            color=(255, 0, 0),  # red for predicted box\n            thickness=1,\n        )\n    return draw_img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:45.550352Z","iopub.execute_input":"2025-05-02T19:11:45.550622Z","iopub.status.idle":"2025-05-02T19:11:45.575174Z","shell.execute_reply.started":"2025-05-02T19:11:45.550602Z","shell.execute_reply":"2025-05-02T19:11:45.574138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\nimport torch\nimport numpy as np\nfrom monai.transforms import (\n    Compose,\n    DeleteItemsd,\n    EnsureChannelFirstd,\n    EnsureTyped,\n    LoadImaged,\n    Orientationd,\n    RandAdjustContrastd,\n    RandGaussianNoised,\n    RandGaussianSmoothd,\n    RandRotated,\n    RandScaleIntensityd,\n    RandShiftIntensityd,\n    RandCropByPosNegLabeld,\n    RandZoomd,\n    RandFlipd,\n    RandRotate90d,\n    MapTransform,\n    CutMixd,\n    MixUpd,\n    Zoomd\n)\nfrom monai.transforms.utility.dictionary import ApplyTransformToPointsd\nfrom monai.transforms.spatial.dictionary import ConvertBoxToPointsd, ConvertPointsToBoxesd\nfrom monai.apps.detection.transforms.dictionary import (\n    AffineBoxToImageCoordinated,\n    AffineBoxToWorldCoordinated,\n    BoxToMaskd,\n    ClipBoxToImaged,\n    ConvertBoxToStandardModed,\n    MaskToBoxd,\n    ConvertBoxModed,\n    StandardizeEmptyBoxd,\n)\nfrom monai.config import KeysCollection\nfrom monai.utils.type_conversion import convert_data_type\nfrom monai.data.box_utils import clip_boxes_to_image\nfrom monai.apps.detection.transforms.box_ops import convert_box_to_mask\n\nclass GenerateExtendedBoxMask(MapTransform):\n    \"\"\"\n    Generate box mask based on the input boxes.\n    \"\"\"\n\n    def __init__(\n        self,\n        keys: KeysCollection,\n        image_key: str,\n        spatial_size: tuple[int, int, int],\n        whole_box: bool,\n        mask_image_key: str = \"mask_image\",\n    ) -> None:\n        \"\"\"\n        Args:\n            keys: keys of the corresponding items to be transformed.\n            image_key: key for the image data in the dictionary.\n            spatial_size: size of the spatial dimensions of the mask.\n            whole_box: whether to use the whole box for generating the mask.\n            mask_image_key: key to store the generated box mask.\n        \"\"\"\n        super().__init__(keys)\n        self.image_key = image_key\n        self.spatial_size = spatial_size\n        self.whole_box = whole_box\n        self.mask_image_key = mask_image_key\n\n    def generate_fg_center_boxes_np(self, boxes, image_size, whole_box=True):\n        # We don't require crop center to be within the boxes.\n        # As along as the cropped patch contains a box, it is considered as a foreground patch.\n        # Positions within extended_boxes are crop centers for foreground patches\n        spatial_dims = len(image_size)\n        boxes_np, *_ = convert_data_type(boxes, np.ndarray)\n\n        extended_boxes = np.zeros_like(boxes_np, dtype=int)\n        boxes_start = np.ceil(boxes_np[:, :spatial_dims]).astype(int)\n        boxes_stop = np.floor(boxes_np[:, spatial_dims:]).astype(int)\n        for axis in range(spatial_dims):\n            if not whole_box:\n                extended_boxes[:, axis] = boxes_start[:, axis] - self.spatial_size[axis] // 2 + 1\n                extended_boxes[:, axis + spatial_dims] = boxes_stop[:, axis] + self.spatial_size[axis] // 2 - 1\n            else:\n                # extended box start\n                extended_boxes[:, axis] = boxes_stop[:, axis] - self.spatial_size[axis] // 2 - 1\n                extended_boxes[:, axis] = np.minimum(extended_boxes[:, axis], boxes_start[:, axis])\n                # extended box stop\n                extended_boxes[:, axis + spatial_dims] = extended_boxes[:, axis] + self.spatial_size[axis] // 2\n                extended_boxes[:, axis + spatial_dims] = np.maximum(\n                    extended_boxes[:, axis + spatial_dims], boxes_stop[:, axis]\n                )\n        extended_boxes, _ = clip_boxes_to_image(extended_boxes, image_size, remove_empty=True)  # type: ignore\n        return extended_boxes\n\n    def generate_mask_img(self, boxes, image_size, whole_box=True):\n        extended_boxes_np = self.generate_fg_center_boxes_np(boxes, image_size, whole_box)\n        mask_img = convert_box_to_mask(\n            extended_boxes_np, np.ones(extended_boxes_np.shape[0]), image_size, bg_label=0, ellipse_mask=True\n        )\n        mask_img = np.amax(mask_img, axis=0, keepdims=True)[0:1, ...]\n        return mask_img\n\n    def __call__(self, data):\n        d = dict(data)\n        for key in self.key_iterator(d):\n            image = d[self.image_key]\n            boxes = d[key]\n            data[self.mask_image_key] = self.generate_mask_img(boxes, image.shape[1:], whole_box=self.whole_box)\n        return data\n    \nif AMP:\n    compute_dtype = torch.float16\nelse:\n    compute_dtype = torch.float32\n\naffine_lps_to_ras = False\n\ntrain_transforms = Compose(\n    [\n        EnsureTyped(keys=[\"box\"], dtype=torch.float32),\n        EnsureTyped(keys=[\"image\"], dtype=torch.uint8),\n        EnsureTyped(keys=[\"label\"], dtype=torch.long),\n        Orientationd(keys=[\"image\"], axcodes=\"RAS\"),\n        ScaleIntensityRanged(\n            keys=[\"image\"],\n            a_min=0,\n            a_max=255,\n            b_min=-1.0,\n            b_max=1.0,\n            clip=False,\n        ),\n        EnsureTyped(keys=[\"image\"], dtype=torch.float16),\n        ConvertBoxToStandardModed(box_keys=[\"box\"], mode=GT_BOX_MODE),\n        StandardizeEmptyBoxd(box_keys=[\"box\"], box_ref_image_keys=\"image\"),\n        ClipBoxToImaged(\n            box_keys=\"box\",\n            label_keys=[\"label\"],\n            box_ref_image_keys=\"image\",\n            remove_empty=True,\n        ),\n        BoxToMaskd(\n            box_keys=[\"box\"],\n            label_keys=[\"label\"],\n            box_mask_keys=[\"box_mask\"],\n            box_ref_image_keys=\"image\",\n            min_fg_label=0,\n            ellipse_mask=True,\n        ),\n        # GenerateExtendedBoxMask(\n        #     keys=\"box\",\n        #     image_key=\"image\",\n        #     spatial_size=[512,512,128],\n        #     whole_box=True,\n        #     mask_image_key=\"box_mask\"\n        # ),\n        # Zoomd(\n        #     keys=[\"image\", \"box_mask\"],\n        #     zoom = 0.6,\n        #     keep_size = False\n        # ),\n        RandCropByPosNegLabeld(\n            keys=[\"image\", \"box_mask\"],\n            label_key=\"box_mask\",\n            spatial_size=PATCH_SIZE,\n            num_samples=4,\n            pos=100,\n            neg=1,\n        ),\n        # CutMixd(\n        #     keys = [\"image\", \"box_mask\"],\n        #     batch_size=BATCH_SIZE,\n        #     label_keys=\"box_mask\"\n        # ),\n        # MixUpd(\n        #     keys = [\"image\", \"box_mask\"],\n        #     batch_size=BATCH_SIZE,\n        #     alpha=0.5\n        # ),\n        # RandZoomd(\n        #     keys=[\"image\", \"box_mask\"],\n        #     prob=0.2,\n        #     min_zoom=0.7,\n        #     max_zoom=1.4,\n        #     padding_mode=\"constant\",\n        #     keep_size=True,\n        # ),\n        RandFlipd(\n            keys=[\"image\", \"box_mask\"],\n            prob=0.5,\n            spatial_axis=0,\n        ),\n        RandFlipd(\n            keys=[\"image\", \"box_mask\"],\n            prob=0.5,\n            spatial_axis=1,\n        ),\n        RandFlipd(\n            keys=[\"image\", \"box_mask\"],\n            prob=0.5,\n            spatial_axis=2,\n        ),\n        RandRotate90d(\n            keys=[\"image\", \"box_mask\"],\n            prob=0.5,\n            max_k=3,\n            spatial_axes=(0, 1),\n        ),\n        # RandRotated(\n        #     keys=[\"image\", \"box_mask\"],\n        #     mode=[\"nearest\", \"nearest\"],\n        #     prob=0.2,\n        #     range_x=np.pi / 6,\n        #     range_y=np.pi / 6,\n        #     range_z=np.pi / 6,\n        #     keep_size=True,\n        #     padding_mode=\"zeros\",\n        # ),\n\n        RandGaussianNoised(keys=[\"image\"], prob=0.1, mean=0.5, std=0.1),\n        # RandGaussianSmoothd(\n        #     keys=[\"image\"],\n        #     prob=0.1,\n        #     sigma_x=(0.5, 1.0),\n        #     sigma_y=(0.5, 1.0),\n        #     sigma_z=(0.5, 1.0),\n        # ),\n        RandScaleIntensityd(keys=[\"image\"], prob=0.15, factors=0.25),\n        RandShiftIntensityd(keys=[\"image\"], prob=0.15, offsets=0.1),\n        RandAdjustContrastd(keys=[\"image\"], prob=0.3, gamma=(0.7, 1.5)),\n        MaskToBoxd(\n            box_keys=[\"box\"],\n            label_keys=[\"label\"],\n            box_mask_keys=[\"box_mask\"],\n            min_fg_label=0,\n        ),\n        ClipBoxToImaged(\n            box_keys=\"box\",\n            label_keys=[\"label\"],\n            box_ref_image_keys=\"image\",\n            remove_empty=True,\n        ),\n        DeleteItemsd(keys=[\"box_mask\"]),\n        EnsureTyped(keys=[\"image\", \"box\"], dtype=compute_dtype),\n        EnsureTyped(keys=[\"label\"], dtype=torch.long),\n    ]\n)\n\nval_transforms = Compose(\n    [\n        EnsureTyped(keys=[\"box\"], dtype=torch.float32),\n        EnsureTyped(keys=[\"image\"], dtype=torch.uint8),\n        EnsureTyped(keys=[\"label\"], dtype=torch.long),\n        ConvertBoxToStandardModed(box_keys=[\"box\"], mode=GT_BOX_MODE),\n        StandardizeEmptyBoxd(box_keys=[\"box\"], box_ref_image_keys=\"image\"),\n        Orientationd(keys=[\"image\"], axcodes=\"RAS\"),\n        ScaleIntensityRanged(\n            keys=[\"image\"],\n            a_min=0,\n            a_max=255,\n            b_min=-1.0,\n            b_max=1.0,\n            clip=False,\n        ),\n        # BoxToMaskd(\n        #     box_keys=[\"box\"],\n        #     label_keys=[\"label\"],\n        #     box_mask_keys=[\"box_mask\"],\n        #     box_ref_image_keys=\"image\",\n        #     min_fg_label=0,\n        #     ellipse_mask=True,\n        # ),\n        # Zoomd(\n        #     keys=[\"image\", \"box_mask\"],\n        #     zoom = 0.6,\n        #     keep_size = False\n        # ),\n        # MaskToBoxd(\n        #     box_keys=[\"box\"],\n        #     label_keys=[\"label\"],\n        #     box_mask_keys=[\"box_mask\"],\n        #     min_fg_label=0,\n        # ),\n        ClipBoxToImaged(\n            box_keys=\"box\",\n            label_keys=[\"label\"],\n            box_ref_image_keys=\"image\",\n            remove_empty=True,\n        ),\n        EnsureTyped(keys=[\"image\", \"box\"], dtype=compute_dtype),\n        EnsureTyped(keys=\"label\", dtype=torch.long),\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:45.576431Z","iopub.execute_input":"2025-05-02T19:11:45.576778Z","iopub.status.idle":"2025-05-02T19:11:45.619447Z","shell.execute_reply.started":"2025-05-02T19:11:45.576710Z","shell.execute_reply":"2025-05-02T19:11:45.618286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from lightning.pytorch import LightningModule\nfrom monai.networks.nets import Unet\nimport wandb\n\nclass RetinaNetLightning(LightningModule):\n    def __init__(self, config):\n        super(RetinaNetLightning, self).__init__()\n        self.save_hyperparameters(config)\n\n        # Build anchor generator\n        self.anchor_generator = AnchorGeneratorWithAnchorShape(\n            feature_map_scales=[2**l for l in range(len(self.hparams.RETURNED_LAYERS) + 1)],\n            base_anchor_shapes=self.hparams.BASE_ANCHOR_SHAPES,\n        )\n\n        # Build backbone\n        conv1_t_size = [max(7, 2 * s + 1) for s in self.hparams.CONVL_T_STRIDE]\n        backbone = resnet.ResNet(\n            block=resnet.ResNetBottleneck,\n            layers=[3, 4, 6, 3],\n            block_inplanes=resnet.get_inplanes(),\n            n_input_channels=self.hparams.N_INPUT_CHANNELS,\n            conv1_t_stride=self.hparams.CONVL_T_STRIDE,\n            conv1_t_size=conv1_t_size,\n        )\n        feature_extractor = resnet_fpn_feature_extractor(\n            backbone=backbone,\n            spatial_dims=self.hparams.SPARTIAL_DIMS,\n            pretrained_backbone=False,\n            trainable_backbone_layers=None,\n            returned_layers=self.hparams.RETURNED_LAYERS,\n        )\n        num_anchors = self.anchor_generator.num_anchors_per_location()[0]\n        size_divisible = [s * 2 * 2 ** max(self.hparams.RETURNED_LAYERS) for s in feature_extractor.body.conv1.stride]\n        net = RetinaNet(\n            spatial_dims=self.hparams.SPARTIAL_DIMS,\n            num_classes=len(self.hparams.FG_LABELS),\n            num_anchors=num_anchors,\n            feature_extractor=feature_extractor,\n            size_divisible=size_divisible,\n        )\n\n        # Build detector\n        self.detector = RetinaNetDetector(network=net, anchor_generator=self.anchor_generator, debug=self.hparams.VERBOSE)\n\n        # Set training components\n        self.detector.set_atss_matcher(num_candidates=4, center_in_gt=False)\n        self.detector.set_hard_negative_sampler(\n            batch_size_per_image=64,\n            positive_fraction=self.hparams.BALANCED_SAMPLER_POS_FRACTION,\n            pool_size=20,\n            min_neg=16,\n        )\n        self.detector.set_target_keys(box_key=\"box\", label_key=\"label\")\n\n        # Set validation components\n        self.detector.set_box_selector_parameters(\n            score_thresh=self.hparams.SCORE_THRESH,\n            topk_candidates_per_level=1000,\n            nms_thresh=self.hparams.NMS_THRESH,\n            detections_per_img=100,\n        )\n        self.detector.set_sliding_window_inferer(\n            roi_size=self.hparams.VAL_PATCH_SIZE,\n            overlap=0.25,\n            sw_batch_size=1,\n            mode=\"constant\",\n            device=\"cpu\",\n        )\n\n        self.hparams.w_cls = 1.0\n\n        self.coco_metric = COCOMetric(classes=[\"motor\"], iou_list=[0.1], max_detection=[100])\n        self.best_val_epoch_metric = 0.0\n        self.best_val_epoch = -1\n\n        self.val_outputs_all = []\n        self.val_targets_all = []\n\n    def forward(self, inputs, targets=None, use_inferer = False):\n        return self.detector(inputs, targets, use_inferer)\n\n    def training_step(self, batch, batch_idx):\n        # inputs = [\n        #     batch_data_i[\"image\"] for batch_data_i in batch\n        # ]\n        # targets = [\n        #     dict(\n        #         label=batch_data_i[\"label\"],\n        #         box=batch_data_i[\"box\"],\n        #     )\n        #     for batch_data_i in batch\n        # ]\n\n        inputs = [\n            batch_data_ii[\"image\"] for batch_data_i in batch for batch_data_ii in batch_data_i\n        ]\n        targets = [\n            dict(\n                label=batch_data_ii[\"label\"],\n                box=batch_data_ii[\"box\"],\n            )\n            for batch_data_i in batch\n            for batch_data_ii in batch_data_i\n        ]\n\n        outputs = self(inputs, targets)\n        loss = self.hparams.w_cls * outputs[self.detector.cls_key] + outputs[self.detector.box_reg_key]\n\n        self.log(\"train_loss\", loss.detach(), on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        self.log(\"train_cls_loss\", outputs[self.detector.cls_key].detach(), on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        self.log(\"train_box_reg_loss\", outputs[self.detector.box_reg_key].detach(), on_step=True, on_epoch=True, prog_bar=True, logger=True)\n\n        torch.cuda.empty_cache()\n\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        use_inferer = not all(\n            [val_data[\"image\"][0, ...].numel() < np.prod(self.hparams.VAL_PATCH_SIZE) for val_data in batch]\n        )\n        val_inputs = [val_data_i.pop(\"image\") for val_data_i in batch]\n        val_outputs = self(val_inputs, use_inferer=use_inferer)\n        #val_outputs = self(val_inputs) #, use_inferer=use_inferer\n\n        if batch_idx == 0:  # log only once per epoch to avoid spam\n            draw_img = visualize_one_xy_slice_in_3d_image(\n                gt_boxes=batch[0][\"box\"].cpu().detach().numpy(),\n                image=val_inputs[0][0, ...].cpu().detach().numpy(),\n                pred_boxes=val_outputs[0][self.detector.target_box_key].cpu().detach().numpy(),\n            )\n            self.logger.experiment.log({\n                \"val_img\": wandb.Image(draw_img, caption=\"Debug\"),\n                \"global_step\": self.global_step\n            })\n\n        self.val_outputs_all.extend(val_outputs)\n        self.val_targets_all.extend(batch)\n\n        return {}\n    \n    def on_validation_epoch_end(self):\n\n        pred_boxes=[val_data_i[self.detector.target_box_key].cpu().detach().numpy() for val_data_i in self.val_outputs_all]\n        pred_classes=[val_data_i[self.detector.target_label_key].cpu().detach().numpy() for val_data_i in self.val_outputs_all]\n        pred_scores=[val_data_i[self.detector.pred_score_key].cpu().detach().numpy() for val_data_i in self.val_outputs_all]\n        gt_boxes=[val_data_i[\"box\"].cpu().detach().numpy() for val_data_i in self.val_targets_all]\n        gt_classes=[val_data_i[\"label\"].cpu().detach().numpy() for val_data_i in self.val_targets_all]\n        \n        # print(\"pred_boxes\", pred_boxes)\n        # print(\"gt_boxes\", gt_boxes)\n\n        # print(\"pred_classes\",pred_classes)\n        # print(\"gt_classes\",gt_classes)\n        # print(\"pred_scores\", pred_scores)\n\n        # Apply your matching metric function across all predictions + targets\n        results_metric = matching_batch(\n            iou_fn=box_utils.box_iou,\n            iou_thresholds=self.coco_metric.iou_thresholds,\n            pred_boxes=pred_boxes,\n            pred_classes=pred_classes,\n            pred_scores=pred_scores,\n            gt_boxes=gt_boxes,\n            gt_classes=gt_classes,\n        )\n\n        val_epoch_metric_dict = self.coco_metric(results_metric)[0]\n        print(val_epoch_metric_dict)\n        self.log_dict(val_epoch_metric_dict, prog_bar=True, on_epoch=True)\n        \n        val_epoch_metric = val_epoch_metric_dict.values()\n        val_epoch_metric = sum(val_epoch_metric) / len(val_epoch_metric)\n        self.log(\"val_metric\", val_epoch_metric, on_epoch=True, prog_bar=True, logger=True)\n        \n        del self.val_targets_all\n        del self.val_outputs_all\n        torch.cuda.empty_cache()\n\n        self.val_targets_all = []\n        self.val_outputs_all = []\n                \n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(\n            self.detector.network.parameters(),\n            lr=self.hparams.LR,\n        )\n        # optimizer = torch.optim.SGD(\n        #     self.detector.network.parameters(),\n        #     lr=self.hparams.LR,\n        #     momentum=0.9,\n        #     weight_decay=3e-5,\n        #     nesterov=True,\n        # )\n        # scheduler = torch.optim.lr_scheduler.StepLR(\n        #     optimizer, step_size=EPOCHS//10*8, gamma=0.1\n        # )\n        after_scheduler = torch.optim.lr_scheduler.StepLR(\n            optimizer, step_size=EPOCHS//10*8, gamma=0.1\n        )\n        scheduler = {\n            'scheduler': GradualWarmupScheduler(\n                optimizer, multiplier=1, total_epoch=3, after_scheduler=after_scheduler\n            ),\n            'interval': 'epoch',\n            'frequency': 1\n        }\n        return [optimizer], [scheduler]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:45.620681Z","iopub.execute_input":"2025-05-02T19:11:45.620986Z","iopub.status.idle":"2025-05-02T19:11:50.441710Z","shell.execute_reply.started":"2025-05-02T19:11:45.620964Z","shell.execute_reply":"2025-05-02T19:11:50.440637Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.data import CacheDataset\n\ntrain_ds = CacheDataset(DetectionDataset3d(train_df), train_transforms, cache_rate=0.0)\nval_ds = CacheDataset(DetectionDataset3d(val_df), val_transforms, cache_rate=0.0)\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=8,\n    collate_fn=no_collation,\n    #persistent_workers=True,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=8,\n    collate_fn=no_collation,\n    #persistent_workers=True,\n)\n\n#print(next(iter(val_loader)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:50.443093Z","iopub.execute_input":"2025-05-02T19:11:50.443516Z","iopub.status.idle":"2025-05-02T19:11:50.456055Z","shell.execute_reply.started":"2025-05-02T19:11:50.443493Z","shell.execute_reply":"2025-05-02T19:11:50.454559Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import wandb\nfrom lightning.pytorch.loggers import WandbLogger\nfrom lightning.pytorch.callbacks import ModelCheckpoint\nfrom lightning.pytorch import Trainer\nfrom lightning.pytorch.callbacks import LearningRateMonitor\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandbkey = user_secrets.get_secret(\"wandbkey\")\n\nwandb.login(key=wandbkey)\n\nlr_monitor = LearningRateMonitor(logging_interval='step')  # or 'epoch'\n\n# Set wandb project and config\nwandb_logger = WandbLogger(project=\"BYU-3D-Detection\")\n\n# Define checkpoint callback\ncheckpoint_callback = ModelCheckpoint(\n    dirpath=\"3d_detection_checkpoints\",\n    filename=\"best_model\",\n    monitor=\"val_metric\",\n    save_top_k=1,\n    mode=\"max\"\n)\n\n# Create config dictionary\nconfig = {\n    \"GT_BOX_MODE\": GT_BOX_MODE,\n    \"LR\": LR,\n    \"SPACING\": SPACING,\n    \"BATCH_SIZE\": BATCH_SIZE,\n    \"PATCH_SIZE\": PATCH_SIZE,\n    \"VAL_PATCH_SIZE\": VAL_PATCH_SIZE,\n    \"FG_LABELS\": FG_LABELS,\n    \"N_INPUT_CHANNELS\": N_INPUT_CHANNELS,\n    \"SPARTIAL_DIMS\": SPARTIAL_DIMS,\n    \"SCORE_THRESH\": SCORE_THRESH,\n    \"NMS_THRESH\": NMS_THRESH,\n    \"RETURNED_LAYERS\": RETURNED_LAYERS,\n    \"CONVL_T_STRIDE\": CONVL_T_STRIDE,\n    \"BASE_ANCHOR_SHAPES\": BASE_ANCHOR_SHAPES,\n    \"BALANCED_SAMPLER_POS_FRACTION\": BALANCED_SAMPLER_POS_FRACTION,\n    \"VERBOSE\": VERBOSE,\n    \"AMP\": AMP,\n}\n\n# Init module\nmodel = RetinaNetLightning(config)\n\n# Init trainer\ntrainer = Trainer(\n    accelerator=(0,1) if torch.cuda.is_available() else \"cpu\",\n    max_epochs=EPOCHS,\n    logger=wandb_logger,\n    callbacks=[checkpoint_callback, lr_monitor],\n    log_every_n_steps=10,\n    check_val_every_n_epoch=2,\n    accumulate_grad_batches=ACCUMULATE,\n    precision= \"16-mixed\" if AMP else \"32-true\",\n    profiler=\"simple\",\n    num_sanity_val_steps = 1,\n)\n\n# Start training\ntrainer.fit(model, train_dataloaders=train_loader, val_dataloaders=val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-02T19:11:50.457263Z","iopub.execute_input":"2025-05-02T19:11:50.457972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}