{"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":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":238372687,"sourceType":"kernelVersion"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\nimport yaml\nfrom PIL import Image\nfrom pathlib import Path\nimport warnings\n!python -c \"import monai\" || pip install -q \"monai-weekly[pillow,tqdm]\"\nwarnings.filterwarnings(\"ignore\")\n\ndata_dir = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025'\nyolo_dir = '/kaggle/input/phase1-byu'\nyolo_train_images = os.path.join(yolo_dir,'images/train')\nyolo_train_labels = os.path.join(yolo_dir,'labels/train')\nyolo_val_images = os.path.join(yolo_dir,'images/val')\nyolo_val_labels = os.path.join(yolo_dir,'labels/val')\nyaml_path = os.path.join(yolo_dir,'dataset.yaml')\nwith open(yaml_path,'r') as file:\n    yaml_data = yaml.safe_load(file)\n\nif 'path' in yaml_data:\n    yaml_data['path'] = yolo_dir\n\nfixed_yaml_path = \"/kaggle/working/fixed_dataset.yaml\"\nwith open(fixed_yaml_path, 'w') as f:\n    yaml.dump(yaml_data, f)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-11T08:58:40.186492Z","iopub.execute_input":"2025-05-11T08:58:40.186786Z","iopub.status.idle":"2025-05-11T08:58:52.615869Z","shell.execute_reply.started":"2025-05-11T08:58:40.186766Z","shell.execute_reply":"2025-05-11T08:58:52.614913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T08:58:52.617848Z","iopub.execute_input":"2025-05-11T08:58:52.618149Z","iopub.status.idle":"2025-05-11T08:58:52.621801Z","shell.execute_reply.started":"2025-05-11T08:58:52.618128Z","shell.execute_reply":"2025-05-11T08:58:52.621053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_label = pd.read_csv('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv')\n\nfrom monai.transforms import (\n    LoadImaged, EnsureChannelFirstd, ScaleIntensityd, RandFlipd,Resized,\n    RandZoomd, RandAffined, Compose, ToTensord\n)\nfrom monai.data import DataLoader,Dataset\nfrom PIL import Image , ImageDraw\nimport matplotlib.pyplot as plt\nfrom monai.transforms import Resize\nfrom monai.config import KeysCollection\nfrom monai.transforms.transform import MapTransform\nimport numpy as np\n\nbox_size=64\n\ndef get_annotated_data(yolo_train_images):\n    data = []\n    \n    for img_file in os.listdir(yolo_train_images):\n        if not img_file.lower().endswith(('.png', '.jpg', '.jpeg')):  # Skip non-image files\n            continue\n\n        img_path = os.path.join(yolo_train_images, img_file)\n        img=Image.open(img_path)\n        #resize_img = img.resize((640, 640), Image.Resampling.LANCZOS)\n        width, height = img.size\n\n        # Extract x and y values directly using split and list comprehension\n        parts = img_file.split('_')\n        y_centre = int([p[1:] for p in parts if p.startswith('y')][0])\n        x_centre = int([p.split('.')[0][1:] for p in parts if p.startswith('x')][0])\n       \n\n        half_w, half_h = box_size / 2, box_size / 2\n\n        # Ensure coordinates are within bounds\n        x1, x2 = max(0, x_centre - half_w), min(width, x_centre + half_w)\n        y1, y2 = max(0, y_centre - half_h), min(height, y_centre + half_h)\n\n        tomo_id = f'tomo_{parts[1]}'\n        motors = train_label[train_label['tomo_id']==tomo_id]['Number of motors'].unique()\n        if motors[0] > 0:\n            label=1\n        else:\n            label=0\n\n        data.append({\n            \"image\": img_path,\n            \"box\": [[x1, y1, x2, y2]],\n            'labels': label\n        })\n\n    return data\n\n\n\nclass ResizeWithBBox(MapTransform):\n    def __init__(self, keys: KeysCollection, target_size, image_key=\"image\", bbox_key=\"box\"):\n        super().__init__(keys)\n        self.resize = Resize(spatial_size=target_size)\n        self.target_size = target_size\n        self.image_key = image_key\n        self.bbox_key = bbox_key\n\n    def __call__(self, data):\n        d = dict(data)\n        \n        # Original image shape\n        original_height, original_width = d[self.image_key].shape[1:]\n        new_width, new_height = self.target_size\n\n        # Resize the image\n        d[self.image_key] = self.resize(d[self.image_key])\n\n        # Resize bounding boxes\n        new_bboxes = []\n        for bbox in d[self.bbox_key]:  # bbox format: [x1, y1, x2, y2]\n            x1, y1, x2, y2 = bbox\n            x1 = int(x1 * new_width / original_width)\n            x2 = int(x2 * new_width / original_width)\n            y1 = int(y1 * new_height / original_height)\n            y2 = int(y2 * new_height / original_height)\n            new_bboxes.append([x1, y1, x2, y2])\n\n        d[self.bbox_key] = np.array(new_bboxes)\n\n        return d\n\n\nclass Ensure3Channelsd(MapTransform):\n    def __init__(self, keys):\n        super().__init__(keys)\n        \n    def __call__(self, data):\n        d = dict(data)\n        img = d[\"image\"]\n        if img.shape[0] == 1:  # grayscale image with shape [1, H, W]\n            d[\"image\"] = img.repeat(3, 1, 1)\n        return d\n\ntrain_df = get_annotated_data(yolo_train_images)\nval_df =get_annotated_data(yolo_val_images)\n\n\n# Define all transforms\nall_transforms = Compose([\n    LoadImaged(keys=[\"image\"]),\n    EnsureChannelFirstd(keys=[\"image\"]),\n    Ensure3Channelsd(keys=[\"image\"]),  # Add this line\n    ResizeWithBBox(keys=[\"image\", \"box\"], target_size=(512, 512)),\n    ScaleIntensityd(keys=[\"image\"]),\n    ToTensord(keys=[\"image\"])\n])\n\n\nimport cv2\nclass FlagellaMotorDataset(torch.utils.data.Dataset):\n    def __init__(self,data,transform=None):\n        self.data=data\n        self.transform =transform\n    def __len__(self):\n        return len(self.data)\n    def __getitem__(self,Index):\n        item= self.data[Index]\n        num_boxes = len(item[\"box\"])\n        num_boxes = len(item[\"box\"])\n        labels = torch.tensor([item['labels']] * num_boxes, dtype=torch.int64)\n        sample ={\n            \"image\":item['image'],\n            \"box\":torch.tensor(item[\"box\"], dtype=torch.float32),\n            'labels':labels\n            \n        }\n        if self.transform:\n            sample = self.transform(sample)\n        return sample\n        \ntrain_data = FlagellaMotorDataset(train_df,all_transforms)\nval_data = FlagellaMotorDataset(val_df,all_transforms)\n\nfrom torch.utils.data import DataLoader\n\ntrain_loader = DataLoader(train_data,batch_size=5,shuffle=True)\nval_loader = DataLoader(val_data,batch_size=2,shuffle=True)\n\nfrom torchvision.models.detection import fasterrcnn_resnet50_fpn\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\n\ndef get_model():\n    model = fasterrcnn_resnet50_fpn(pretrained=True)\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes=2)  # [background, motor]\n    return model\n\n\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = get_model().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\nmodel.train()\nfor epoch in range(5):\n    for batch in train_loader:\n        images = [img.to(device) for img in batch[\"image\"]]\n        targets = []\n        for i in range(len(images)):\n            targets.append({\n                \"boxes\": batch[\"box\"][i].to(device),\n                \"labels\": batch[\"labels\"][i].to(device)\n            })\n        loss_dict = model(images, targets)\n        \n        loss = sum(loss for loss in loss_dict.values())\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n    print(f\"Epoch {epoch + 1}, Loss: {loss.item():.4f}\")\n\n\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\n\ndef show_image_with_boxes(img, boxes):\n    fig, ax = plt.subplots(1)\n    ax.imshow(img.permute(1, 2, 0).squeeze(), cmap=\"gray\")\n    for box in boxes:\n        print(box)\n        x1, y1, x2, y2 = box\n        rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1,\n                                 linewidth=2, edgecolor='r', facecolor='none')\n        ax.add_patch(rect)\n    plt.show()\n\nmodel.eval()\nwith torch.no_grad():\n    for batch in train_loader:\n        images = [img.to(device) for img in batch[\"image\"]]\n        \n        outputs = model(images)\n        print(outputs)\n        for img, out in zip(images, outputs):\n            show_image_with_boxes(img.cpu(), out[\"boxes\"].cpu())\n        break\n            \n       \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T09:56:34.953382Z","iopub.execute_input":"2025-05-11T09:56:34.954346Z","iopub.status.idle":"2025-05-11T10:27:59.646642Z","shell.execute_reply.started":"2025-05-11T09:56:34.954309Z","shell.execute_reply":"2025-05-11T10:27:59.645778Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.transforms import Transform\nfrom torchvision.utils import make_grid\nimport torch\n\ndef plot_image_with_box(img, box, title=\"Image\"):\n    \"\"\"Visualize grayscale image with bounding box\"\"\"\n    img_np = img.squeeze().numpy() if isinstance(img, torch.Tensor) else img.squeeze()\n    fig, ax = plt.subplots(1, 1)\n    ax.imshow(img_np, cmap='gray')\n    for b in box:\n        x1, y1, x2, y2 = b\n        rect = plt.Rectangle((x1, y1), x2 - x1, y2 - y1,\n                             linewidth=2, edgecolor='r', facecolor='none')\n        ax.add_patch(rect)\n    ax.set_title(title)\n    plt.axis(\"off\")\n    plt.show()\n# Start from original\nitem = train_df[99].copy()\n# Step 1: Load image\nload = LoadImaged(keys=[\"image\"])\nitem = load(item)\n\nplot_image_with_box(item[\"image\"], item[\"box\"], \"Step 1: Loaded Image\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T08:59:02.600742Z","iopub.status.idle":"2025-05-11T08:59:02.600956Z","shell.execute_reply.started":"2025-05-11T08:59:02.600853Z","shell.execute_reply":"2025-05-11T08:59:02.600863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}