{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        os.path.join(dirname, filename)\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:24:23.984739Z","iopub.execute_input":"2025-07-18T13:24:23.985018Z","iopub.status.idle":"2025-07-18T13:34:59.694965Z","shell.execute_reply.started":"2025-07-18T13:24:23.984997Z","shell.execute_reply":"2025-07-18T13:34:59.694376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install monai nibabel torch torchvision\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:39:44.071037Z","iopub.execute_input":"2025-07-18T13:39:44.071628Z","iopub.status.idle":"2025-07-18T13:41:15.322443Z","shell.execute_reply.started":"2025-07-18T13:39:44.071602Z","shell.execute_reply":"2025-07-18T13:41:15.321627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# utils/dataset.py\nimport os\nimport numpy as np\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nimport pandas as pd\nimport torch\n\nclass TomoDataset(Dataset):\n    def __init__(self, data_root, csv_path, transform=None):\n        self.data_root = data_root\n        self.labels_df = pd.read_csv(csv_path)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.labels_df)\n\n    def __getitem__(self, idx):\n        row = self.labels_df.iloc[idx]\n        tomo_id = row[\"tomo_id\"]\n        tomo_folder = os.path.join(self.data_root, tomo_id)\n        slices = sorted(os.listdir(tomo_folder))\n        volume = [np.array(Image.open(os.path.join(tomo_folder, s))) for s in slices]\n        volume = np.stack(volume, axis=0).astype(np.float32)  # [D, H, W]\n\n        if self.transform:\n            volume = self.transform(volume)\n\n        label = np.array([row[\"Motor axis 0\"], row[\"Motor axis 1\"], row[\"Motor axis 2\"]], dtype=np.float32)\n        has_motor = not np.isnan(label).any()\n        label = label if has_motor else np.array([-1, -1, -1], dtype=np.float32)\n\n        return torch.from_numpy(volume).unsqueeze(0), torch.tensor(label), torch.tensor(has_motor, dtype=torch.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:44:55.110339Z","iopub.execute_input":"2025-07-18T13:44:55.111134Z","iopub.status.idle":"2025-07-18T13:45:00.567083Z","shell.execute_reply.started":"2025-07-18T13:44:55.111103Z","shell.execute_reply":"2025-07-18T13:45:00.566368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.transforms import Compose, ScaleIntensity, Resize, ToTensor\n\ntransform = Compose([\n    ScaleIntensity(),\n    Resize((64, 128, 128)),  # D, H, W\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:45:15.093319Z","iopub.execute_input":"2025-07-18T13:45:15.093753Z","iopub.status.idle":"2025-07-18T13:45:54.050072Z","shell.execute_reply.started":"2025-07-18T13:45:15.093730Z","shell.execute_reply":"2025-07-18T13:45:54.049324Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# utils/dataset.py\nimport os\nimport numpy as np\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nimport pandas as pd\nimport torch\n\nclass TomoDataset(Dataset):\n    def __init__(self, data_root, csv_path, transform=None):\n        self.data_root = data_root\n        self.labels_df = pd.read_csv(csv_path)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.labels_df)\n\n    def __getitem__(self, idx):\n        row = self.labels_df.iloc[idx]\n        tomo_id = row[\"tomo_id\"]\n        tomo_folder = os.path.join(self.data_root, tomo_id)\n        slices = sorted(os.listdir(tomo_folder))\n        volume = [np.array(Image.open(os.path.join(tomo_folder, s))) for s in slices]\n        volume = np.stack(volume, axis=0).astype(np.float32)  # [D, H, W]\n\n        if self.transform:\n            volume = self.transform(volume)\n\n        label = np.array([row[\"Motor axis 0\"], row[\"Motor axis 1\"], row[\"Motor axis 2\"]], dtype=np.float32)\n        has_motor = not np.isnan(label).any()\n        label = label if has_motor else np.array([-1, -1, -1], dtype=np.float32)\n\n        return torch.from_numpy(volume).unsqueeze(0), torch.tensor(label), torch.tensor(has_motor, dtype=torch.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:45:59.625357Z","iopub.execute_input":"2025-07-18T13:45:59.626726Z","iopub.status.idle":"2025-07-18T13:45:59.635618Z","shell.execute_reply.started":"2025-07-18T13:45:59.626688Z","shell.execute_reply":"2025-07-18T13:45:59.634511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nimport numpy as np\n\nclass TomoDataset(Dataset):\n    def __init__(self, root_dir, label_csv, target_shape=(64, 128, 128)):\n        self.root_dir = root_dir\n        self.labels = pd.read_csv(label_csv)\n\n        # Only include rows with motor annotations and folders that exist\n        self.labels = self.labels.dropna(subset=['Motor axis 0', 'Motor axis 1', 'Motor axis 2'])\n        self.labels = self.labels[self.labels['tomo_id'].apply(\n            lambda tid: os.path.isdir(os.path.join(root_dir, tid))\n        )].reset_index(drop=True)\n\n        self.target_depth, self.target_height, self.target_width = target_shape\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n        row = self.labels.iloc[idx]\n        tomo_id = row['tomo_id']\n        tomo_path = os.path.join(self.root_dir, tomo_id)\n\n        # Load and resize all slices\n        slice_files = sorted([\n            f for f in os.listdir(tomo_path)\n            if f.lower().endswith(('.jpg', '.png'))\n        ])\n\n        resized_slices = []\n        for fname in slice_files:\n            img_path = os.path.join(tomo_path, fname)\n            img = Image.open(img_path).convert('L')\n            img = img.resize((self.target_width, self.target_height))\n            img_np = np.array(img, dtype=np.float32) / 255.0\n            resized_slices.append(img_np)\n\n        # Stack to [D, H, W]\n        volume = np.stack(resized_slices, axis=0)\n\n        # === Handle depth: pad or crop to self.target_depth ===\n        depth = volume.shape[0]\n        if depth < self.target_depth:\n            pad_before = (self.target_depth - depth) // 2\n            pad_after = self.target_depth - depth - pad_before\n            volume = np.pad(volume, ((pad_before, pad_after), (0, 0), (0, 0)), mode='constant', constant_values=0)\n        elif depth > self.target_depth:\n            start = (depth - self.target_depth) // 2\n            volume = volume[start:start + self.target_depth]\n\n        # Final shape [1, D, H, W]\n        volume = np.expand_dims(volume, axis=0)\n        volume_tensor = torch.tensor(volume, dtype=torch.float32)\n\n        # Label: 3D motor position\n        target = row[['Motor axis 0', 'Motor axis 1', 'Motor axis 2']].values.astype(np.float32)\n        target_tensor = torch.tensor(target)\n\n        return volume_tensor, target_tensor\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:46:03.449039Z","iopub.execute_input":"2025-07-18T13:46:03.449397Z","iopub.status.idle":"2025-07-18T13:46:03.462512Z","shell.execute_reply.started":"2025-07-18T13:46:03.449371Z","shell.execute_reply":"2025-07-18T13:46:03.461691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.networks.nets import resnet\n\ndef get_model():\n    model = resnet.resnet18(spatial_dims=3, n_input_channels=1, num_classes=3)  # For x, y, z regression\n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:46:10.262540Z","iopub.execute_input":"2025-07-18T13:46:10.263137Z","iopub.status.idle":"2025-07-18T13:46:10.267405Z","shell.execute_reply.started":"2025-07-18T13:46:10.263111Z","shell.execute_reply":"2025-07-18T13:46:10.266482Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\n# Dataset & loader\ntrain_dataset = TomoDataset(\n    root_dir=\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\",\n    label_csv=\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv\",\n    target_shape=(64, 64, 64)\n)\ntrain_loader = DataLoader(train_dataset, batch_size=2, shuffle=True)\n\n# Model, loss, optimizer\nmodel = get_model().cuda()\ncriterion = nn.MSELoss()\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\n\n# Training loop\nepochs = 5\nmodel.train()\nfor epoch in range(epochs):\n    epoch_loss = 0\n    for i, (x, y) in enumerate(train_loader):\n        x, y = x.cuda(), y.cuda()\n        optimizer.zero_grad()\n        out = model(x)\n        loss = criterion(out, y)\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss.item()\n    print(f\"Epoch {epoch+1}/{epochs}, Loss: {epoch_loss/len(train_loader):.4f}\")\ntorch.save(model.state_dict(), \"best_model.pth\")    \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-18T13:46:13.397795Z","iopub.execute_input":"2025-07-18T13:46:13.398664Z","execution_failed":"2025-07-18T17:13:46.521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestTomoDataset(torch.utils.data.Dataset):\n    def __init__(self, root_dir, target_shape=(64, 128, 128)):\n        self.root_dir = root_dir\n        self.tomo_ids = [d for d in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, d))]\n        self.tomo_ids.sort()  # Ensure consistent ordering\n        self.target_depth, self.target_height, self.target_width = target_shape\n\n    def __len__(self):\n        return len(self.tomo_ids)\n\n    def __getitem__(self, idx):\n        tomo_id = self.tomo_ids[idx]\n        tomo_path = os.path.join(self.root_dir, tomo_id)\n\n        # Load and resize all slices\n        slice_files = sorted([\n            f for f in os.listdir(tomo_path)\n            if f.lower().endswith(('.jpg', '.png'))\n        ])\n\n        resized_slices = []\n        for fname in slice_files:\n            img_path = os.path.join(tomo_path, fname)\n            img = Image.open(img_path).convert('L')\n            img = img.resize((self.target_width, self.target_height))\n            img_np = np.array(img, dtype=np.float32) / 255.0\n            resized_slices.append(img_np)\n\n        volume = np.stack(resized_slices, axis=0)\n\n        # Pad/crop to target depth\n        depth = volume.shape[0]\n        if depth < self.target_depth:\n            pad_before = (self.target_depth - depth) // 2\n            pad_after = self.target_depth - depth - pad_before\n            volume = np.pad(volume, ((pad_before, pad_after), (0, 0), (0, 0)), mode='constant')\n        elif depth > self.target_depth:\n            start = (depth - self.target_depth) // 2\n            volume = volume[start:start + self.target_depth]\n\n        volume = np.expand_dims(volume, axis=0)  # [1, D, H, W]\n        volume_tensor = torch.tensor(volume, dtype=torch.float32)\n\n        return volume_tensor, tomo_id\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-07-18T17:13:46.694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset=TestTomoDataset(\n    root_dir='/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test',\n    target_shape=(64,64,64)\n)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n\n# Load model\nmodel = get_model().cuda()\nmodel.load_state_dict(torch.load(\"best_model.pth\"))  # load your trained weights\nmodel.eval()\n\n# Predict\nresults = []\nwith torch.no_grad():\n    for volume, tomo_id in test_loader:\n        volume = volume.cuda()\n        pred = model(volume)  # [1, 3]\n        pred = pred.squeeze().cpu().numpy()  # [3]\n        results.append({\n            \"tomo_id\": tomo_id[0],\n            \"Motor axis 0\": pred[0],\n            \"Motor axis 1\": pred[1],\n            \"Motor axis 2\": pred[2]\n        })\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-07-18T17:13:46.706Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_df = pd.DataFrame(results)\nsubmission_df.to_csv(\"submission.csv\", index=False)\nprint(submission_df.head())","metadata":{"trusted":true,"execution":{"execution_failed":"2025-07-18T17:13:46.748Z"}},"outputs":[],"execution_count":null}]}