{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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        print(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\n\n# Use the kagglehub client library to attach Kaggle resources like competitions, datasets, and models to your session\n# Learn more about kagglehub: https://github.com/Kaggle/kagglehub/blob/main/README.md\n\nimport kagglehub\n# kagglehub.dataset_download('<owner>/<dataset-slug>')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport gc\nimport torch\nimport numpy as np\nimport pandas as pd\nimport zarr\nfrom scipy.ndimage import maximum_filter\nfrom monai.networks.nets import UNet\n\n# ------------------------------------------------------------------\n# 1. 基本設定\n# ------------------------------------------------------------------\nVOXEL_SCALE = {\n    \"z\": 10.012444196428572,\n    \"y\": 10.012444196428572,\n    \"x\": 10.012444537618887\n}\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nMODEL_DIR = \"/kaggle/input/datasets/sokhandan/3dunet-best-negative\"\nINFERENCE_THRESH = 0.1 \n\nfrom scipy.ndimage import maximum_filter\n\ndef extract_particle_coordinates_3d(\n    heatmap,\n    threshold=0.10,\n    min_distance=4,\n    max_particles=200\n):\n    heatmap = heatmap.astype(np.float32)\n\n    local_max = (\n        maximum_filter(\n            heatmap,\n            size=min_distance * 2 + 1\n        ) == heatmap\n    )\n\n    mask = (\n        heatmap > threshold\n    ) & local_max\n\n    coords = np.argwhere(mask)\n\n    if len(coords) == 0:\n        return np.empty((0, 3), dtype=np.int32)\n\n    scores = heatmap[\n        coords[:, 0],\n        coords[:, 1],\n        coords[:, 2]\n    ]\n\n    order = np.argsort(scores)[::-1]\n\n    return coords[\n        order[:max_particles]\n    ]\n\nKAGGLE_INPUT_DIR = \"/kaggle/input/competitions/czii-cryo-et-object-identification\"\ntest_zarr_paths = glob.glob(os.path.join(KAGGLE_INPUT_DIR, \"**\", \"test\", \"**\", \"*.zarr\"), recursive=True)\n\nif len(test_zarr_paths) == 0:\n    all_zarrs = glob.glob(os.path.join(KAGGLE_INPUT_DIR, \"**\", \"*.zarr\"), recursive=True)\n    test_zarr_paths = [p for p in all_zarrs if \"test\" in p or \"ExperimentRuns\" in p]\n\nprint(\"===== TEST DATA =====\")\nprint(\"件数:\", len(test_zarr_paths))\n\nfor p in test_zarr_paths:\n    print(p)\n    \nsubmission_path = \"submission.csv\"\n\npd.DataFrame(\n    columns=[\"experiment\", \"particle_type\", \"x\", \"y\", \"z\"]\n).to_csv(submission_path, index=False)\n\n# ==================================================================\n# 🌟 STEP 1: Group 0 (大型粒子 3種) の推論\n# ==================================================================\ng0_path = os.path.join(MODEL_DIR, \"3dunet_group0_best (1).pth\")\nif os.path.exists(g0_path):\n    print(\"🚀 Group 0 の推論を開始します...\")\n    model_g0 = UNet(spatial_dims=3, in_channels=1, out_channels=3, channels=(8, 16, 32, 64), strides=(2, 2, 2), num_res_units=1, norm=\"batch\").to(device)\n    model_g0.load_state_dict(torch.load(g0_path, map_location=device))\n    model_g0.eval()\n\n    if hasattr(torch, \"compile\"):\n        model_g0 = torch.compile(model_g0)\n\n    g0_particles = [\"ribosome\", \"virus-like-particle\", \"thyroglobulin\"]\n\n    for z_path in test_zarr_paths:\n        import re\n\n        m = re.search(r\"TS_\\d+_\\d+\", z_path)\n        \n        if m:\n            exp_id = m.group(0)\n        else:\n            exp_id = \"TS_sample\"\n\n        print(\"Processing:\", exp_id)\n        \n        z_root = zarr.open(z_path, mode='r')\n        first_key = '0' if '0' in z_root else (list(z_root.keys())[0] if hasattr(z_root, 'keys') else None)\n        tomo_data = np.array(z_root[first_key] if first_key else z_root, dtype=np.float32)\n        \n        std_val = np.std(tomo_data)\n        tomo_data = (tomo_data - np.mean(tomo_data)) / (std_val if std_val > 1e-7 else 1.0)\n        \n        tz, ty, tx = tomo_data.shape\n        pz, py, px = 96, 96, 96\n        stride = 64  # :星2: ストライドを 80 に広げて推論回数とメモリ保持を極減\n\n        prob_map = np.zeros((3, tz, ty, tx), dtype=np.float16)\n        count_map = np.zeros((tz, ty, tx), dtype=np.float16)\n\n        z_steps = list(range(0, max(1, tz - pz + 1), stride))\n        if z_steps[-1] + pz < tz: z_steps.append(tz - pz)\n        y_steps = list(range(0, max(1, ty - py + 1), stride))\n        if y_steps[-1] + py < ty: y_steps.append(ty - py)\n        x_steps = list(range(0, max(1, tx - px + 1), stride))\n        if x_steps[-1] + px < tx: x_steps.append(tx - px)\n\n        with torch.inference_mode():\n            coords = [\n                (z,y,x)\n                for z in z_steps\n                for y in y_steps\n                for x in x_steps\n            ]\n            \n            BATCH_SIZE = 16\n            \n            with torch.inference_mode():\n                with torch.cuda.amp.autocast():\n            \n                    for i in range(\n                        0,\n                        len(coords),\n                        BATCH_SIZE\n                    ):\n\n                        batch_coords = coords[\n                            i:i+BATCH_SIZE\n                        ]\n            \n                        patches = [\n                            tomo_data[\n                                z:z+pz,\n                                y:y+py,\n                                x:x+px\n                            ]\n                            for z,y,x in batch_coords\n                        ]\n            \n                        batch = np.stack(\n                            patches\n                        ).astype(np.float32)\n\n                        tensor = (\n                            torch.from_numpy(batch)\n                            .unsqueeze(1)\n                            .to(device)\n                        )\n            \n                        raw_preds = model_g0(tensor)\n\n                        print(\n                            \"RAW:\",\n                            float(raw_preds.min()),\n                            float(raw_preds.max()),\n                            float(raw_preds.mean())\n                        )\n                        \n                        sig_preds = torch.sigmoid(raw_preds)\n                        \n                        print(\n                            \"SIG:\",\n                            float(sig_preds.min()),\n                            float(sig_preds.max()),\n                            float(sig_preds.mean())\n                        )\n                        \n                        preds = sig_preds.cpu().numpy()\n            \n                        for pred,(z,y,x) in zip(\n                            preds,\n                            batch_coords\n                        ):\n\n                            prob_map[\n                                :,\n                                z:z+pz,\n                                y:y+py,\n                                x:x+px\n                            ] += pred.astype(\n                                np.float16\n                            )\n            \n                            count_map[\n                                z:z+pz,\n                                y:y+py,\n                                x:x+px\n                            ] += np.float16(1)\n                        \n        count_map = np.maximum(count_map, np.float16(1.0))\n        prob_map /= count_map\n\n        print(f\"\\n===== {exp_id} =====\")\n\n        for c_idx, p_name in enumerate(g0_particles):\n            pmap = prob_map[c_idx].astype(np.float32)\n        \n            print(\n                p_name,\n                \"min=\", float(pmap.min()),\n                \"max=\", float(pmap.max()),\n                \"mean=\", float(pmap.mean())\n            )\n            \n        for c_idx, p_name in enumerate(g0_particles):\n            pmap = prob_map[c_idx].astype(np.float32)\n        \n            coords = extract_particle_coordinates_3d(\n                pmap,\n                threshold=0.15,\n                min_distance=6,\n                max_particles=150\n            )\n\n            print(\n                p_name,\n                \"detected=\",\n                len(coords)\n            )\n\n            print(f\"\\n===== {p_name} =====\")\n            print(\"PRED\")\n            print(coords[:20])\n                    \n            if len(coords) == 0:\n                flat_idx = np.argsort(pmap.ravel())[-5:]\n                coords = np.array(\n                    np.unravel_index(flat_idx, pmap.shape)\n                ).T\n\n            rows = []\n        \n            for c in coords:\n                rows.append({\n                    'experiment': exp_id,\n                    'particle_type': p_name,\n                    'x': c[2] * VOXEL_SCALE['x'],\n                    'y': c[1] * VOXEL_SCALE['y'],\n                    'z': c[0] * VOXEL_SCALE['z']\n                })\n        \n            if rows:\n                pd.DataFrame(rows).to_csv(\n                    submission_path,\n                    mode=\"a\",\n                    header=False,\n                    index=False\n                )\n        \n            del rows\n            del coords\n            del pmap\n\n        del tomo_data\n        del prob_map\n        del count_map\n        \n    # 🌟 Group 0 モデルを完全に解放\n    del model_g0\n\n# ==================================================================\n# 🌟 STEP 2: Group 1 (小型粒子 3種) の推論\n# ==================================================================\ng1_path = os.path.join(MODEL_DIR, \"3dunet_group1_best (1).pth\")\nif os.path.exists(g1_path):\n    print(\"🚀 Group 1 の推論を開始します...\")\n    model_g1 = UNet(spatial_dims=3, in_channels=1, out_channels=3, channels=(8, 16, 32, 64), strides=(2, 2, 2), num_res_units=1, norm=\"batch\").to(device)\n    model_g1.load_state_dict(torch.load(g1_path, map_location=device))\n    model_g1.eval()\n\n    if hasattr(torch, \"compile\"):\n        model_g1 = torch.compile(model_g1)\n\n    g1_particles = [\"beta-galactosidase\", \"beta-amylase\", \"apo-ferritin\"]\n\n    for z_path in test_zarr_paths:\n        import re\n\n        m = re.search(r\"TS_\\d+_\\d+\", z_path)\n        \n        if m:\n            exp_id = m.group(0)\n        else:\n            exp_id = \"TS_sample\"\n\n        print(\"Processing:\", exp_id)\n        \n        z_root = zarr.open(z_path, mode='r')\n        first_key = '0' if '0' in z_root else (list(z_root.keys())[0] if hasattr(z_root, 'keys') else None)\n        tomo_data = np.array(z_root[first_key] if first_key else z_root, dtype=np.float32)\n        \n        std_val = np.std(tomo_data)\n        tomo_data = (tomo_data - np.mean(tomo_data)) / (std_val if std_val > 1e-7 else 1.0)\n        \n        tz, ty, tx = tomo_data.shape\n        pz, py, px = 96, 96, 96\n        stride = 64\n\n        prob_map = np.zeros((3, tz, ty, tx), dtype=np.float16)\n        count_map = np.zeros((tz, ty, tx), dtype=np.float16)\n\n        z_steps = list(range(0, max(1, tz - pz + 1), stride))\n        if z_steps[-1] + pz < tz: z_steps.append(tz - pz)\n        y_steps = list(range(0, max(1, ty - py + 1), stride))\n        if y_steps[-1] + py < ty: y_steps.append(ty - py)\n        x_steps = list(range(0, max(1, tx - px + 1), stride))\n        if x_steps[-1] + px < tx: x_steps.append(tx - px)\n\n        with torch.inference_mode():\n            coords = [\n                (z,y,x)\n                for z in z_steps\n                for y in y_steps\n                for x in x_steps\n            ]\n\n            BATCH_SIZE = 16\n            \n            with torch.inference_mode():\n                with torch.cuda.amp.autocast():\n            \n                    for i in range(\n                        0,\n                        len(coords),\n                        BATCH_SIZE\n                    ):\n\n                        batch_coords = coords[\n                            i:i+BATCH_SIZE\n                        ]\n            \n                        patches = [\n                            tomo_data[\n                                z:z+pz,\n                                y:y+py,\n                                x:x+px\n                            ]\n                            for z,y,x in batch_coords\n                        ]\n\n                        batch = np.stack(\n                            patches\n                        ).astype(np.float32)\n            \n                        tensor = (\n                            torch.from_numpy(batch)\n                            .unsqueeze(1)\n                            .to(device)\n                        )\n\n                        raw_preds = model_g1(tensor)\n\n                        print(\n                            \"RAW:\",\n                            float(raw_preds.min()),\n                            float(raw_preds.max()),\n                            float(raw_preds.mean())\n                        )\n                        \n                        sig_preds = torch.sigmoid(raw_preds)\n                        \n                        print(\n                            \"SIG:\",\n                            float(sig_preds.min()),\n                            float(sig_preds.max()),\n                            float(sig_preds.mean())\n                        )\n                        \n                        preds = sig_preds.cpu().numpy()\n\n                        for pred,(z,y,x) in zip(\n                            preds,\n                            batch_coords\n                        ):\n            \n                            prob_map[\n                                :,\n                                z:z+pz,\n                                y:y+py,\n                                x:x+px\n                            ] += pred.astype(\n                                np.float16\n                            )\n\n                            count_map[\n                                z:z+pz,\n                                y:y+py,\n                                x:x+px\n                            ] += np.float16(1)\n\n        count_map = np.maximum(count_map, np.float16(1.0))\n        prob_map /= count_map\n\n        print(f\"\\n===== {exp_id} =====\")\n\n        for c_idx, p_name in enumerate(g1_particles):\n            pmap = prob_map[c_idx].astype(np.float32)\n        \n            print(\n                p_name,\n                \"min=\", float(pmap.min()),\n                \"max=\", float(pmap.max()),\n                \"mean=\", float(pmap.mean())\n            )\n\n        for c_idx, p_name in enumerate(g1_particles):\n            pmap = prob_map[c_idx].astype(np.float32)\n        \n            coords = extract_particle_coordinates_3d(\n                pmap,\n                threshold=0.10,\n                min_distance=4,\n                max_particles=200\n            )\n\n            print(\n                p_name,\n                \"detected=\",\n                len(coords)\n            )\n\n            print(f\"\\n===== {p_name} =====\")\n            print(\"PRED\")\n            print(coords[:20])\n        \n            if len(coords) == 0:\n                flat_idx = np.argsort(pmap.ravel())[-5:]\n                coords = np.array(\n                    np.unravel_index(flat_idx, pmap.shape)\n                ).T\n        \n            rows = []\n\n            for c in coords:\n                rows.append({\n                    'experiment': exp_id,\n                    'particle_type': p_name,\n                    'x': c[2] * VOXEL_SCALE['x'],\n                    'y': c[1] * VOXEL_SCALE['y'],\n                    'z': c[0] * VOXEL_SCALE['z']\n                })\n\n            if rows:\n                pd.DataFrame(rows).to_csv(\n                    submission_path,\n                    mode=\"a\",\n                    header=False,\n                    index=False\n                )\n        \n            del rows\n            del coords\n            del pmap\n\n        del tomo_data\n        del prob_map\n        del count_map\n\n    del model_g1\n\n# ------------------------------------------------------------------\n# 3. 出力\n# ------------------------------------------------------------------\nsub_df = pd.read_csv(submission_path)\n\nsub_df = sub_df.dropna()\n\nsub_df = sub_df[\n    [\"experiment\",\n     \"particle_type\",\n     \"x\",\n     \"y\",\n     \"z\"]\n]\n\nsub_df.insert(\n    0,\n    \"id\",\n    np.arange(\n        len(sub_df),\n        dtype=np.int64\n    )\n)\n\nsub_df.to_csv(\n    submission_path,\n    index=False\n)\n\nprint(sub_df.head())\nprint(f\"🎉 処理完了！ 生成件数: {len(sub_df)} 件\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}