Joint Lung + Nodule Segmentation — 3D U-Net / DynUNet (ex-LIDC)

Single-model, end-to-end 3-class segmentation of chest CT: predicts background / lung / nodule in one forward pass, directly from the full CT (no ROI stage, no bbox crop). Trained on the NLST + NSCLC-Radiomics subset of the unified corpus — LIDC-IDRI is excluded because it lacks ground-truth 2D lung labels.

Companion to joint-dynunet-3d-pseudo-lidc (same architecture, LIDC included via ROI-model-derived pseudo-lung labels), and to joint-segresnet-3d-{ex,pseudo}-lidc (same task, different architecture).

Model details

  • Architecture: MONAI DynUNet — nnU-Net-style dynamic 3D U-Net with per-level kernels/strides and instance norm
  • Trainable parameters: 31,181,763
  • Depth: 6 encoder levels (256 → 4 bottleneck)
  • Kernels / strides: 3³ at every level; strides [[1,1,1], [2,2,2]×5]
  • Input: (1, 256, 256, 256) full CT resampled to 256³ (no bbox crop), intensity-normalised to [0, 1]
  • Output: (3, 256, 256, 256) softmax logits — class 0 = background, class 1 = lung, class 2 = nodule
  • Framework: PyTorch + MONAI

Nodule is treated as a class distinct from lung: a voxel that is both lung tissue and nodule is assigned exclusively to class 2 (nodule takes precedence over lung).

Data

Trained on the unified_ex_lidc split (patient-grouped, dataset-stratified, LIDC excluded):

Source Role
NLST train + val
NSCLC-Radiomics train + val
LIDC-IDRI excluded

Split sizes: 1 110 train / 196 val / 129 test (held out).

Lung labels are ground-truth per-slice 2D masks (from roi_sem_seg_2d/), stacked into 3D per series. Nodule labels come from the corpus's own 3D nodule annotations.

Validation metrics (best-epoch, val split)

Class Dice Recall Precision
Lung 0.9800 0.984 0.976
Nodule 0.6752 0.746 0.617
Combined (mean) 0.8276 — —

Combined score = 0.5 · (Lung Dice + Nodule Dice). Both class Dices are per-case-averaged on the val split.

Best joint model in this collection — beats the SegResNet ex-LIDC counterpart by 0.010 on nodule Dice at slightly higher recall / lower precision.

How to load & run inference

import yaml, torch
from monai.networks.nets import DynUNet

cfg = yaml.safe_load(open("config.yaml"))["model"]
default_strides = [[1, 1, 1]] + [[2, 2, 2]] * 5
strides = cfg.get("strides", default_strides)
kernel  = cfg.get("kernel_size", [[3, 3, 3]] * len(strides))
model = DynUNet(
    spatial_dims         = cfg["spatial_dims"],
    in_channels          = cfg["in_channels"],
    out_channels         = cfg["out_channels"],           # 3
    kernel_size          = kernel,
    strides              = strides,
    upsample_kernel_size = cfg.get("upsample_kernel_size", strides[1:]),
    norm_name            = cfg.get("norm_name", "instance"),
    deep_supervision     = cfg.get("deep_supervision", False),
)
state = torch.load("model.pth", map_location="cpu", weights_only=True)
model.load_state_dict(state)
model.eval()

with torch.no_grad():
    x = torch.randn(1, 1, 256, 256, 256)             # (B, C, H, W, D)
    logits = model(x)                                # (B, 3, D, H, W)
    pred_class = logits.argmax(dim=1)                # (B, D, H, W) in {0, 1, 2}
    lung_mask   = (pred_class == 1).to(torch.uint8)
    nodule_mask = (pred_class == 2).to(torch.uint8)

Unlike the two-stage nodule-* checkpoints in this collection, this model does not need a lung-bbox crop — feed it the whole CT resampled to 256³.

Training recipe

  • Loss: Multi-class Focal Tversky + weighted CE (α=0.3, β=0.7, γ=2.0, λ_CE=0.1, class weights = [1.0, 1.0, 100.0] for [bg, lung, nodule])
  • Optimizer: Adam (lr = 1e-5, weight decay = 1e-5)
  • Scheduler: CosineAnnealingLR (T_max = 400, η_min = 1e-6)
  • Batch size: 2 (larger model than SegResNet — less VRAM headroom)
  • Epochs: 400
  • Mixed precision: bf16
  • Augmentation: 3D flips, 90° rotations, elastic rotation, zoom, intensity scale/shift, Gaussian noise/blur, contrast
  • Seed: 42
  • Hardware: 1 × NVIDIA H100 94 GB
  • Wall-clock: ≈ 3.5 days

Full config is included in this repo as config.yaml.

Ablation: ex-LIDC vs pseudo-LIDC

The joint task cannot use LIDC directly because LIDC lacks GT lung labels. Two variants were trained:

  • This model (ex-LIDC): drop LIDC entirely → 1 110 train cases, all with full GT.
  • joint-dynunet-3d-pseudo-lidc: include LIDC with lung labels produced by the 2D SegResNet ROI model → 1 683 train cases, mixed GT
    • pseudo.

On the val split, ex-LIDC beats pseudo-LIDC by 0.022 combined Dice (0.8276 vs 0.8055) — a larger gap than the SegResNet pair. The pseudo-labels' noise slightly hurts the lung head's supervision signal and the additional LIDC diversity does not compensate.

License & intended use

Model weights released under Apache 2.0. Training data was public but covered by dataset-specific terms (NLST, NSCLC-Radiomics) — users must comply with those separately when using the model on comparable data.

Not a medical device. Not intended for clinical use. Research only.

Citation

Paper in preparation.

Downloads last month
10
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Collections including HalmosiL/joint-dynunet-3d-ex-lidc