{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"isSourceIdPinned":false,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q 'git+https://github.com/facebookresearch/detectron2.git'","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imports and configs","metadata":{}},{"cell_type":"code","source":"from detectron2.data import detection_utils as utils, build_detection_train_loader, DatasetCatalog, MetadataCatalog\nfrom detectron2.checkpoint import DetectionCheckpointer\nfrom detectron2.engine import DefaultTrainer, hooks\nfrom detectron2.utils.visualizer import Visualizer\nfrom detectron2.utils.logger import setup_logger\nfrom detectron2.evaluation import COCOEvaluator\nfrom detectron2.structures import BoxMode\nfrom detectron2.config import get_cfg\nfrom detectron2 import model_zoo\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom tqdm.notebook import tqdm\nimport detectron2.data.transforms as T\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport pandas as pd\nimport numpy as np\nimport warnings\nimport random\nimport torch\nimport json\nimport copy\nimport cv2\nimport os\n\nsetup_logger()\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    dataset_path = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/\"\n    train_image_path = os.path.join(dataset_path, \"train\")\n    train_label_path = os.path.join(dataset_path, \"train_labels.csv\")\n    sample_sub_path = os.path.join(dataset_path, \"sample_submission.csv\")\n\n    seed = 42\n    n_folds = 5\n    current_fold = 0\n    \n    model_name = \"COCO-Detection/faster_rcnn_R_50_FPN_3x.yaml\"\n    output_dir = model_name.split(\"/\")[1].split(\".\")[0] + \"_fold_\" + str(current_fold)\n    \n    box_size = 64\n    checkpoint_period = 500\n    eval_period = 500\n    warmup_iters = 1000\n    max_iter = 10000\n    learning_rate = 0.001\n    batch_size = 8\n    threshold = 0.4","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.makedirs(CFG.output_dir, exist_ok=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"random.seed(CFG.seed)\nnp.random.seed(CFG.seed)\ntorch.manual_seed(CFG.seed)\ntorch.cuda.manual_seed_all(CFG.seed)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data preprocessing","metadata":{}},{"cell_type":"code","source":"def create_dataset(tomogram_ids, labels):\n    dataset_dicts = []\n    \n    for tomo_id in tqdm(tomogram_ids):\n        tomo_motors = labels[labels['tomo_id'] == tomo_id]\n        \n        if len(tomo_motors) == 0:\n            continue\n            \n        array_shape = [\n            tomo_motors['Array shape (axis 0)'].iloc[0],\n            tomo_motors['Array shape (axis 1)'].iloc[0],\n            tomo_motors['Array shape (axis 2)'].iloc[0]\n        ]\n        \n        tomo_path = os.path.join(CFG.train_image_path, str(tomo_id))\n        \n        for _, motor in tomo_motors.iterrows():\n            z_pos = int(motor['Motor axis 0'])\n            y_pos = int(motor['Motor axis 1'])\n            x_pos = int(motor['Motor axis 2'])\n            \n            box_size = CFG.box_size\n            x1 = max(0, x_pos - box_size//2)\n            y1 = max(0, y_pos - box_size//2)\n            x2 = min(array_shape[2], x_pos + box_size//2)\n            y2 = min(array_shape[1], y_pos + box_size//2)\n            \n            slice_path = os.path.join(tomo_path, f\"slice_{z_pos:04d}.jpg\")\n            if not os.path.exists(slice_path):\n                continue\n\n            img = cv2.imread(slice_path, cv2.IMREAD_GRAYSCALE)\n            height, width = img.shape  \n            record = {\n                \"file_name\": slice_path,\n                \"image_id\": f\"{tomo_id}_{z_pos}\",\n                \"height\": height,\n                \"width\": width,\n                \"annotations\": [\n                    {\n                        \"bbox\": [x1, y1, x2, y2],\n                        \"bbox_mode\": BoxMode.XYXY_ABS,\n                        \"category_id\": 0,\n                    }\n                ]\n            }\n            dataset_dicts.append(record)\n            \n            z_range = 5  # Include 5 slices above and below\n            for z_offset in range(-z_range, z_range + 1):\n                if z_offset == 0:  # Skip the original slice\n                    continue\n                    \n                adj_z_pos = z_pos + z_offset\n                if 0 <= adj_z_pos < array_shape[0]:\n                    adj_slice_path = os.path.join(tomo_path, f\"{adj_z_pos:04d}.jpg\")\n                    if os.path.exists(adj_slice_path):\n                        adj_record = {\n                            \"file_name\": adj_slice_path,\n                            \"image_id\": f\"{tomo_id}_{adj_z_pos}\",\n                            \"height\": height,\n                            \"width\": width,\n                            \"annotations\": [\n                                {\n                                    \"bbox\": [x1, y1, x2, y2],\n                                    \"bbox_mode\": BoxMode.XYXY_ABS,\n                                    \"category_id\": 0,\n                                }\n                            ]\n                        }\n                        dataset_dicts.append(adj_record)\n    \n    return dataset_dicts","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels = pd.read_csv(CFG.train_label_path)\ntrain_labels = train_labels[train_labels[\"Number of motors\"] == 1].reset_index(drop=True)\ntrain_labels[\"fold\"] = -1\n\nsplit = StratifiedGroupKFold(n_splits=CFG.n_folds, random_state=CFG.seed, shuffle=True).split(train_labels, train_labels['Number of motors'], groups=train_labels[\"tomo_id\"])\nfor fold_idx, (train_idx, val_idx) in enumerate(split):\n    train_labels.loc[val_idx, \"fold\"] = fold_idx","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_tomo_ids = train_labels[train_labels[\"fold\"] != CFG.current_fold]['tomo_id'].unique()\nval_tomo_ids = train_labels[train_labels[\"fold\"] == CFG.current_fold]['tomo_id'].unique()\n\nprint(f\"Number of training tomograms:   {len(train_tomo_ids)}\")\nprint(f\"Number of validation tomograms: {len(val_tomo_ids)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DatasetCatalog.clear()\nMetadataCatalog.clear()\n\nDatasetCatalog.register(\"train\", lambda: create_dataset(train_tomo_ids, train_labels))\nMetadataCatalog.get(\"train\").set(thing_classes=[\"motor\"])\n\nDatasetCatalog.register(\"val\", lambda: create_dataset(val_tomo_ids, train_labels))\nMetadataCatalog.get(\"val\").set(thing_classes=[\"motor\"])\n\ntrain_dataset = DatasetCatalog.get(\"train\")\ntrain_metadata = MetadataCatalog.get(\"train\")\n\nval_dataset = DatasetCatalog.get(\"val\")\nval_metadata = MetadataCatalog.get(\"val\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(15, 5))\n\nsamples = random.sample(train_dataset, 3)\nfor i, sample in enumerate(samples):\n    img = cv2.imread(sample[\"file_name\"])\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n    visualizer = Visualizer(img, metadata=train_metadata, scale=1.0)\n    vis = visualizer.draw_dataset_dict(sample)\n\n    plt.subplot(1, 3, i + 1)\n    plt.imshow(vis.get_image())\n    plt.title(\"/\".join(sample[\"file_name\"].split(\"/\")[-2:]), fontsize=10)\n    plt.axis(\"off\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training configs","metadata":{}},{"cell_type":"code","source":"def setup_cfg():\n    cfg = get_cfg()\n    cfg.merge_from_file(model_zoo.get_config_file(CFG.model_name))\n\n    cfg.DATASETS.TRAIN = (\"train\",)\n    cfg.DATASETS.TEST = (\"val\",)\n\n    cfg.INPUT.RANDOM_FLIP = \"none\"\n    cfg.INPUT.MIN_SIZE_TRAIN_SAMPLING = \"choice\"\n    \n    cfg.DATALOADER.NUM_WORKERS = 2\n\n    cfg.MODEL.WEIGHTS = model_zoo.get_checkpoint_url(CFG.model_name)\n    cfg.MODEL.ROI_HEADS.NUM_CLASSES = 1\n\n    cfg.SOLVER.IMS_PER_BATCH = CFG.batch_size\n    cfg.SOLVER.BASE_LR = CFG.learning_rate\n    cfg.SOLVER.MAX_ITER = CFG.max_iter\n    cfg.SOLVER.STEPS = [2000, 4000, 6000, 8000]\n    cfg.SOLVER.GAMMA = 0.1\n    cfg.SOLVER.WARMUP_ITERS = CFG.warmup_iters\n    cfg.SOLVER.CHECKPOINT_PERIOD = CFG.checkpoint_period\n\n    cfg.MODEL.ROI_HEADS.BATCH_SIZE_PER_IMAGE = 128\n    cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = CFG.threshold\n\n    cfg.TEST.EVAL_PERIOD = CFG.eval_period\n\n    cfg.OUTPUT_DIR = CFG.output_dir\n\n    return cfg","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cfg = setup_cfg()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"class Trainer(DefaultTrainer):\n    @classmethod\n    def build_evaluator(cls, cfg, dataset_name, output_folder=None):\n        if output_folder is None:\n            output_folder = os.path.join(cfg.OUTPUT_DIR, \"evaluation\")\n        return COCOEvaluator(dataset_name, cfg, False, output_folder)\n    \n    def build_hooks(self):\n        hooks_list = super().build_hooks()\n        \n        for idx, hook in enumerate(hooks_list):\n            if isinstance(hook, hooks.PeriodicCheckpointer):\n                hooks_list.pop(idx)\n                break\n        \n        hooks_list.append(\n            hooks.BestCheckpointer(\n                eval_period=self.cfg.TEST.EVAL_PERIOD,\n                checkpointer=DetectionCheckpointer(self.model, self.cfg.OUTPUT_DIR),\n                val_metric=\"bbox/AP\",\n                mode=\"max\",\n                file_prefix=\"best_checkpoint\"\n            )\n        )\n        return hooks_list","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer = Trainer(cfg)","metadata":{"_kg_hide-output":true,"scrolled":true,"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.resume_or_load(resume=False)\ntrainer.train()","metadata":{"_kg_hide-output":true,"scrolled":true,"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metrics = []\nfor line in open(f\"{CFG.output_dir}/metrics.json\"):\n    metrics.append(json.loads(line))\n    \nmetrics = pd.DataFrame(metrics)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sns.set_style(\"whitegrid\")\n\nfig, (ax1, ax2, ax3) = plt.subplots(3, 1, figsize=(16, 12), sharex=True)\n\ncolors1 = sns.color_palette(\"Set2\", n_colors=3)\nfor idx, col in enumerate([\"total_loss\", \"loss_box_reg\", \"loss_cls\"]):\n    min_val = metrics[col].min()\n    min_row = metrics.loc[metrics[col].idxmin()]\n    label = f\"{col.replace('_', ' ').title()} (min: {min_val:.6f}@i={int(min_row.iteration)})\"\n    sns.lineplot(x=\"iteration\", y=col, ax=ax1, data=metrics, linewidth=2, label=label, color=colors1[idx])\n    ax1.plot(min_row[\"iteration\"], min_row[col], marker=\"X\", markersize=8, color=colors1[idx])\n\nax1.set_ylabel(\"Loss\", fontsize=14, weight='bold')\nax1.legend(loc=\"upper right\", frameon=True)\nax1.grid(True)\n\ncolors2 = sns.color_palette(\"Dark2\", n_colors=3)\nfor idx, col in enumerate([\"bbox/AP\", \"bbox/AP50\", \"bbox/AP75\"]):\n    max_val = metrics[col].max()\n    max_row = metrics.loc[metrics[col].idxmax()]\n    label = f\"{col} (max: {max_val:.2f}@i={int(max_row.iteration)})\"\n    sns.lineplot(x=\"iteration\", y=col, ax=ax2, data=metrics, linewidth=2, label=label, color=colors2[idx])\n    ax2.plot(max_row[\"iteration\"], max_row[col], marker=\"x\", markersize=6, color=colors2[idx])\n\nax2.set_ylabel(\"Average Precision\", fontsize=14, weight='bold')\nax2.legend(loc=\"lower right\", frameon=True)\nax2.grid(True)\n\nsns.lineplot(x=\"iteration\", y=\"lr\", ax=ax3, linewidth=2, data=metrics, color=\"tab:blue\")\nax3.set_ylabel(\"Learning Rate\", fontsize=14, weight='bold')\nax3.grid(True)\n\n\nplt.xlabel(\"Iteration\", fontsize=14, weight='bold')\nplt.tight_layout()\nplt.subplots_adjust(hspace=0.4)\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}