{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport random\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Circle\nimport albumentations as A\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-18T17:33:20.234537Z","iopub.execute_input":"2025-03-18T17:33:20.234845Z","iopub.status.idle":"2025-03-18T17:33:31.998656Z","shell.execute_reply.started":"2025-03-18T17:33:20.234817Z","shell.execute_reply":"2025-03-18T17:33:31.997338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# paths and variables\ntrain_root_path = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\"\ntrain_labels = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv\"\n\ntest_root_path = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test\"\n\nval_ratio = 0.2\nbatch_size = 4\nnum_workers = 4","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-18T17:33:31.999629Z","iopub.execute_input":"2025-03-18T17:33:32.000238Z","iopub.status.idle":"2025-03-18T17:33:32.007031Z","shell.execute_reply.started":"2025-03-18T17:33:32.000198Z","shell.execute_reply":"2025-03-18T17:33:32.004551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TomogramDataset(Dataset):\n    def __init__(\n        self, \n        root_path: str, \n        metadata: pd.DataFrame,\n        num_negative_samples: int = 0,\n        only_with_motor: bool = True,\n        transforms: callable = None,\n        augmentations: callable = None\n    ):\n        \"\"\"\n        Args:\n            root_path (str): Path to the directory containing tomogram folder.\n            metadata (pd.DataFrame): Label file.\n            num_negative_samples (int): Number of negative samples to generate for the dataset from each tomogram. \n                                        If pssitive, random images are selected.\n            only_with_motor (bool): Whether to filter out samples without motor data. If True, random image from tomogram is selected.\n            transforms (callable): Torch transformation functions to apply to each image.\n            augmentations (callable): Albumentation augmentation to apply random to the data.\n        \"\"\"\n        self.data = []\n        self.labels = []\n        self.transforms = transforms\n        self.augmentations = augmentations\n\n        if only_with_motor:\n            metadata = metadata[metadata['Number of motors'] != 0]\n\n        tomo_ids = list(set(metadata[\"tomo_id\"]))\n        for tomo_id in tqdm(tomo_ids):\n            tomo_data = metadata[metadata['tomo_id'] == tomo_id]\n            tomo_folder = os.listdir(os.path.join(root_path, tomo_id))\n            tomo_folder.sort()\n\n            z_pos = [int(z) for z in list(tomo_data['Motor axis 0']) if z >= 0]\n            z_avail = [i for i in range(len(tomo_folder)) if i not in z_pos]\n\n            # add positive samples (and without motor)\n            for ri, row in tomo_data.iterrows():\n                z = int(row['Motor axis 0'])\n                y = int(row['Motor axis 1'])\n                x = int(row['Motor axis 2'])\n                vs = float(row['Voxel spacing'])\n\n                if z < 0:\n                    z = np.random.choice(z_avail)\n                    z_avail.remove(z)\n                    \n                _data = os.path.join(root_path, tomo_id, tomo_folder[z])\n                _label = (x, y, vs)\n                self.data.append(_data)\n                self.labels.append(_label)\n\n            # add negative samples\n            if num_negative_samples > 0:\n                vs = list(tomo_data['Voxel spacing'])[0]\n                z = np.random.choice(z_avail, num_negative_samples)\n                for _z in z:\n                    _data = os.path.join(root_path, tomo_id, tomo_folder[_z])\n                    _label = (-1, -1, vs)\n                    self.data.append(_data)\n                    self.labels.append(_label)\n\n    def __len__(self):\n        return len(self.data)\n\n    def load_data(self, path: str) -> np.ndarray:\n        image = cv2.imread(path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n        \n        if len(image.shape) == 2:\n            image = np.expand_dims(image, axis=2)\n            \n        return image\n        \n    def __getitem__(self, idx: int):\n        \"\"\"\n        Returns:\n            tuple: A tuple (image, xy, motor_label) where:\n                - image (torch.Tensor): The transformed image tensor. Grayscale image is used.\n                - xy (torch.Tensor): The coordinates (x, y) of the motor, if motor is not present (-1,-1) is returned.\n                - motor_label (torch.Tensor): A label indicating whether the sample has a motor - one or zero.\n        \"\"\"\n        data = self.data[idx]\n        label = self.labels[idx]\n        motor_label = torch.tensor(1 if label[0] >= 0 else 0)\n\n        # load data\n        xy = label[:2]\n        image = self.load_data(data)\n        # image = cv2.resize(image, (256,256))\n\n        # apply augmentations\n        if self.augmentations is not None:\n            _xy = [xy] if xy[0] >= 0 else []\n            augmented = self.augmentations(image=image, keypoints=_xy)\n            image = augmented['image']\n            xy = augmented['keypoints'][0] if _xy else xy\n        \n        # apply transforms\n        if self.transforms is not None:\n            image = self.transforms(image)\n            h, w = image.shape[1:]\n            xy = torch.tensor(xy).float() \n            if xy[0] >= 0:\n                xy /= torch.tensor([w, h])\n        \n        return image, xy, motor_label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-18T17:33:32.008651Z","iopub.execute_input":"2025-03-18T17:33:32.009164Z","iopub.status.idle":"2025-03-18T17:33:32.030520Z","shell.execute_reply.started":"2025-03-18T17:33:32.009119Z","shell.execute_reply":"2025-03-18T17:33:32.029475Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Visualization","metadata":{}},{"cell_type":"code","source":"dataset = TomogramDataset(\n    train_root_path, \n    pd.read_csv(train_labels),\n    only_with_motor = True,\n    num_negative_samples = 0\n)\n\nfig, ax = plt.subplots(4,4, figsize=(4*5,4*5))\nax = ax.flatten()\nfor ai, _ax in enumerate(ax):\n    image, xy, motor_label = dataset[ai]\n    x, y = xy\n    if x >= 0:\n        circle = Circle((x, y), radius=40, fill=False, edgecolor=\"tab:red\", lw=3)\n        ax[ai].add_patch(circle)\n    ax[ai].imshow(image)\nplt.show()\n\nfig, ax = plt.subplots(4,4, figsize=(4*5,4*5))\nax = ax.flatten()\nfor ai, _ax in enumerate(ax):\n    image, xy, motor_label = dataset[ai]\n    x, y = xy\n    if x >= 0:\n        h, w = image[0].shape\n        circle = Circle((x, y), radius=40, fill=False, edgecolor=\"tab:red\", lw=4)\n        ax[ai].add_patch(circle)\n        ax[ai].set_ylim(y-100, y+100)\n        ax[ai].set_xlim(x-100, x+100)\n    ax[ai].imshow(image)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-18T17:33:32.032635Z","iopub.execute_input":"2025-03-18T17:33:32.033055Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train and Validation Dataset ","metadata":{}},{"cell_type":"code","source":"# split data\nall_data = pd.read_csv(train_labels)\ntomo_ids = list(set(all_data[\"tomo_id\"]))\nrandom.shuffle(tomo_ids)\n\nidx = int(np.ceil(len(tomo_ids) * val_ratio))\nval_tomo_ids = tomo_ids[:idx]\ntrain_tomo_ids = tomo_ids[idx:]\n\nval_data = all_data[all_data[\"tomo_id\"].isin(val_tomo_ids)]\ntrain_data = all_data[all_data[\"tomo_id\"].isin(train_tomo_ids)]","metadata":{"trusted":true,"execution":{"iopub.status.idle":"2025-03-18T17:33:45.649523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# prepare transforms\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485], std=[0.229])\n])\n\nresize = [\n    A.LongestMaxSize(512, p=1),\n    A.PadIfNeeded(min_height=512, min_width=512, border_mode=0, value=0, p=1)\n    # A.Resize(height=256, width=256, p=1)\n]\n\naugmentation_train = A.Compose([\n    A.RandomRotate90(p=1),\n    A.HorizontalFlip(p=0.5),\n    *resize\n], keypoint_params=A.KeypointParams(format='xy'))\n\naugmentation_val = A.Compose([\n    *resize\n], keypoint_params=A.KeypointParams(format='xy'))\n\n\n# create datasets and dataloaders\ndataset_train = TomogramDataset(\n    train_root_path, \n    train_data,\n    only_with_motor = False,\n    num_negative_samples = 2,\n    transforms = transform,\n    augmentations = augmentation_train\n)\n\ndataset_val = TomogramDataset(\n    train_root_path, \n    val_data,\n    only_with_motor = True,\n    num_negative_samples = 0,\n    transforms = transform,\n    augmentations = augmentation_val\n)\n\ndataloader_train = DataLoader(dataset_train, batch_size=batch_size, num_workers=num_workers, shuffle=True)\ndataloader_val = DataLoader(dataset_val, batch_size=batch_size, num_workers=num_workers, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-18T17:33:45.650802Z","iopub.execute_input":"2025-03-18T17:33:45.651206Z","iopub.status.idle":"2025-03-18T17:33:49.380816Z","shell.execute_reply.started":"2025-03-18T17:33:45.651180Z","shell.execute_reply":"2025-03-18T17:33:49.379907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}