{
  "id": 612246,
  "title": "3D Circle of Willis Segmentation with a Stacked U-Net ",
  "url": "/competitions/rsna-intracranial-aneurysm-detection/discussion/612246",
  "author_name": "suguuuuu",
  "post_date": "2025-10-18T01:25:49.227000",
  "votes": 12,
  "comment_count": 1,
  "views": 0,
  "content": "<h1>Introduction</h1>\n<p>I struggled with the <strong>3D multi-class segmentation task</strong>, so this post documents my process, experiments, and key takeaways.</p>\n<hr>\n<h1>Background</h1>\n<p>Some participants might have succeeded with a straightforward one-stage training pipeline — if so, please share your method with me!<br>\nThe following describes my workaround approach that eventually stabilized training.</p>\n<hr>\n<h1>Preprocessing</h1>\n<p>Due to the nature of the distributed data, the <strong>imaging range</strong> (number of slices, spatial resolution, etc.) varies significantly between cases.<br>\nTo stabilize training, I added a preprocessing step to <strong>extract and crop the arterial region</strong>.</p>\n<p>This normalization reduced variation in pixel intensity and anatomical alignment, improving segmentation accuracy.<br>\nAll aneurysm segmentation was performed <strong>within the cropped arterial region</strong>.</p>\n<p>Figure: Example of cropped area (bounding box around the provided arterial segmentation mask).<br>\n<img src=\"https://storage.googleapis.com/zenn-user-upload/d6fb453b9eb2-20251018.png\" alt=\"\"></p>\n<hr>\n<h1>Final Approach</h1>\n<p>The final setup used a <strong>two-stage cascaded 3D U-Net architecture</strong>.<br>\nAt the intermediate stage, I added an auxiliary output that predicts a unified “artery” foreground class for additional supervision.<br>\nThis multi-level supervision stabilized the final multi-class aneurysm segmentation.</p>\n<p><strong>Architecture Overview:</strong></p>\n<ol>\n<li><p><strong>Stage 1:</strong> A large model pretrained on brain datasets (public MONAI model).</p>\n<ul>\n<li>Fixed choice to leverage pretrained weights due to limited labeled data.</li></ul></li>\n<li><p><strong>Stage 2:</strong> A smaller, scratch-trained model for detailed segmentation.</p></li>\n</ol>\n<p><strong>Training procedure:</strong></p>\n<ol>\n<li>Train the Stage 1 model with <strong>binary segmentation</strong> (foreground vs. background).</li>\n<li>Train the Stage 2 model with <strong>multi-class segmentation</strong> (13 arteries + background).</li>\n<li>Use different <strong>loss weights</strong> for Stage 1 and Stage 2 outputs during training.</li>\n</ol>\n<p>Figure: Model architecture overview<br>\n<img src=\"https://storage.googleapis.com/zenn-user-upload/293784878644-20251018.png\" alt=\"\"></p>\n<p>This <strong>two-phase (or staged fine-tuning)</strong> approach yielded stable and reproducible results.</p>\n<hr>\n<h1>Iterative Development Process</h1>\n<p>Here’s how the final design evolved through experiments:</p>\n<ol>\n<li><strong>Direct multi-class training from the start</strong><br>\n→ Result: Unstable, poor convergence, scattered false positives.</li>\n<li><strong>Merged all arteries into a single “foreground” class</strong> (binary setup)<br>\n→ Result: The model learned vessel structures well.</li>\n<li><strong>Separated vessel detection and class discrimination</strong> into two stages (cascade)<br>\n→ Result: Stable and consistent performance.</li>\n</ol>\n<ul>\n<li>Inspiration: Multi-stage refinement designs like <strong>OpenPose</strong>, which progressively refine outputs using previous predictions.</li>\n<li>A single-stage end-to-end setup was too unstable, so splitting into <strong>Stage 1 → Stage 2</strong> was essential.</li>\n</ul>\n<p>Ultimately, the best results came from <strong>two-stage sequential training</strong> rather than one end-to-end process.</p>\n<hr>\n<h1>Example Code (Model Definition)</h1>\n<p>Below is a simplified example of how I connected a pretrained MONAI segmentation model and a custom U-Net into a cascade structure.</p>\n<pre><code> monai.networks.nets  UNet\n monai.bundle  load\n\n ():\n    device = config.device\n    K = (config, , )  \n\n    \n    stage1 = load(\n        name=,\n        source=,\n        device=device,\n        net_override={\n            : config.in_channels,\n            : ,\n        },\n    )\n\n    \n    stage2_in_ch = config.in_channels + \n    stage2 = UNet(\n        spatial_dims=,\n        in_channels=stage2_in_ch,\n        out_channels=K,\n        channels=(, , , ),\n        strides=(, , ),\n        num_res_units=,\n        norm=Norm.BATCH,\n    ).to(device)\n\n     (nn.Module):\n         ():\n            ().__init__()\n            .stage1 = s1\n            .stage2 = s2\n\n         ():\n             torch.is_tensor(out):\n                 out\n             (out, ):\n                 v  out.values():\n                     torch.is_tensor(v):\n                         v\n             RuntimeError()\n\n         ():\n            \n            logit1 = ._take_tensor(.stage1(x))  \n            prob1  = F.softmax(logit1, dim=)\n            fg1    = prob1[:, :, ...]                 \n\n            \n            x2 = torch.cat([x, fg1], dim=)\n            logit2 = .stage2(x2)\n\n             {\n                : logit1,   \n                : logit2,     \n            }\n\n    model = CascadeSeg(stage1.to(device), stage2).to(device)\n     model\n</code></pre>\n<h3>Notes</h3>\n<ul>\n<li>Since <code>stage1</code> may return a dictionary, <code>_take_tensor</code> extracts the first tensor it finds.</li>\n<li>The Stage 1 output is converted to a <strong>foreground probability map</strong>, concatenated with the input, and passed to Stage 2.</li>\n<li>Stage 2 is a standard U-Net.</li>\n<li>Both <strong>coarse (2-class)</strong> and <strong>fine (multi-class)</strong> outputs can contribute to the overall loss.</li>\n</ul>",
  "messages": [
    {
      "id": 3303421,
      "postDate": "2025-10-18T01:25:49.227Z",
      "content": "<h1>Introduction</h1>\n<p>I struggled with the <strong>3D multi-class segmentation task</strong>, so this post documents my process, experiments, and key takeaways.</p>\n<hr>\n<h1>Background</h1>\n<p>Some participants might have succeeded with a straightforward one-stage training pipeline — if so, please share your method with me!<br>\nThe following describes my workaround approach that eventually stabilized training.</p>\n<hr>\n<h1>Preprocessing</h1>\n<p>Due to the nature of the distributed data, the <strong>imaging range</strong> (number of slices, spatial resolution, etc.) varies significantly between cases.<br>\nTo stabilize training, I added a preprocessing step to <strong>extract and crop the arterial region</strong>.</p>\n<p>This normalization reduced variation in pixel intensity and anatomical alignment, improving segmentation accuracy.<br>\nAll aneurysm segmentation was performed <strong>within the cropped arterial region</strong>.</p>\n<p>Figure: Example of cropped area (bounding box around the provided arterial segmentation mask).<br>\n<img src=\"https://storage.googleapis.com/zenn-user-upload/d6fb453b9eb2-20251018.png\" alt=\"\"></p>\n<hr>\n<h1>Final Approach</h1>\n<p>The final setup used a <strong>two-stage cascaded 3D U-Net architecture</strong>.<br>\nAt the intermediate stage, I added an auxiliary output that predicts a unified “artery” foreground class for additional supervision.<br>\nThis multi-level supervision stabilized the final multi-class aneurysm segmentation.</p>\n<p><strong>Architecture Overview:</strong></p>\n<ol>\n<li><p><strong>Stage 1:</strong> A large model pretrained on brain datasets (public MONAI model).</p>\n<ul>\n<li>Fixed choice to leverage pretrained weights due to limited labeled data.</li></ul></li>\n<li><p><strong>Stage 2:</strong> A smaller, scratch-trained model for detailed segmentation.</p></li>\n</ol>\n<p><strong>Training procedure:</strong></p>\n<ol>\n<li>Train the Stage 1 model with <strong>binary segmentation</strong> (foreground vs. background).</li>\n<li>Train the Stage 2 model with <strong>multi-class segmentation</strong> (13 arteries + background).</li>\n<li>Use different <strong>loss weights</strong> for Stage 1 and Stage 2 outputs during training.</li>\n</ol>\n<p>Figure: Model architecture overview<br>\n<img src=\"https://storage.googleapis.com/zenn-user-upload/293784878644-20251018.png\" alt=\"\"></p>\n<p>This <strong>two-phase (or staged fine-tuning)</strong> approach yielded stable and reproducible results.</p>\n<hr>\n<h1>Iterative Development Process</h1>\n<p>Here’s how the final design evolved through experiments:</p>\n<ol>\n<li><strong>Direct multi-class training from the start</strong><br>\n→ Result: Unstable, poor convergence, scattered false positives.</li>\n<li><strong>Merged all arteries into a single “foreground” class</strong> (binary setup)<br>\n→ Result: The model learned vessel structures well.</li>\n<li><strong>Separated vessel detection and class discrimination</strong> into two stages (cascade)<br>\n→ Result: Stable and consistent performance.</li>\n</ol>\n<ul>\n<li>Inspiration: Multi-stage refinement designs like <strong>OpenPose</strong>, which progressively refine outputs using previous predictions.</li>\n<li>A single-stage end-to-end setup was too unstable, so splitting into <strong>Stage 1 → Stage 2</strong> was essential.</li>\n</ul>\n<p>Ultimately, the best results came from <strong>two-stage sequential training</strong> rather than one end-to-end process.</p>\n<hr>\n<h1>Example Code (Model Definition)</h1>\n<p>Below is a simplified example of how I connected a pretrained MONAI segmentation model and a custom U-Net into a cascade structure.</p>\n<pre><code> monai.networks.nets  UNet\n monai.bundle  load\n\n ():\n    device = config.device\n    K = (config, , )  \n\n    \n    stage1 = load(\n        name=,\n        source=,\n        device=device,\n        net_override={\n            : config.in_channels,\n            : ,\n        },\n    )\n\n    \n    stage2_in_ch = config.in_channels + \n    stage2 = UNet(\n        spatial_dims=,\n        in_channels=stage2_in_ch,\n        out_channels=K,\n        channels=(, , , ),\n        strides=(, , ),\n        num_res_units=,\n        norm=Norm.BATCH,\n    ).to(device)\n\n     (nn.Module):\n         ():\n            ().__init__()\n            .stage1 = s1\n            .stage2 = s2\n\n         ():\n             torch.is_tensor(out):\n                 out\n             (out, ):\n                 v  out.values():\n                     torch.is_tensor(v):\n                         v\n             RuntimeError()\n\n         ():\n            \n            logit1 = ._take_tensor(.stage1(x))  \n            prob1  = F.softmax(logit1, dim=)\n            fg1    = prob1[:, :, ...]                 \n\n            \n            x2 = torch.cat([x, fg1], dim=)\n            logit2 = .stage2(x2)\n\n             {\n                : logit1,   \n                : logit2,     \n            }\n\n    model = CascadeSeg(stage1.to(device), stage2).to(device)\n     model\n</code></pre>\n<h3>Notes</h3>\n<ul>\n<li>Since <code>stage1</code> may return a dictionary, <code>_take_tensor</code> extracts the first tensor it finds.</li>\n<li>The Stage 1 output is converted to a <strong>foreground probability map</strong>, concatenated with the input, and passed to Stage 2.</li>\n<li>Stage 2 is a standard U-Net.</li>\n<li>Both <strong>coarse (2-class)</strong> and <strong>fine (multi-class)</strong> outputs can contribute to the overall loss.</li>\n</ul>",
      "rawMarkdown": "# Introduction\n I struggled with the **3D multi-class segmentation task**, so this post documents my process, experiments, and key takeaways.\n\n---\n\n# Background\n\nSome participants might have succeeded with a straightforward one-stage training pipeline — if so, please share your method with me!\nThe following describes my workaround approach that eventually stabilized training.\n\n---\n\n# Preprocessing\n\nDue to the nature of the distributed data, the **imaging range** (number of slices, spatial resolution, etc.) varies significantly between cases.\nTo stabilize training, I added a preprocessing step to **extract and crop the arterial region**.\n\nThis normalization reduced variation in pixel intensity and anatomical alignment, improving segmentation accuracy.\nAll aneurysm segmentation was performed **within the cropped arterial region**.\n\nFigure: Example of cropped area (bounding box around the provided arterial segmentation mask).\n![](https://storage.googleapis.com/zenn-user-upload/d6fb453b9eb2-20251018.png)\n\n---\n\n# Final Approach\n\nThe final setup used a **two-stage cascaded 3D U-Net architecture**.\nAt the intermediate stage, I added an auxiliary output that predicts a unified “artery” foreground class for additional supervision.\nThis multi-level supervision stabilized the final multi-class aneurysm segmentation.\n\n**Architecture Overview:**\n\n1. **Stage 1:** A large model pretrained on brain datasets (public MONAI model).\n\n   * Fixed choice to leverage pretrained weights due to limited labeled data.\n2. **Stage 2:** A smaller, scratch-trained model for detailed segmentation.\n\n**Training procedure:**\n\n1. Train the Stage 1 model with **binary segmentation** (foreground vs. background).\n2. Train the Stage 2 model with **multi-class segmentation** (13 arteries + background).\n3. Use different **loss weights** for Stage 1 and Stage 2 outputs during training.\n\nFigure: Model architecture overview\n![](https://storage.googleapis.com/zenn-user-upload/293784878644-20251018.png)\n\nThis **two-phase (or staged fine-tuning)** approach yielded stable and reproducible results.\n\n---\n\n# Iterative Development Process\n\nHere’s how the final design evolved through experiments:\n\n1. **Direct multi-class training from the start**\n   → Result: Unstable, poor convergence, scattered false positives.\n2. **Merged all arteries into a single “foreground” class** (binary setup)\n   → Result: The model learned vessel structures well.\n3. **Separated vessel detection and class discrimination** into two stages (cascade)\n   → Result: Stable and consistent performance.\n\n* Inspiration: Multi-stage refinement designs like **OpenPose**, which progressively refine outputs using previous predictions.\n* A single-stage end-to-end setup was too unstable, so splitting into **Stage 1 → Stage 2** was essential.\n\nUltimately, the best results came from **two-stage sequential training** rather than one end-to-end process.\n\n---\n\n# Example Code (Model Definition)\n\nBelow is a simplified example of how I connected a pretrained MONAI segmentation model and a custom U-Net into a cascade structure.\n\n```python\n\nfrom monai.networks.nets import UNet\nfrom monai.bundle import load\n\ndef create_model(config):\n    device = config.device\n    K = getattr(config, \"num_artery_classes\", 14)  # total classes\n\n    # ----- Stage 1: Binary (background / foreground) -----\n    stage1 = load(\n        name=\"brats_mri_segmentation\",\n        source=\"monaihosting\",\n        device=device,\n        net_override={\n            \"in_channels\": config.in_channels,\n            \"out_channels\": 2,\n        },\n    )\n\n    # ----- Stage 2: input = image + Stage1 foreground prob (1 channel added) -----\n    stage2_in_ch = config.in_channels + 1\n    stage2 = UNet(\n        spatial_dims=3,\n        in_channels=stage2_in_ch,\n        out_channels=K,\n        channels=(16, 32, 64, 128),\n        strides=(2, 2, 2),\n        num_res_units=1,\n        norm=Norm.BATCH,\n    ).to(device)\n\n    class CascadeSeg(nn.Module):\n        def __init__(self, s1, s2):\n            super().__init__()\n            self.stage1 = s1\n            self.stage2 = s2\n\n        def _take_tensor(self, out):\n            if torch.is_tensor(out):\n                return out\n            if isinstance(out, dict):\n                for v in out.values():\n                    if torch.is_tensor(v):\n                        return v\n            raise RuntimeError(\"Stage1 output tensor not found\")\n\n        def forward(self, x):\n            # ---- Stage 1 ----\n            logit1 = self._take_tensor(self.stage1(x))  # [B,2,D,H,W]\n            prob1  = F.softmax(logit1, dim=1)\n            fg1    = prob1[:, 1:2, ...]                 # foreground prob\n\n            # ---- Stage 2 ----\n            x2 = torch.cat([x, fg1], dim=1)\n            logit2 = self.stage2(x2)\n\n            return {\n                \"coarse\": logit1,   # Stage 1 logits\n                \"fine\": logit2,     # Stage 2 logits\n            }\n\n    model = CascadeSeg(stage1.to(device), stage2).to(device)\n    return model\n```\n\n### Notes\n\n* Since `stage1` may return a dictionary, `_take_tensor` extracts the first tensor it finds.\n* The Stage 1 output is converted to a **foreground probability map**, concatenated with the input, and passed to Stage 2.\n* Stage 2 is a standard U-Net.\n* Both **coarse (2-class)** and **fine (multi-class)** outputs can contribute to the overall loss.\n",
      "votes": 12
    },
    {
      "id": 3452364,
      "postDate": "2026-05-03T07:14:23.740Z",
      "content": "<p><a href=\"https://www.kaggle.com/sugupoko\" target=\"_blank\">@sugupoko</a> hi, thx for this cascaded solution. is there any code that i can reproduce? i am currently interested in medical imaging domain.</p>",
      "rawMarkdown": "@sugupoko hi, thx for this cascaded solution. is there any code that i can reproduce? i am currently interested in medical imaging domain."
    }
  ],
  "comments": [
    {
      "id": 3452364,
      "author_name": "Simon Beck",
      "author_url": "",
      "post_date": "2026-05-03T07:14:23.740000",
      "content": "<p><a href=\"https://www.kaggle.com/sugupoko\" target=\"_blank\">@sugupoko</a> hi, thx for this cascaded solution. is there any code that i can reproduce? i am currently interested in medical imaging domain.</p>",
      "votes": 0,
      "replies": []
    }
  ],
  "raw_markdown_by_id": {
    "3303421": "# Introduction\n I struggled with the **3D multi-class segmentation task**, so this post documents my process, experiments, and key takeaways.\n\n---\n\n# Background\n\nSome participants might have succeeded with a straightforward one-stage training pipeline — if so, please share your method with me!\nThe following describes my workaround approach that eventually stabilized training.\n\n---\n\n# Preprocessing\n\nDue to the nature of the distributed data, the **imaging range** (number of slices, spatial resolution, etc.) varies significantly between cases.\nTo stabilize training, I added a preprocessing step to **extract and crop the arterial region**.\n\nThis normalization reduced variation in pixel intensity and anatomical alignment, improving segmentation accuracy.\nAll aneurysm segmentation was performed **within the cropped arterial region**.\n\nFigure: Example of cropped area (bounding box around the provided arterial segmentation mask).\n![](https://storage.googleapis.com/zenn-user-upload/d6fb453b9eb2-20251018.png)\n\n---\n\n# Final Approach\n\nThe final setup used a **two-stage cascaded 3D U-Net architecture**.\nAt the intermediate stage, I added an auxiliary output that predicts a unified “artery” foreground class for additional supervision.\nThis multi-level supervision stabilized the final multi-class aneurysm segmentation.\n\n**Architecture Overview:**\n\n1. **Stage 1:** A large model pretrained on brain datasets (public MONAI model).\n\n   * Fixed choice to leverage pretrained weights due to limited labeled data.\n2. **Stage 2:** A smaller, scratch-trained model for detailed segmentation.\n\n**Training procedure:**\n\n1. Train the Stage 1 model with **binary segmentation** (foreground vs. background).\n2. Train the Stage 2 model with **multi-class segmentation** (13 arteries + background).\n3. Use different **loss weights** for Stage 1 and Stage 2 outputs during training.\n\nFigure: Model architecture overview\n![](https://storage.googleapis.com/zenn-user-upload/293784878644-20251018.png)\n\nThis **two-phase (or staged fine-tuning)** approach yielded stable and reproducible results.\n\n---\n\n# Iterative Development Process\n\nHere’s how the final design evolved through experiments:\n\n1. **Direct multi-class training from the start**\n   → Result: Unstable, poor convergence, scattered false positives.\n2. **Merged all arteries into a single “foreground” class** (binary setup)\n   → Result: The model learned vessel structures well.\n3. **Separated vessel detection and class discrimination** into two stages (cascade)\n   → Result: Stable and consistent performance.\n\n* Inspiration: Multi-stage refinement designs like **OpenPose**, which progressively refine outputs using previous predictions.\n* A single-stage end-to-end setup was too unstable, so splitting into **Stage 1 → Stage 2** was essential.\n\nUltimately, the best results came from **two-stage sequential training** rather than one end-to-end process.\n\n---\n\n# Example Code (Model Definition)\n\nBelow is a simplified example of how I connected a pretrained MONAI segmentation model and a custom U-Net into a cascade structure.\n\n```python\n\nfrom monai.networks.nets import UNet\nfrom monai.bundle import load\n\ndef create_model(config):\n    device = config.device\n    K = getattr(config, \"num_artery_classes\", 14)  # total classes\n\n    # ----- Stage 1: Binary (background / foreground) -----\n    stage1 = load(\n        name=\"brats_mri_segmentation\",\n        source=\"monaihosting\",\n        device=device,\n        net_override={\n            \"in_channels\": config.in_channels,\n            \"out_channels\": 2,\n        },\n    )\n\n    # ----- Stage 2: input = image + Stage1 foreground prob (1 channel added) -----\n    stage2_in_ch = config.in_channels + 1\n    stage2 = UNet(\n        spatial_dims=3,\n        in_channels=stage2_in_ch,\n        out_channels=K,\n        channels=(16, 32, 64, 128),\n        strides=(2, 2, 2),\n        num_res_units=1,\n        norm=Norm.BATCH,\n    ).to(device)\n\n    class CascadeSeg(nn.Module):\n        def __init__(self, s1, s2):\n            super().__init__()\n            self.stage1 = s1\n            self.stage2 = s2\n\n        def _take_tensor(self, out):\n            if torch.is_tensor(out):\n                return out\n            if isinstance(out, dict):\n                for v in out.values():\n                    if torch.is_tensor(v):\n                        return v\n            raise RuntimeError(\"Stage1 output tensor not found\")\n\n        def forward(self, x):\n            # ---- Stage 1 ----\n            logit1 = self._take_tensor(self.stage1(x))  # [B,2,D,H,W]\n            prob1  = F.softmax(logit1, dim=1)\n            fg1    = prob1[:, 1:2, ...]                 # foreground prob\n\n            # ---- Stage 2 ----\n            x2 = torch.cat([x, fg1], dim=1)\n            logit2 = self.stage2(x2)\n\n            return {\n                \"coarse\": logit1,   # Stage 1 logits\n                \"fine\": logit2,     # Stage 2 logits\n            }\n\n    model = CascadeSeg(stage1.to(device), stage2).to(device)\n    return model\n```\n\n### Notes\n\n* Since `stage1` may return a dictionary, `_take_tensor` extracts the first tensor it finds.\n* The Stage 1 output is converted to a **foreground probability map**, concatenated with the input, and passed to Stage 2.\n* Stage 2 is a standard U-Net.\n* Both **coarse (2-class)** and **fine (multi-class)** outputs can contribute to the overall loss.\n",
    "3452364": "@sugupoko hi, thx for this cascaded solution. is there any code that i can reproduce? i am currently interested in medical imaging domain."
  }
}