{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":414283,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":338167,"modelId":359121}],"dockerImageVersionId":31040,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image, ImageFilter\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\nimport gc\n\n# 定数\nDATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"\nMODEL_PATH = \"/kaggle/input/cnn0527/pytorch/default/1/cnn_2d_fold5.pth\"\nOUTPUT_DIR = \"/kaggle/working\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-27T13:09:24.079973Z","iopub.execute_input":"2025-05-27T13:09:24.080378Z","iopub.status.idle":"2025-05-27T13:09:27.966205Z","shell.execute_reply.started":"2025-05-27T13:09:24.080342Z","shell.execute_reply":"2025-05-27T13:09:27.965203Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CNN_2D(nn.Module):\n    def __init__(self, img_size=1249, conv_drop=0.2, dense_drop=0.4):\n        super().__init__()\n        self.img_size = img_size  # 手動指定されたサイズを使う\n\n        self.features = nn.Sequential(\n            nn.Conv2d(1, 32, 3, padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.Dropout2d(conv_drop), nn.MaxPool2d(2),\n            nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.Dropout2d(conv_drop), nn.MaxPool2d(2),\n            nn.Conv2d(64,128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(), nn.Dropout2d(conv_drop), nn.MaxPool2d(2),\n            nn.Conv2d(128,256, 3, padding=1), nn.BatchNorm2d(256), nn.ReLU(), nn.Dropout2d(conv_drop), nn.MaxPool2d(2),\n            nn.Conv2d(256,256, 3, padding=1), nn.BatchNorm2d(256), nn.ReLU(), nn.Dropout2d(conv_drop), nn.MaxPool2d(2),\n        )\n\n        with torch.no_grad():\n            dummy = torch.zeros(1, 1, img_size, img_size)\n            f_dummy = self.features(dummy)\n            B, C, H, W = f_dummy.shape\n            feat_dim = C * H * W\n            h_dim = C * H\n            v_dim = C * W\n\n        self.cls_head = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(feat_dim, 512), nn.ReLU(), nn.Dropout(dense_drop),\n            nn.Linear(512, 1), nn.Sigmoid()\n        )\n\n        self.reg_head = nn.Sequential(\n            nn.Linear(h_dim + v_dim, 256), nn.ReLU(), nn.Dropout(dense_drop),\n            nn.Linear(256, 2), nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        f = self.features(x)\n        class_prob = self.cls_head(f)\n\n        h_pool = f.max(dim=3)[0]\n        v_pool = f.max(dim=2)[0]\n        h_flat = h_pool.view(h_pool.size(0), -1)\n        v_flat = v_pool.view(v_pool.size(0), -1)\n        dir_feat = torch.cat([h_flat, v_flat], dim=1)\n\n        coords = self.reg_head(dir_feat)\n        return class_prob, coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T13:09:27.967595Z","iopub.execute_input":"2025-05-27T13:09:27.967961Z","iopub.status.idle":"2025-05-27T13:09:27.980022Z","shell.execute_reply.started":"2025-05-27T13:09:27.967940Z","shell.execute_reply":"2025-05-27T13:09:27.979035Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # モデルの準備\n# model = CNN_2D().to(device)\n# model.load_state_dict(torch.load(MODEL_PATH, map_location=device))\n# model.eval()\n\n# img_size = model.img_size\n# transform = transforms.Compose([\n#     transforms.Resize((img_size, img_size), antialias=True),\n#     transforms.ToTensor(),\n#     transforms.Normalize([0.5], [0.5])\n# ])\n\n# results = []\n# test_dir = os.path.join(DATA_DIR, \"test\")\n# print(f\"🔍 test_dir: {test_dir}\")\n\n# for tomo_id in tqdm(os.listdir(test_dir), desc=\"📂 Processing tomo_ids\"):\n#     tomo_path = os.path.join(test_dir, tomo_id)\n#     if not os.path.isdir(tomo_path):\n#         print(f\"⏭️ Skipping non-directory: {tomo_path}\")\n#         continue\n\n#     print(f\"📁 Now processing: {tomo_id} ({tomo_path})\")\n\n#     for fname in sorted(os.listdir(tomo_path)):\n#         if not fname.endswith(\".jpg\"):\n#             continue\n\n#         z = int(fname.replace(\"slice_\", \"\").replace(\".jpg\", \"\"))\n#         img_path = os.path.join(tomo_path, fname)\n\n#         try:\n#             print(f\"🖼️ Loading image: {img_path}\")\n#             img = Image.open(img_path).convert(\"L\").filter(ImageFilter.SHARPEN)\n#             edge = cv2.Canny(np.array(img), 50, 150)\n#             img_tensor = transform(Image.fromarray(edge)).unsqueeze(0).to(device)\n\n#             with torch.no_grad():\n#                 class_prob, coords = model(img_tensor)\n#                 pred_cls = (class_prob.squeeze().item() > 0.5)\n\n#                 if pred_cls:\n#                     x_norm, y_norm = coords.squeeze().tolist()\n#                     axis_0 = z\n#                     axis_1 = int(round(y_norm * img_size))\n#                     axis_2 = int(round(x_norm * img_size))\n#                     print(f\"✅ Motor detected: prob=1 → (x={axis_2}, y={axis_1}, z={axis_0})\")\n#                 else:\n#                     axis_0, axis_1, axis_2 = -1, -1, -1\n#                     print(f\"❌ No motor detected: prob=0 at z={z}\")\n\n#                 results.append([tomo_id, axis_0, axis_1, axis_2])\n\n#         except Exception as e:\n#             print(f\"[ERROR] Failed to process {img_path}: {e}\")\n#             continue","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T13:09:27.981057Z","iopub.execute_input":"2025-05-27T13:09:27.981417Z","iopub.status.idle":"2025-05-27T13:09:28.007884Z","shell.execute_reply.started":"2025-05-27T13:09:27.981386Z","shell.execute_reply":"2025-05-27T13:09:28.006607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ✅ 高速化：OpenCVベース + ログ付きの推論処理\nimport cv2\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# モデルの準備\nmodel = CNN_2D(img_size=1249).to(device)\nmodel.load_state_dict(torch.load(MODEL_PATH, map_location=device))\nmodel.eval()\n\nimg_size = model.img_size\nresults = []\ntest_dir = os.path.join(DATA_DIR, \"test\")\nprint(f\"🔍 test_dir: {test_dir}\")\n\nfor tomo_id in tqdm(os.listdir(test_dir), desc=\"📂 Processing tomo_ids\"):\n    tomo_path = os.path.join(test_dir, tomo_id)\n    if not os.path.isdir(tomo_path):\n        print(f\"⏭️ Skipping non-directory: {tomo_path}\")\n        continue\n\n    print(f\"📁 Now processing: {tomo_id} ({tomo_path})\")\n\n    for fname in sorted(os.listdir(tomo_path)):\n        if not fname.endswith(\".jpg\"):\n            continue\n\n        z = int(fname.replace(\"slice_\", \"\").replace(\".jpg\", \"\"))\n        img_path = os.path.join(tomo_path, fname)\n\n        try:\n            print(f\"🖼️ Loading image: {img_path}\")\n            img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            edge = cv2.Canny(img, 50, 150)\n            edge_resized = cv2.resize(edge, (img_size, img_size), interpolation=cv2.INTER_LINEAR)\n\n            img_np = edge_resized.astype(np.float32) / 255.0\n            img_np = (img_np - 0.5) / 0.5\n            img_tensor = torch.tensor(img_np).unsqueeze(0).unsqueeze(0).to(device)\n\n            with torch.no_grad():\n                class_prob, coords = model(img_tensor)\n                pred_cls = (class_prob.squeeze().item() > 0.5)\n\n                if pred_cls:\n                    x_norm, y_norm = coords.squeeze().tolist()\n                    axis_0 = z\n                    axis_1 = int(round(y_norm * img_size))\n                    axis_2 = int(round(x_norm * img_size))\n                    print(f\"✅ Motor detected: prob=1 → (x={axis_2}, y={axis_1}, z={axis_0})\")\n                else:\n                    axis_0, axis_1, axis_2 = -1, -1, -1\n                    print(f\"❌ No motor detected: prob=0 at z={z}\")\n\n                results.append([tomo_id, axis_0, axis_1, axis_2])\n\n        except Exception as e:\n            print(f\"[ERROR] Failed to process {img_path}: {e}\")\n            continue\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T13:12:04.772344Z","iopub.execute_input":"2025-05-27T13:12:04.772710Z","iopub.status.idle":"2025-05-27T13:43:21.691686Z","shell.execute_reply.started":"2025-05-27T13:12:04.772681Z","shell.execute_reply":"2025-05-27T13:43:21.690737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.DataFrame(results, columns=[\"tomo_id\", \"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"])\nsubmission.to_csv(\"/kaggle/working/submission.csv\", index=False)\nprint(\"✅ submission.csv saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T13:09:31.797194Z","iopub.status.idle":"2025-05-27T13:09:31.797523Z","shell.execute_reply.started":"2025-05-27T13:09:31.797362Z","shell.execute_reply":"2025-05-27T13:09:31.797377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(\"submission.csv exists:\", os.path.exists(\"/kaggle/working/submission.csv\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T13:09:31.798635Z","iopub.status.idle":"2025-05-27T13:09:31.799025Z","shell.execute_reply.started":"2025-05-27T13:09:31.798830Z","shell.execute_reply":"2025-05-27T13:09:31.798850Z"}},"outputs":[],"execution_count":null}]}