Introduce the third segmentation pipeline using MONAI (1.6.x) to allow direct comparison with Pipelines A and B. This includes the implementation of the iterative pseudo-labeling workflow, training scripts, and inference protocols. - Add `scripts_monai/` directory containing the MONAI pipeline scripts. - Update documentation in `README.md` and `AGENTS.md` to include MONAI package requirements and pipeline details. - Configure `.gitignore` to exclude MONAI-specific run directories. - Update data directory descriptions to include MONAI pseudo-labels.
208 lines
No EOL
8.1 KiB
Python
208 lines
No EOL
8.1 KiB
Python
"""Shared helpers for the MONAI iterative pseudo-labeling pipeline (scripts_monai/).
|
||
|
||
Backbone: MONAI 1.6.x (pip) — monai.networks.nets.UNet, MONAI transforms /
|
||
Dataset / DataLoader, and MONAI sliding-window inference. Everything
|
||
study-specific (splits, gates, loss weighting, longitudinal consistency) is
|
||
shared with the other pipelines so round results stay directly comparable:
|
||
|
||
* per-sample weighted Dice+CE loss — Pipeline A's convention (row "w";
|
||
pseudo-label rows down-weighted sample-by-sample)
|
||
* selection gates + head-relative consistency filter — Pipeline B's
|
||
implementation (scripts_nnu/nnu_common; absolute patient-space resampling
|
||
is unreliable across acquisitions, see README)
|
||
* DDP via torchrun for training, per-rank row sharding for inference —
|
||
Pipeline A's parallelism model
|
||
|
||
Network: monai UNet 3D, 1 in / 2 out, channels 16→128 (~1.2M params) — the
|
||
MONAI analogue of Pipeline A's Unet3D(base=16, depth=4).
|
||
"""
|
||
import os
|
||
import sys
|
||
|
||
# Data/run dirs: overridable for scratch runs. Code (src/, scripts_nnu/): next to this file.
|
||
ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal")
|
||
_REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||
sys.path.insert(0, _REPO)
|
||
sys.path.insert(0, os.path.join(_REPO, "scripts_nnu"))
|
||
|
||
import numpy as np
|
||
import torch
|
||
import torch.nn as nn
|
||
|
||
from src.common import dice # noqa: E402
|
||
|
||
# Selection gates / consistency filter shared with Pipeline B (pure numpy/scipy)
|
||
from nnu_common import pos_mask, neg_frac_bg, consistency_filter, load_voxel_stats # noqa: E402,F401
|
||
|
||
WIN = 96
|
||
|
||
|
||
# ---------------- network / loss ----------------
|
||
|
||
def build_model(device):
|
||
from monai.networks.nets import UNet
|
||
net = UNet(
|
||
spatial_dims=3,
|
||
in_channels=1,
|
||
out_channels=2,
|
||
channels=(16, 32, 64, 128),
|
||
strides=(2, 2, 2),
|
||
num_res_units=2,
|
||
act=("LEAKYRELU", {"inplace": True, "negative_slope": 0.01}),
|
||
norm="batch",
|
||
)
|
||
return net.to(device)
|
||
|
||
|
||
class WeightedDiceCELoss(nn.Module):
|
||
"""Per-sample (soft-dice + CE), weighted batch mean.
|
||
|
||
Same convention as Pipeline A's src/losses.py per_sample_loss (dice over
|
||
background+foreground averaged, dice_weight 0.5), built on the MONAI
|
||
DiceLoss. Per-sample reduction is required so accepted pseudo-label rows
|
||
can be down-weighted individually via row "w"; a batch-level loss
|
||
(MONAI DiceCELoss / nnU-Net CE+Dice) cannot express that.
|
||
"""
|
||
|
||
def __init__(self, dice_weight=0.5):
|
||
super().__init__()
|
||
from monai.losses import DiceLoss
|
||
self.dice_weight = dice_weight
|
||
self.dice = DiceLoss(include_background=True, softmax=False, reduction="none")
|
||
|
||
def forward(self, logits, label, w):
|
||
logits = logits.float()
|
||
c = logits.size(1)
|
||
one = nn.functional.one_hot(label, c).permute(0, -1, *range(1, label.ndim)).float()
|
||
p = nn.functional.softmax(logits, dim=1)
|
||
dsc = self.dice(p, one) # (B, C, 1, 1, 1)
|
||
dsc = dsc.squeeze(-1).squeeze(-1).squeeze(-1).mean(dim=1) # (B,)
|
||
ce = nn.functional.cross_entropy(logits, label, reduction="none") # (B, z, y, x)
|
||
ce = ce.mean(dim=tuple(range(1, ce.ndim))) # (B,)
|
||
per = (1 - self.dice_weight) * ce + self.dice_weight * dsc
|
||
if torch.isnan(per).any() or torch.isinf(per).any():
|
||
per = torch.zeros_like(per) # zero-gradient fallback: keeps DDP collectives in sync
|
||
return (per * w).sum() / w.sum().clamp(min=1e-6)
|
||
|
||
|
||
# ---------------- data ----------------
|
||
|
||
def train_transform(win=WIN):
|
||
"""Channel-first MONAI chain: load → min-pad → augment → pos/neg crop → typed tensors.
|
||
|
||
Mirrors Pipeline A's augmentation set (per-axis flip p=0.5, rot90 in the
|
||
(y,x) plane p=0.1, brightness 1±0.1 p=0.3, Gaussian noise σ=0.01 p=0.4).
|
||
Patch sampling uses MONAI's foreground-aware RandCropByPosNegLabeld
|
||
(pos=neg=1) instead of Pipeline A's uniform random crop.
|
||
"""
|
||
from monai.transforms import (
|
||
Compose, LoadImaged, EnsureChannelFirstd, EnsureTyped, SpatialPadd,
|
||
RandFlipd, RandRotate90d, RandScaleIntensityd, RandGaussianNoised,
|
||
RandCropByPosNegLabeld)
|
||
return Compose([
|
||
LoadImaged(keys=["pimg", "plabel"]),
|
||
EnsureChannelFirstd(keys=["pimg", "plabel"]),
|
||
SpatialPadd(keys=["pimg", "plabel"], spatial_size=(win, win, win)),
|
||
RandFlipd(keys=["pimg", "plabel"], prob=0.5, spatial_axis=0),
|
||
RandFlipd(keys=["pimg", "plabel"], prob=0.5, spatial_axis=1),
|
||
RandFlipd(keys=["pimg", "plabel"], prob=0.5, spatial_axis=2),
|
||
RandRotate90d(keys=["pimg", "plabel"], prob=0.1, max_k=3, spatial_axes=(1, 2)),
|
||
RandScaleIntensityd(keys=["pimg"], factors=0.1, prob=0.3),
|
||
RandGaussianNoised(keys=["pimg"], std=0.01, prob=0.4),
|
||
RandCropByPosNegLabeld(keys=["pimg", "plabel"], label_key="plabel",
|
||
spatial_size=(win, win, win), pos=1, neg=1, num_samples=1),
|
||
EnsureTyped(keys=["pimg", "plabel"], dtype=[torch.float32, torch.long]),
|
||
])
|
||
|
||
|
||
def rows_for_dataset(rows):
|
||
out = []
|
||
for r in rows:
|
||
out.append({"key": r["key"], "pimg": r["pimg"], "plabel": r["plabel"],
|
||
"w": float(r.get("w", 1.0))})
|
||
return out
|
||
|
||
|
||
def _collate(batch):
|
||
# RandCropByPosNegLabeld(num_samples=1) yields a 1-item list per sample
|
||
items = [b[0] if isinstance(b, (list, tuple)) else b for b in batch]
|
||
return torch.utils.data.default_collate(items)
|
||
|
||
|
||
def make_dataloader(rows, win=WIN, batch=3, workers=4):
|
||
from monai.data import Dataset
|
||
ds = Dataset(data=rows_for_dataset(rows), transform=train_transform(win))
|
||
return torch.utils.data.DataLoader(
|
||
ds, batch_size=batch, shuffle=True, num_workers=workers,
|
||
collate_fn=_collate, drop_last=len(ds) > batch, pin_memory=True,
|
||
persistent_workers=workers > 0)
|
||
|
||
|
||
# ---------------- inference ----------------
|
||
|
||
@torch.no_grad()
|
||
def predict_probs(model, vol, device, win=WIN, overlap=0.5, tta=True, sw_batch=8):
|
||
"""Tumor probability map (C, z, y, x → (z, y, x) numpy) for one volume.
|
||
|
||
MONAI sliding_window_inference (gaussian blend, `overlap`), plus the same
|
||
4-view flip TTA (identity + 3 axis flips) as Pipeline A.
|
||
"""
|
||
from monai.inferers import sliding_window_inference
|
||
a = np.nan_to_num(np.asarray(vol, dtype=np.float32), nan=0.0, posinf=1.5, neginf=0.0)
|
||
a = np.clip(a, 0.0, 1.5)
|
||
t = torch.from_numpy(a).unsqueeze(0).unsqueeze(0).to(device)
|
||
|
||
def run(vt):
|
||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||
logits = sliding_window_inference(
|
||
vt, roi_size=(win, win, win), sw_batch_size=sw_batch,
|
||
predictor=model, mode="gaussian", overlap=overlap,
|
||
sw_device=device, device=device)
|
||
return torch.softmax(logits.float(), dim=1)[0, 1].cpu().numpy()
|
||
|
||
acc = run(t)
|
||
if not tta:
|
||
return acc
|
||
for ax in (0, 1, 2):
|
||
acc += np.flip(run(t.flip(ax + 2)), axis=ax)
|
||
return acc / 4.0
|
||
|
||
|
||
def prob_dice(probs, lab, thr=0.5):
|
||
"""Dice at probability threshold 0.5 — the convention of Pipeline A's
|
||
08_eval and Pipeline B's primary metric."""
|
||
return dice((probs >= thr).astype("uint8"), (lab > 0).astype("uint8"))
|
||
|
||
|
||
def load_model(ckpt, device):
|
||
sd = torch.load(ckpt, map_location=device, weights_only=True)
|
||
net = build_model(device)
|
||
net.load_state_dict(sd.get("model", sd))
|
||
net.eval()
|
||
return net, sd
|
||
|
||
|
||
# ---------------- distributed (torchrun, as in Pipeline A) ----------------
|
||
|
||
def rank_info():
|
||
return (int(os.environ.get("RANK", 0)), int(os.environ.get("WORLD_SIZE", 1)),
|
||
int(os.environ.get("LOCAL_RANK", 0)))
|
||
|
||
|
||
def init_dist():
|
||
if int(os.environ.get("WORLD_SIZE", 1)) > 1:
|
||
from datetime import timedelta
|
||
import torch.distributed as dist
|
||
dist.init_process_group("nccl", timeout=timedelta(minutes=30))
|
||
|
||
|
||
def barrier():
|
||
if int(os.environ.get("WORLD_SIZE", 1)) > 1:
|
||
import torch.distributed as dist
|
||
dist.barrier()
|
||
|
||
|
||
def destroy_dist():
|
||
if int(os.environ.get("WORLD_SIZE", 1)) > 1:
|
||
import torch.distributed as dist
|
||
dist.destroy_process_group() |