| """ |
| Pulmo — Two-stage explainable lung-nodule analysis pipeline. |
| |
| This module defines BOTH models of the pipeline and the glue code that chains |
| them into a single `volume -> findings` call: |
| |
| Stage 1 (detector) HeatmapUNet3D -> find nodule centres in a full CT volume |
| Stage 2 (characteriser) Student2p5D -> per-candidate diagnosis + explanation |
| |
| Module keys here MUST match the released checkpoints exactly: |
| |
| Stage 1 (`stage1_detector_v2.pth`) |
| e1..e4, bott, u1..u4, d1..d4, out -> 3D U-Net heatmap detector |
| |
| Stage 2 (`student_2p5d_best.pth`) |
| cnn.* -> 2D U-Net backbone (shared trunk) |
| detection_head.* -> binary nodule / non-nodule |
| concept_head.* -> 8 LIDC radiological concepts (regression) |
| malignancy_head.*-> Linear(8 -> 2) (the concept bottleneck) |
| cnn.final.* -> segmentation logits of the middle slice |
| |
| Only `torch`, `numpy`, and `scipy` are required. |
| """ |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from scipy.ndimage import maximum_filter |
|
|
| CONCEPT_NAMES = [ |
| "subtlety", "internalStructure", "calcification", "sphericity", |
| "margin", "lobulation", "spiculation", "texture", |
| ] |
|
|
| |
| HU_CLIP = (-1000, 1000) |
| ROI = 64 |
| N_SLICES = 7 |
| STAGE1_PATCH = (64, 128, 128) |
|
|
|
|
| def normalize_hu(x): |
| """Clip to HU_CLIP and scale to [0, 1] (same for both stages).""" |
| x = np.clip(x.astype(np.float32), HU_CLIP[0], HU_CLIP[1]) |
| return (x - HU_CLIP[0]) / (HU_CLIP[1] - HU_CLIP[0]) |
|
|
|
|
| |
| |
| |
| class ResBlock2D(nn.Module): |
| def __init__(self, i, o): |
| super().__init__() |
| self.conv1 = nn.Conv2d(i, o, 3, padding=1, bias=False) |
| self.norm1 = nn.InstanceNorm2d(o) |
| self.conv2 = nn.Conv2d(o, o, 3, padding=1, bias=False) |
| self.norm2 = nn.InstanceNorm2d(o) |
| self.act = nn.LeakyReLU(0.1, inplace=True) |
| self.skip = nn.Conv2d(i, o, 1, bias=False) if i != o else nn.Identity() |
|
|
| def forward(self, x): |
| idt = self.skip(x) |
| out = self.act(self.norm1(self.conv1(x))) |
| out = self.norm2(self.conv2(out)) |
| return self.act(out + idt) |
|
|
|
|
| class UNet2D(nn.Module): |
| def __init__(self, in_channels, base=24): |
| super().__init__() |
| self.stem = ResBlock2D(in_channels, base) |
| self.down1 = nn.Sequential(nn.MaxPool2d(2), ResBlock2D(base, base * 2)) |
| self.down2 = nn.Sequential(nn.MaxPool2d(2), ResBlock2D(base * 2, base * 4)) |
| self.down3 = nn.Sequential(nn.MaxPool2d(2), ResBlock2D(base * 4, base * 8)) |
| self.bottom = nn.Sequential(nn.MaxPool2d(2), ResBlock2D(base * 8, base * 16)) |
| self.global_pool = nn.AdaptiveAvgPool2d(1) |
| self.up4 = nn.ConvTranspose2d(base * 16, base * 8, 2, 2) |
| self.dec4 = ResBlock2D(base * 16, base * 8) |
| self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, 2) |
| self.dec3 = ResBlock2D(base * 8, base * 4) |
| self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, 2) |
| self.dec2 = ResBlock2D(base * 4, base * 2) |
| self.up1 = nn.ConvTranspose2d(base * 2, base, 2, 2) |
| self.dec1 = ResBlock2D(base * 2, base) |
| self.final = nn.Conv2d(base, 1, 1) |
| self.out_dim = base * 16 |
|
|
| def forward(self, x): |
| s0 = self.stem(x) |
| s1 = self.down1(s0) |
| s2 = self.down2(s1) |
| s3 = self.down3(s2) |
| b = self.bottom(s3) |
| gf = self.global_pool(b).flatten(1) |
| u4 = self.up4(b); d4 = self.dec4(torch.cat([u4, s3], 1)) |
| u3 = self.up3(d4); d3 = self.dec3(torch.cat([u3, s2], 1)) |
| u2 = self.up2(d3); d2 = self.dec2(torch.cat([u2, s1], 1)) |
| u1 = self.up1(d2); d1 = self.dec1(torch.cat([u1, s0], 1)) |
| return gf, self.final(d1) |
|
|
|
|
| class Student2p5D(nn.Module): |
| """Stage 2: 2.5D concept-bottleneck multi-task characteriser.""" |
|
|
| def __init__(self, n_slices=7, n_concepts=8, base=24, head_dropout=0.1): |
| super().__init__() |
| self.n_slices = n_slices |
| self.n_concepts = n_concepts |
| self.cnn = UNet2D(n_slices, base=base) |
| cd = self.cnn.out_dim |
| self.detection_head = nn.Sequential( |
| nn.LayerNorm(cd), nn.Linear(cd, 256), nn.GELU(), |
| nn.Dropout(head_dropout), nn.Linear(256, 2), |
| ) |
| self.concept_head = nn.Sequential( |
| nn.LayerNorm(cd), nn.Linear(cd, 256), nn.GELU(), |
| nn.Dropout(0.3), nn.Linear(256, n_concepts), |
| ) |
| |
| self.malignancy_head = nn.Linear(n_concepts, 2) |
|
|
| def forward(self, x): |
| gf, seg = self.cnn(x) |
| concepts = self.concept_head(gf) |
| return { |
| "detection": self.detection_head(gf), |
| "concepts": concepts, |
| "malignancy": self.malignancy_head(concepts), |
| "segmentation": seg, |
| } |
|
|
|
|
| |
| |
| |
| class CB3(nn.Module): |
| def __init__(self, i, o): |
| super().__init__() |
| self.c1 = nn.Conv3d(i, o, 3, padding=1, bias=False) |
| self.n1 = nn.InstanceNorm3d(o) |
| self.c2 = nn.Conv3d(o, o, 3, padding=1, bias=False) |
| self.n2 = nn.InstanceNorm3d(o) |
| self.a = nn.LeakyReLU(0.1, inplace=True) |
|
|
| def forward(self, x): |
| return self.a(self.n2(self.c2(self.a(self.n1(self.c1(x)))))) |
|
|
|
|
| class HeatmapUNet3D(nn.Module): |
| """Stage 1: 3D U-Net that outputs a nodule-centre probability heatmap.""" |
|
|
| def __init__(self, base=16): |
| super().__init__() |
| self.e1 = CB3(1, base); self.e2 = CB3(base, base * 2) |
| self.e3 = CB3(base * 2, base * 4); self.e4 = CB3(base * 4, base * 8) |
| self.pool = nn.MaxPool3d(2); self.bott = CB3(base * 8, base * 16) |
| self.u4 = nn.ConvTranspose3d(base * 16, base * 8, 2, 2); self.d4 = CB3(base * 16, base * 8) |
| self.u3 = nn.ConvTranspose3d(base * 8, base * 4, 2, 2); self.d3 = CB3(base * 8, base * 4) |
| self.u2 = nn.ConvTranspose3d(base * 4, base * 2, 2, 2); self.d2 = CB3(base * 4, base * 2) |
| self.u1 = nn.ConvTranspose3d(base * 2, base, 2, 2); self.d1 = CB3(base * 2, base) |
| self.out = nn.Conv3d(base, 1, 1) |
|
|
| def forward(self, x): |
| e1 = self.e1(x); e2 = self.e2(self.pool(e1)) |
| e3 = self.e3(self.pool(e2)); e4 = self.e4(self.pool(e3)) |
| b = self.bott(self.pool(e4)) |
| d = self.d4(torch.cat([self.u4(b), e4], 1)); d = self.d3(torch.cat([self.u3(d), e3], 1)) |
| d = self.d2(torch.cat([self.u2(d), e2], 1)); d = self.d1(torch.cat([self.u1(d), e1], 1)) |
| return self.out(d) |
|
|
|
|
| |
| |
| |
| def load_stage1(ckpt_path, device="cpu"): |
| """Load the Stage-1 detector (HeatmapUNet3D).""" |
| ck = torch.load(ckpt_path, map_location=device, weights_only=False) |
| base = ck.get("base", 16) if isinstance(ck, dict) else 16 |
| model = HeatmapUNet3D(base=base).to(device) |
| state = ck["model_state_dict"] if isinstance(ck, dict) and "model_state_dict" in ck else ck |
| model.load_state_dict(state, strict=True) |
| model.eval() |
| return model |
|
|
|
|
| def load_stage2(ckpt_path, device="cpu", n_slices=7, n_concepts=8, base=24): |
| """Load the Stage-2 characteriser (Student2p5D).""" |
| model = Student2p5D(n_slices=n_slices, n_concepts=n_concepts, base=base).to(device) |
| ck = torch.load(ckpt_path, map_location=device, weights_only=False) |
| state = ck["model_state_dict"] if isinstance(ck, dict) and "model_state_dict" in ck else ck |
| model.load_state_dict(state, strict=True) |
| model.eval() |
| return model |
|
|
|
|
| |
| |
| |
| @torch.no_grad() |
| def _heatmap_volume(model, volume, device, stride=(32, 64, 64), batch=4): |
| """Run the 3D detector over a full volume with an overlapping sliding window. |
| Returns a per-voxel nodule-centre probability heatmap, same shape as `volume`. |
| """ |
| PZ, PY, PX = STAGE1_PATCH |
| Z0, Y0, X0 = volume.shape |
| Zp, Yp, Xp = max(Z0, PZ), max(Y0, PY), max(X0, PX) |
| if (Zp, Yp, Xp) != (Z0, Y0, X0): |
| v2 = np.full((Zp, Yp, Xp), HU_CLIP[0], dtype=volume.dtype) |
| v2[:Z0, :Y0, :X0] = volume |
| volume = v2 |
| Z, Y, X = volume.shape |
| acc = np.zeros((Z, Y, X), np.float32); cnt = np.zeros((Z, Y, X), np.float32) |
| sz, sy, sx = stride |
| zs = list(range(0, Z - PZ + 1, sz)) + ([Z - PZ] if (Z - PZ) % sz else []) |
| ys = list(range(0, Y - PY + 1, sy)) + ([Y - PY] if (Y - PY) % sy else []) |
| xs = list(range(0, X - PX + 1, sx)) + ([X - PX] if (X - PX) % sx else []) |
| coords = [(z, y, x) for z in zs for y in ys for x in xs] |
| buf, pos = [], [] |
| use_amp = (str(device) != "cpu") |
|
|
| def flush(): |
| if not buf: |
| return |
| xb = torch.from_numpy(np.stack(buf)[:, None]).to(device) |
| if use_amp: |
| with torch.autocast("cuda", dtype=torch.bfloat16): |
| hm = torch.sigmoid(model(xb).float()).cpu().numpy() |
| else: |
| hm = torch.sigmoid(model(xb)).cpu().numpy() |
| for k, (z, y, x) in enumerate(pos): |
| acc[z:z + PZ, y:y + PY, x:x + PX] += hm[k, 0] |
| cnt[z:z + PZ, y:y + PY, x:x + PX] += 1 |
| buf.clear(); pos.clear() |
|
|
| for (z, y, x) in coords: |
| p = np.clip(volume[z:z + PZ, y:y + PY, x:x + PX].astype(np.float32), HU_CLIP[0], HU_CLIP[1]) |
| if p.mean() < HU_CLIP[0] + 20: |
| continue |
| buf.append((p - HU_CLIP[0]) / (HU_CLIP[1] - HU_CLIP[0])) |
| pos.append((z, y, x)) |
| if len(buf) >= batch: |
| flush() |
| flush() |
| cnt[cnt == 0] = 1 |
| return (acc / cnt)[:Z0, :Y0, :X0] |
|
|
|
|
| def find_candidates(model, volume, spacing, device="cpu", |
| peak_thresh=0.1, cluster_mm=6.0, return_scores=False): |
| """Stage 1: locate nodule centres in a full CT volume. |
| |
| Args: |
| model : a loaded HeatmapUNet3D (Stage 1). |
| volume : (Z, Y, X) numpy array of raw HU values. |
| spacing : (sz, sy, sx) voxel spacing in mm, [z, y, x] order. |
| peak_thresh : heatmap threshold. Lower -> higher recall, more false |
| positives (Stage 2 then filters them). 0.1 is a balanced |
| default; ~0.05 maximises recall. |
| cluster_mm : merge peaks closer than this (mm). |
| |
| Returns: |
| list of (z, y, x) voxel coordinates, highest-scoring first. |
| If return_scores, returns (coords, scores). |
| """ |
| hm = _heatmap_volume(model, volume, device) |
| mx = maximum_filter(hm, size=(5, 9, 9)) |
| peaks = np.argwhere((hm == mx) & (hm >= peak_thresh)) |
| if len(peaks) == 0: |
| return ([], []) if return_scores else [] |
| sc = hm[peaks[:, 0], peaks[:, 1], peaks[:, 2]] |
| order = np.argsort(-sc); peaks = peaks[order]; sc = sc[order] |
| sz, sy, sx = spacing |
| mm = peaks * np.array([sz, sy, sx]) |
| taken = np.zeros(len(peaks), bool); out = []; outs = [] |
| for i in range(len(peaks)): |
| if taken[i]: |
| continue |
| out.append(tuple(int(v) for v in peaks[i])); outs.append(float(sc[i])) |
| taken |= (np.linalg.norm(mm - mm[i], axis=1) < cluster_mm) |
| return (out, outs) if return_scores else out |
|
|
|
|
| |
| |
| |
| def crop_stage2_input(volume, center_zyx): |
| """Crop a 64^3 patch centred on a candidate and return the Stage-2 input |
| tensor `(1, 7, 64, 64)` (7 central axial slices, normalized to [0, 1]). |
| Out-of-bounds regions are padded with air (HU_CLIP[0]). |
| """ |
| z, y, x = (int(round(c)) for c in center_zyx) |
| Z, Y, X = volume.shape |
| h = ROI // 2 |
| patch = np.full((ROI, ROI, ROI), HU_CLIP[0], dtype=np.float32) |
| z0, y0, x0 = z - h, y - h, x - h |
| zz0, zz1 = max(0, z0), min(Z, z0 + ROI) |
| yy0, yy1 = max(0, y0), min(Y, y0 + ROI) |
| xx0, xx1 = max(0, x0), min(X, x0 + ROI) |
| patch[zz0 - z0:zz1 - z0, yy0 - y0:yy1 - y0, xx0 - x0:xx1 - x0] = \ |
| volume[zz0:zz1, yy0:yy1, xx0:xx1] |
| patch = normalize_hu(patch) |
| c = ROI // 2; hs = N_SLICES // 2 |
| slices = patch[c - hs:c + hs + 1] |
| return torch.from_numpy(slices[None]).float() |
|
|
|
|
| def explain_malignancy(stage2_model, out): |
| """Concept-bottleneck attribution for one Stage-2 output. |
| |
| malignancy logit = sum_i w_net[i] * concept[i], w_net = W[malign] - W[benign]. |
| Returns a list of (concept_name, concept_value, contribution) sorted by |
| contribution (most malignancy-driving first). |
| """ |
| concepts = out["concepts"][0].detach().cpu().numpy() |
| W = stage2_model.malignancy_head.weight.detach().cpu().numpy() |
| w_net = W[1] - W[0] |
| contrib = w_net * concepts |
| order = np.argsort(contrib)[::-1] |
| return [(CONCEPT_NAMES[i], float(concepts[i]), float(contrib[i])) for i in order] |
|
|
|
|
| if __name__ == "__main__": |
| s1 = HeatmapUNet3D(base=16) |
| s2 = Student2p5D() |
| print("Stage 1 HeatmapUNet3D : %.2fM params" % (sum(p.numel() for p in s1.parameters()) / 1e6)) |
| print("Stage 2 Student2p5D : %.2fM params" % (sum(p.numel() for p in s2.parameters()) / 1e6)) |
| out = s2(torch.randn(2, 7, 64, 64)) |
| for k, v in out.items(): |
| print(f" stage2 {k:13s}: {tuple(v.shape)}") |
| hm = s1(torch.randn(1, 1, 64, 128, 128)) |
| print(f" stage1 heatmap : {tuple(hm.shape)}") |
|
|