jdmayfield commited on
Commit
dcc1df6
Β·
verified Β·
1 Parent(s): b36bf80

Create AAM v3 extraction

Browse files

Add AAM feature extraction pipeline (aam_v3_extraction.py)

Computes three segmentation-free asymmetry channels per case (airway
lesion, nodal/soft-tissue asymmetry, necrosis) via spine-anchored
midline detection and mirrored left-right comparison, producing a QC
plot, a 4D heatmap NIfTI, and a rich per-case feature CSV in a single
parallelized, resumable pass. Used in the current AsymmetryNet
manuscript for HPV status prediction in oropharyngeal squamous cell
carcinoma.

All paths are generic placeholders; no institutional data, patient
identifiers, or working-environment paths are included.

Files changed (1) hide show
  1. AAM v3 extraction +619 -0
AAM v3 extraction ADDED
@@ -0,0 +1,619 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ AsymmetryNet: Asymmetry Attention Mechanism (AAM) β€” segmentation-free
3
+ feature extraction for radiologist-inspired left-right asymmetry in
4
+ head and neck CT, developed for HPV status prediction in oropharyngeal
5
+ squamous cell carcinoma (OPSCC).
6
+
7
+ Given a cropped neck CT volume (NIfTI), this pipeline computes three
8
+ per-slice asymmetry channels across the full volume in a single pass:
9
+
10
+ Ch1 (airway lesion): seeded region growing from the airway lumen,
11
+ bounded by Sobel-detected soft-tissue edges,
12
+ constrained to the side of airway deviation
13
+ from a spine-anchored midline.
14
+ Ch2 (nodal/soft-tissue asymmetry): mirrored left-right comparison of
15
+ fat-plane and soft-tissue density about the
16
+ same midline.
17
+ Ch3 (necrosis): per-case, histogram-derived HU-windowed fluid
18
+ detection, restricted to the union of the
19
+ Ch1 and Ch2 regions.
20
+
21
+ Midline is detected per case via spine-anchored bone-mass/compactness
22
+ scoring, not image center or a fixed landmark.
23
+
24
+ Outputs per case: a QC plot (5 representative slices), a 4D heatmap
25
+ NIfTI (D x H x W x 3, channels = [airway_lesion, nodal, necrosis]) for
26
+ downstream radiomic feature extraction restricted to these regions,
27
+ and a row of per-case rich features appended to a CSV.
28
+
29
+ Parallelized via ProcessPoolExecutor and resumable: interrupting and
30
+ restarting skips cases already written to the output CSV.
31
+ """
32
+
33
+ import numpy as np
34
+ import nibabel as nib
35
+ import pandas as pd
36
+ import matplotlib
37
+ matplotlib.use('Agg') # required for plotting inside worker subprocesses
38
+ import matplotlib.pyplot as plt
39
+ from pathlib import Path
40
+ from scipy import ndimage as ndi
41
+ from scipy.signal import find_peaks
42
+ import os, time, csv
43
+ from concurrent.futures import ProcessPoolExecutor, as_completed
44
+
45
+ # ── Paths & Output ───────────────────────────────────────────────────
46
+ # >>> CHANGE THESE PATHS for your own environment before running <<<
47
+ BASE = Path("/path/to/project")
48
+ CROP_DIR = BASE / "crops" # expects {CROP_DIR}/{dataset}/{case_id}/crop224.nii.gz
49
+ MASTER = BASE / "master_index.csv" # expects columns: case_id, dataset, extracted, [hpv_norm]
50
+
51
+ OUTPUT_DIR = BASE / "AAM_results"
52
+ PLOTS_DIR = OUTPUT_DIR / "plots"
53
+ HEATMAP_DIR = OUTPUT_DIR / "heatmaps"
54
+ FEATURES_CSV_PATH = OUTPUT_DIR / "aam_features_rich.csv"
55
+ PLOTS_DIR.mkdir(parents=True, exist_ok=True)
56
+ HEATMAP_DIR.mkdir(parents=True, exist_ok=True)
57
+
58
+ print(f"Results will be saved to: {OUTPUT_DIR}")
59
+
60
+ # ── Constants ──────────────────────────────────────────────────────────
61
+ HU_MIN, HU_MAX = -200.0, 300.0
62
+ # z-gating already happened during crop generation (z_ctr/thick_mm
63
+ # centering + shift/flip/reshape corrections) -- gating further here
64
+ # would discard real oropharyngeal content across the full crop depth
65
+ Z_GATE_FRAC = 0.0
66
+ AIR_THRESH = 0.08
67
+ BONE_THRESH = 0.88
68
+ SOFT_LO = 0.34
69
+ # Fraction-of-width is the right approach after all: it automatically
70
+ # scales with per-case zoom/FOV (a physically-zoomed-in crop makes the
71
+ # same vertebra occupy more pixels of the same W, which a percentage
72
+ # tracks correctly), and vertebral size itself scales somewhat with
73
+ # overall neck/body size (sex, habitus) -- both handled better by a
74
+ # relative measure than a fixed pixel constant.
75
+ SPINE_RADIUS_FRAC = 0.1
76
+ MIN_FLUID_PX = 15
77
+ CENTROID_HEAT_THRESH = 0.3 # consistent with the 0.3 heat-significance convention used throughout
78
+
79
+ PARAMS_CONTRAST = {'sobel_sigma': 0.5, 'sobel_edge_thresh': 0.22, 'grow_hu_tol': 0.10, 'grow_iters': 15, 'nodal_thresh': 0.18}
80
+ PARAMS_NONCONTRAST = {'sobel_sigma': 1.0, 'sobel_edge_thresh': 0.20, 'grow_hu_tol': 0.12, 'grow_iters': 12, 'nodal_thresh': 0.12}
81
+
82
+ CSV_FIELDNAMES = [
83
+ 'case_id', 'dataset', 'hpv', 'is_contrast', 'contrast_conf', 'midline',
84
+ 'n_valid_slices',
85
+ 'max_nodal_ratio', 'mean_nodal_ratio', 'total_nodal_area',
86
+ 'nodal_laterality_bias', 'n_slices_with_nodal',
87
+ 'max_node_short_axis_px', 'bilateral_nodal_slices', 'bilateral_nodal_frac',
88
+ 'mean_node_centroid_y_vs_airway',
89
+ 'max_necrosis_area', 'mean_necrosis_frac', 'max_necrosis_frac',
90
+ 'n_slices_with_necrosis',
91
+ 'max_airway_eff', 'mean_airway_eff', 'n_slices_with_airway_eff',
92
+ 'ch1_max_area_px', 'ch1_mean_area_px', 'ch1_max_extent_px', 'n_slices_with_ch1',
93
+ 'ch1_airway_z', 'ch1_airway_y', 'ch1_airway_x', 'ch1_airway_side',
94
+ 'ch1_airway_volume_vox', 'ch1_airway_peak_intensity',
95
+ 'ch2_nodal_z', 'ch2_nodal_y', 'ch2_nodal_x', 'ch2_nodal_side',
96
+ 'ch2_nodal_volume_vox', 'ch2_nodal_peak_intensity',
97
+ 'plot_path', 'heatmap_path',
98
+ ]
99
+
100
+ # ── Helper functions (identical detection logic throughout) ───────────
101
+
102
+ def hu_to_norm(hu):
103
+ return (hu - HU_MIN) / (HU_MAX - HU_MIN)
104
+
105
+ def detect_contrast(vol, z_gate):
106
+ oropharynx = vol[z_gate:]
107
+ soft = oropharynx[(oropharynx > 0.15) & (oropharynx < BONE_THRESH)]
108
+ if soft.size == 0: return True, 0.5
109
+ enhance_frac = (soft > 0.65).sum() / soft.size
110
+ is_contrast = enhance_frac > 0.03
111
+ confidence = min(abs(enhance_frac - 0.03) / 0.03, 1.0)
112
+ return is_contrast, confidence
113
+
114
+ def get_params(is_contrast):
115
+ return PARAMS_CONTRAST if is_contrast else PARAMS_NONCONTRAST
116
+
117
+ def find_spine_center(sl, H, W):
118
+ bone = ndi.binary_fill_holes(ndi.binary_closing(sl >= BONE_THRESH, iterations=2))
119
+ lower_half = sl[H//2:, :]
120
+ is_air_heavy = (lower_half < 0.1).mean() > 0.6
121
+ start_y = H // 2
122
+ best_y, best_x = int(0.75*H), W//2
123
+ best_score = -1.0
124
+ directions = [range(start_y, H-20)]
125
+ if is_air_heavy:
126
+ directions.append(range(start_y, max(30, int(H*0.35)), -1))
127
+ for y_range in directions:
128
+ for y_start in y_range:
129
+ posterior = np.zeros((H, W), bool)
130
+ posterior[y_start:, :] = True
131
+ spine_bone = bone & posterior
132
+ if spine_bone.sum() < 25: continue
133
+ ys, xs = np.where(spine_bone)
134
+ cy, cx = int(ys.mean()), int(xs.mean())
135
+ local = spine_bone[max(0,cy-35):cy+35, max(0,cx-30):cx+30]
136
+ if local.size == 0: continue
137
+ compactness = local.sum() / local.size
138
+ score = sl[spine_bone].mean() * compactness * np.sqrt(spine_bone.sum())
139
+ if score > best_score:
140
+ best_score = score; best_y, best_x = cy, cx
141
+ return best_y, best_x
142
+
143
+ def body_mask(sl):
144
+ return ndi.binary_erosion(ndi.binary_fill_holes(sl >= AIR_THRESH), iterations=2)
145
+
146
+ def bone_mask(sl):
147
+ bone = ndi.binary_closing(sl >= BONE_THRESH, iterations=2)
148
+ return ndi.binary_dilation(ndi.binary_fill_holes(bone), iterations=3)
149
+
150
+ def airway_mask(sl, H, W):
151
+ rl, rh = int(0.15*H), int(0.72*H)
152
+ lbl, nf = ndi.label(sl < AIR_THRESH)
153
+ edge = set(lbl[0,:])|set(lbl[-1,:])|set(lbl[:,0])|set(lbl[:,-1])
154
+ aw = np.zeros(sl.shape, bool)
155
+ for i in range(1, nf+1):
156
+ if i in edge: continue
157
+ c = (lbl==i)
158
+ if c.sum() < 20: continue
159
+ ys, xs = np.where(c)
160
+ if ys.mean()<rl or ys.mean()>rh: continue
161
+ if (xs.max()-xs.min()+1)/W > 0.45: continue
162
+ aw |= c
163
+ return aw
164
+
165
+ def global_midline(vol, z_gate):
166
+ D, H, W = vol.shape
167
+ cx_list = []
168
+ for z in range(z_gate, D):
169
+ sl = vol[z]
170
+ if is_skull_base(sl): continue
171
+ _, spx = find_spine_center(sl, H, W)
172
+ cx_list.append(spx)
173
+ return int(np.median(cx_list)) if cx_list else W//2
174
+
175
+ def sobel_clean(sl, body, bone_m, spine_m, sigma=0.8):
176
+ clean = sl * body * (~bone_m) * (~spine_m)
177
+ clean = np.clip(clean, 0, hu_to_norm(80))
178
+ sm = ndi.gaussian_filter(clean, sigma)
179
+ gy = ndi.sobel(sm, axis=0); gx = ndi.sobel(sm, axis=1)
180
+ mag = np.sqrt(gy**2 + gx**2)
181
+ mag *= body * (~bone_m) * (~spine_m)
182
+ return mag
183
+
184
+ def mirror_about(arr, midline, W):
185
+ cw = min(midline, W-midline)
186
+ if cw < 5: return arr.copy()
187
+ out = np.zeros_like(arr)
188
+ left = arr[:, midline-cw:midline]; right = arr[:, midline:midline+cw]
189
+ out[:, midline-cw:midline] = np.flip(right, axis=1)
190
+ out[:, midline:midline+cw] = np.flip(left, axis=1)
191
+ return out
192
+
193
+ def anterior_mask(H, W, spine_y):
194
+ mask = np.zeros((H, W), bool)
195
+ mask[:spine_y-2, :] = True
196
+ return mask
197
+
198
+ def spine_cylinder_mask(H, W, spine_y, spine_x, radius):
199
+ yy, xx = np.ogrid[:H, :W]
200
+ return np.sqrt((yy-spine_y)**2 + (xx-spine_x)**2) < radius
201
+
202
+ def is_skull_base(sl, thresh=0.25):
203
+ return (sl >= BONE_THRESH).sum() / sl.size > thresh
204
+
205
+ def ch1_airway_lesion(sl, sobel_map, midline, work_mask, H, W, params):
206
+ aw = airway_mask(sl, H, W)
207
+ if aw.sum() < 20:
208
+ return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
209
+ aw_cx = np.where(aw)[1].mean()
210
+ shift_px = aw_cx - midline
211
+ if abs(shift_px) < 2:
212
+ lesion_side = 'left' if aw[:,:midline].sum() < aw[:,midline:].sum() else 'right'
213
+ else:
214
+ lesion_side = 'left' if shift_px > 0 else 'right'
215
+ aw_mirror = mirror_about(aw.astype(np.float32), midline, W) > 0.5
216
+ tissue = (sl >= SOFT_LO) & work_mask
217
+ seed = tissue & aw_mirror & (~aw)
218
+ if lesion_side == 'left': seed[:, midline:] = False
219
+ else: seed[:, :midline] = False
220
+ if seed.sum() < 3:
221
+ return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
222
+ sob_norm = sobel_map / (sobel_map.max() + 1e-6)
223
+ barrier = sob_norm > params['sobel_edge_thresh']
224
+ grown = seed.copy()
225
+ for _ in range(params['grow_iters']):
226
+ ref_hu = sl[grown].mean()
227
+ expanded = ndi.binary_dilation(grown, iterations=1)
228
+ candidates = expanded & (~grown) & work_mask & (~barrier)
229
+ candidates &= np.abs(sl - ref_hu) < params['grow_hu_tol']
230
+ if lesion_side == 'left': candidates[:, midline:] = False
231
+ else: candidates[:, :midline] = False
232
+ if candidates.sum() == 0: break
233
+ grown |= candidates
234
+ heat = ndi.gaussian_filter(grown.astype(np.float32), 2.0)
235
+ return (heat/(heat.max()+1e-6) if heat.max()>0 else heat), grown.copy()
236
+
237
+ def ch2_nodal(sl, sobel_map, midline, work_mask, H, W, params):
238
+ body = ndi.binary_erosion(
239
+ ndi.binary_fill_holes(sl >= AIR_THRESH),
240
+ iterations=max(4, int(0.04*max(H,W))))
241
+ bone_int = bone_mask(sl)
242
+ aw = airway_mask(sl, H, W)
243
+ central = ndi.binary_dilation(aw, iterations=10)
244
+ spy, _ = np.where(bone_int)
245
+ spine_y = int(np.percentile(spy, 75)) if len(spy)>0 else int(0.7*H)
246
+ deep_roi = body & (~bone_int) & (~central)
247
+ deep_roi[spine_y:, :] = False
248
+ if deep_roi.sum() < 50:
249
+ return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
250
+ FAT_LO = hu_to_norm(-150); FAT_HI = hu_to_norm(-30)
251
+ soft = (sl >= SOFT_LO) & (sl < BONE_THRESH) & deep_roi
252
+ fat = (sl >= FAT_LO) & (sl <= FAT_HI) & deep_roi
253
+ cw = min(midline, W-midline)
254
+ if cw < 10: return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
255
+ mir_fat = mirror_about(fat.astype(np.float32), midline, W) > 0.5
256
+ A = (soft & mir_fat).astype(np.float32)
257
+ xor = fat ^ (mirror_about(fat.astype(np.float32), midline, W) > 0.5)
258
+ softness = np.clip((sl - SOFT_LO) / (0.64 - SOFT_LO), 0, 1)
259
+ B = xor.astype(np.float32) * softness * soft
260
+ heat = ndi.gaussian_filter(np.maximum(A, B), 4.0) * deep_roi
261
+ mask = (heat > heat.max()*0.3) if heat.max()>0 else np.zeros((H,W), bool)
262
+ return (heat/(heat.max()+1e-6) if heat.max()>0 else heat), mask
263
+
264
+ def compute_fluid_window(vol, z_gate):
265
+ oropharynx = vol[z_gate:]
266
+ body_vox = oropharynx[(oropharynx > 0.05) & (oropharynx < 0.85)]
267
+ counts, edges = np.histogram(body_vox, bins=200)
268
+ centers = 0.5*(edges[:-1]+edges[1:])
269
+ smooth = ndi.gaussian_filter1d(counts.astype(float), sigma=3)
270
+ peaks, _ = find_peaks(smooth, height=smooth.max()*0.05, distance=15)
271
+ soft_cands = [(i, smooth[p]) for i,p in enumerate(peaks) if centers[p]>0.30]
272
+ soft_peak = centers[peaks[max(soft_cands, key=lambda x:x[1])[0]]] if soft_cands else 0.48
273
+ return soft_peak-0.08, soft_peak-0.02, soft_peak
274
+
275
+ def ch3_necrosis(sl, midline, work_mask, H, W, fluid_lo, fluid_hi, m1, m2):
276
+ lesion_roi = ndi.binary_dilation(m1|m2, iterations=1)
277
+ if lesion_roi.sum() < 5: return np.zeros((H,W), np.float32)
278
+ aw_dil = ndi.binary_dilation(airway_mask(sl,H,W), iterations=3)
279
+ fluid = (sl>=fluid_lo)&(sl<=fluid_hi)&work_mask&(~aw_dil)&lesion_roi
280
+ lbl, n = ndi.label(fluid)
281
+ heat = np.zeros((H,W), np.float32)
282
+ for i in range(1,n+1):
283
+ c=(lbl==i)
284
+ if c.sum()>=MIN_FLUID_PX: heat[c]=1.0
285
+ heat = ndi.gaussian_filter(heat, 3.0)
286
+ return heat/(heat.max()+1e-6) if heat.max()>0 else heat
287
+
288
+ def volume_weighted_centroid(heat_vol, midline, thresh=CENTROID_HEAT_THRESH):
289
+ mask = heat_vol > thresh
290
+ n_vox = int(mask.sum())
291
+ if n_vox == 0:
292
+ return {"z": np.nan, "y": np.nan, "x": np.nan,
293
+ "volume_vox": 0, "side": None, "peak_intensity": 0.0}
294
+ zs, ys, xs = np.where(mask)
295
+ weights = heat_vol[mask]
296
+ z_c = float(np.average(zs, weights=weights))
297
+ y_c = float(np.average(ys, weights=weights))
298
+ x_c = float(np.average(xs, weights=weights))
299
+ side = "left" if x_c > midline else "right"
300
+ return {"z": z_c, "y": y_c, "x": x_c,
301
+ "volume_vox": n_vox, "side": side,
302
+ "peak_intensity": float(heat_vol.max())}
303
+
304
+
305
+ # ══════════════════════════════════════════════════════════════════════
306
+ # SINGLE consolidated per-case worker: one pass over all slices produces
307
+ # the plot, the heatmap NIfTI, AND the rich feature row together.
308
+ # ══════════════════════════════════════════════════════════════════════
309
+
310
+ def process_case_full(case_id, dataset, hpv_label=None):
311
+ path = CROP_DIR / dataset / case_id / "crop224.nii.gz"
312
+ if not path.exists():
313
+ return None, "missing_crop"
314
+
315
+ vol = nib.load(str(path)).get_fdata().astype(np.float32)
316
+ if vol.max() > 1.5:
317
+ vol = np.clip(vol, HU_MIN, HU_MAX)
318
+ vol = (vol - HU_MIN) / (HU_MAX - HU_MIN)
319
+ vol = np.clip(vol, 0, 1).astype(np.float32)
320
+
321
+ D, H, W = vol.shape
322
+ z_gate = int(D * Z_GATE_FRAC)
323
+ midline = global_midline(vol, z_gate)
324
+ is_contrast, conf = detect_contrast(vol, z_gate)
325
+ params = get_params(is_contrast)
326
+ fluid_lo, fluid_hi, _ = compute_fluid_window(vol, z_gate)
327
+
328
+ # which slices get cached for the QC plot (5 representative slices)
329
+ plot_zs = set(np.linspace(0, D - 1, 5).round().astype(int).tolist())
330
+ plot_cache = {}
331
+
332
+ # rich-feature accumulators
333
+ nodal_ratios, airway_effs = [], []
334
+ total_nodal = left_total = right_total = 0
335
+ max_necrosis = 0
336
+ n_valid_slices = 0
337
+ max_node_short_axis_px = 0
338
+ necrosis_fracs = []
339
+ node_centroid_ys = []
340
+ bilateral_slices = 0
341
+ ch1_areas = []
342
+ ch1_max_extent_px = 0
343
+
344
+ # full-volume heatmap accumulators
345
+ ch1_vol = np.zeros((D, H, W), dtype=np.float32)
346
+ ch2_vol = np.zeros((D, H, W), dtype=np.float32)
347
+ ch3_vol = np.zeros((D, H, W), dtype=np.float32)
348
+
349
+ for z in range(D):
350
+ sl = vol[z]
351
+ if is_skull_base(sl):
352
+ if z in plot_zs:
353
+ plot_cache[z] = None # signal "skull base, plain image only"
354
+ continue
355
+
356
+ body = body_mask(sl)
357
+ bone_m = bone_mask(sl)
358
+ spine_y, spine_x = find_spine_center(sl, H, W)
359
+ spine_m = spine_cylinder_mask(H, W, spine_y, spine_x, int(W * SPINE_RADIUS_FRAC))
360
+ work_mask = body & (~bone_m) & (~spine_m)
361
+ ant_work_mask = work_mask & anterior_mask(H, W, spine_y)
362
+
363
+ sob = sobel_clean(sl, body, bone_m, spine_m, sigma=params['sobel_sigma'])
364
+ h1, m1 = ch1_airway_lesion(sl, sob, midline, work_mask, H, W, params)
365
+ h2, m2 = ch2_nodal(sl, sob, midline, ant_work_mask, H, W, params)
366
+ h3 = ch3_necrosis(sl, midline, ant_work_mask, H, W, fluid_lo, fluid_hi, m1, m2)
367
+
368
+ # store raw per-slice heat into the full-volume accumulators
369
+ ch1_vol[z] = h1
370
+ ch2_vol[z] = h2
371
+ ch3_vol[z] = h3
372
+
373
+ # cache for the QC plot -- avoids any recompute later
374
+ if z in plot_zs:
375
+ plot_cache[z] = dict(sob=sob, h1=h1, h2=h2, h3=h3,
376
+ spine_y=spine_y, spine_x=spine_x)
377
+
378
+ n_valid_slices += 1
379
+
380
+ # ── rich per-slice stats (unchanged logic) ──────────────────────
381
+ if m2.sum() > 30:
382
+ l = int(m2[:, :midline].sum())
383
+ r = int(m2[:, midline:].sum())
384
+ total_nodal += l + r; left_total += l; right_total += r
385
+ if l + r > 0:
386
+ nodal_ratios.append(abs(l-r) / (l+r))
387
+ if l > 20 and r > 20:
388
+ bilateral_slices += 1
389
+ lbl_n, n_comp = ndi.label(m2)
390
+ if n_comp > 0:
391
+ comp_sizes = [(i, (lbl_n==i).sum()) for i in range(1, n_comp+1)]
392
+ best_comp = (lbl_n == max(comp_sizes, key=lambda x:x[1])[0])
393
+ ys_c, xs_c = np.where(best_comp)
394
+ h_comp = ys_c.max() - ys_c.min() + 1
395
+ w_comp = xs_c.max() - xs_c.min() + 1
396
+ max_node_short_axis_px = max(max_node_short_axis_px, min(h_comp, w_comp))
397
+ aw_here = airway_mask(sl, H, W)
398
+ if aw_here.sum() > 0:
399
+ aw_y = float(np.where(aw_here)[0].mean())
400
+ node_centroid_ys.append(float(ys_c.mean()) - aw_y)
401
+ raw_fluid = (ndi.binary_dilation(m2, iterations=1)
402
+ & (sl >= fluid_lo) & (sl <= fluid_hi) & work_mask)
403
+ necrosis_fracs.append(float(raw_fluid.sum() / max(m2.sum(), 1)))
404
+
405
+ raw_fluid_all = ndi.binary_dilation((m1 | m2), iterations=1) & (sl >= fluid_lo) & (sl <= fluid_hi)
406
+ max_necrosis = max(max_necrosis, int(raw_fluid_all.sum()))
407
+
408
+ aw = airway_mask(sl, H, W)
409
+ if aw.sum() > 20:
410
+ l_aw = int(aw[:, :midline].sum()); r_aw = int(aw[:, midline:].sum())
411
+ if l_aw + r_aw > 0:
412
+ airway_effs.append(abs(l_aw - r_aw) / (l_aw + r_aw))
413
+
414
+ if m1.sum() > 10:
415
+ ch1_areas.append(int(m1.sum()))
416
+ ys_m1, xs_m1 = np.where(m1)
417
+ ch1_max_extent_px = max(ch1_max_extent_px, int(np.max(np.abs(xs_m1 - midline))))
418
+
419
+ # ── 3D coordinates ──────────────────────────────────────────────────
420
+ ch1_coords = volume_weighted_centroid(ch1_vol, midline)
421
+ ch2_coords = volume_weighted_centroid(ch2_vol, midline)
422
+
423
+ # ── save heatmap channels ────────────────────────────────────────────
424
+ heatmap_stack = np.stack([ch1_vol, ch2_vol, ch3_vol], axis=-1)
425
+ case_heatmap_dir = HEATMAP_DIR / dataset
426
+ case_heatmap_dir.mkdir(parents=True, exist_ok=True)
427
+ heatmap_path = case_heatmap_dir / f"{case_id}_heatmaps.nii.gz"
428
+ nib.save(nib.Nifti1Image(heatmap_stack, np.eye(4)), str(heatmap_path))
429
+
430
+ # ── QC plot, using cached slices (no recompute) ─────────────────────
431
+ zs_sorted = sorted(plot_zs)
432
+ fig, axes = plt.subplots(len(zs_sorted), 5, figsize=(18, 4*len(zs_sorted)))
433
+ if len(zs_sorted) == 1:
434
+ axes = axes[np.newaxis, :]
435
+ titles = ["CT + midline + spine", "Sobel (clean)", "Ch1: Airway lesion",
436
+ "Ch2: Nodal disease", "Ch3: Necrosis"]
437
+
438
+ for row_idx, z in enumerate(zs_sorted):
439
+ sl = vol[z]
440
+ cached = plot_cache.get(z)
441
+ if cached is None:
442
+ for c in range(5):
443
+ ax = axes[row_idx, c]
444
+ ax.imshow(sl, cmap='gray', vmin=0, vmax=1)
445
+ ax.axis('off')
446
+ continue
447
+ panels = [None, cached['sob'], cached['h1'], cached['h2'], cached['h3']]
448
+ cmaps = [None, "hot", "Reds", "YlOrRd", "Blues"]
449
+ for c in range(5):
450
+ ax = axes[row_idx, c]
451
+ ax.imshow(sl, cmap='gray', vmin=0, vmax=1)
452
+ if c == 0:
453
+ ax.axvline(midline, color='cyan', lw=2, ls='--')
454
+ theta = np.linspace(0, 2*np.pi, 60)
455
+ r = int(W * SPINE_RADIUS_FRAC)
456
+ ax.plot(cached['spine_x'] + r*np.cos(theta),
457
+ cached['spine_y'] + r*np.sin(theta), 'r-', lw=1.5, alpha=0.7)
458
+ ax.axhline(cached['spine_y'] - 2, color='yellow', lw=1.5, ls='--', alpha=0.6)
459
+ if panels[c] is not None:
460
+ pmax = panels[c].max()
461
+ if pmax > 0:
462
+ disp = panels[c] / pmax
463
+ masked = np.ma.masked_where(disp < 0.25, disp)
464
+ ax.imshow(masked, cmap=cmaps[c], alpha=0.65, vmin=0, vmax=1)
465
+ if row_idx == 0:
466
+ ax.set_title(titles[c], fontsize=11)
467
+ ax.axis('off')
468
+
469
+ title = (f"{case_id} [{dataset}] | {'CONTRAST' if is_contrast else 'NON-CONTRAST'} "
470
+ f"| midline={midline} | conf={conf:.2f}")
471
+ fig.suptitle(title, fontsize=14, fontweight='bold')
472
+ plt.tight_layout()
473
+ plot_path = PLOTS_DIR / f"{case_id}.png"
474
+ plt.savefig(plot_path, dpi=160, bbox_inches='tight')
475
+ plt.close(fig)
476
+
477
+ row = {
478
+ 'case_id': case_id, 'dataset': dataset,
479
+ 'hpv': hpv_label, 'is_contrast': is_contrast, 'contrast_conf': round(conf, 3),
480
+ 'midline': midline, 'n_valid_slices': n_valid_slices,
481
+ 'max_nodal_ratio': round(max(nodal_ratios), 4) if nodal_ratios else 0.0,
482
+ 'mean_nodal_ratio': round(float(np.mean(nodal_ratios)), 4) if nodal_ratios else 0.0,
483
+ 'total_nodal_area': total_nodal,
484
+ 'nodal_laterality_bias': round(abs(left_total-right_total) / max(left_total+right_total, 1), 4),
485
+ 'n_slices_with_nodal': len(nodal_ratios),
486
+ 'max_node_short_axis_px': max_node_short_axis_px,
487
+ 'bilateral_nodal_slices': bilateral_slices,
488
+ 'bilateral_nodal_frac': round(bilateral_slices / max(n_valid_slices, 1), 4),
489
+ 'mean_node_centroid_y_vs_airway': round(float(np.mean(node_centroid_ys)), 2) if node_centroid_ys else 0.0,
490
+ 'max_necrosis_area': max_necrosis,
491
+ 'mean_necrosis_frac': round(float(np.mean(necrosis_fracs)), 4) if necrosis_fracs else 0.0,
492
+ 'max_necrosis_frac': round(float(np.max(necrosis_fracs)), 4) if necrosis_fracs else 0.0,
493
+ 'n_slices_with_necrosis': sum(1 for f in necrosis_fracs if f > 0.05),
494
+ 'max_airway_eff': round(max(airway_effs), 4) if airway_effs else 0.0,
495
+ 'mean_airway_eff': round(float(np.mean(airway_effs)), 4) if airway_effs else 0.0,
496
+ 'n_slices_with_airway_eff': len(airway_effs),
497
+ 'ch1_max_area_px': max(ch1_areas) if ch1_areas else 0,
498
+ 'ch1_mean_area_px': round(float(np.mean(ch1_areas)), 1) if ch1_areas else 0.0,
499
+ 'ch1_max_extent_px': ch1_max_extent_px,
500
+ 'n_slices_with_ch1': len(ch1_areas),
501
+ 'ch1_airway_z': ch1_coords['z'], 'ch1_airway_y': ch1_coords['y'], 'ch1_airway_x': ch1_coords['x'],
502
+ 'ch1_airway_side': ch1_coords['side'], 'ch1_airway_volume_vox': ch1_coords['volume_vox'],
503
+ 'ch1_airway_peak_intensity': ch1_coords['peak_intensity'],
504
+ 'ch2_nodal_z': ch2_coords['z'], 'ch2_nodal_y': ch2_coords['y'], 'ch2_nodal_x': ch2_coords['x'],
505
+ 'ch2_nodal_side': ch2_coords['side'], 'ch2_nodal_volume_vox': ch2_coords['volume_vox'],
506
+ 'ch2_nodal_peak_intensity': ch2_coords['peak_intensity'],
507
+ 'plot_path': str(plot_path.relative_to(BASE)),
508
+ 'heatmap_path': str(heatmap_path.relative_to(BASE)),
509
+ }
510
+ return row, None
511
+
512
+
513
+ # ══════════════════════════════════════════════════════════════════════
514
+ # Parallel + resumable driver
515
+ # ══════════════════════════════════════════════════════════════════════
516
+
517
+ def _process_one(args):
518
+ """Top-level, picklable wrapper required by ProcessPoolExecutor."""
519
+ case_id, dataset, hpv_label = args
520
+ try:
521
+ row, err = process_case_full(case_id, dataset, hpv_label)
522
+ return row, err
523
+ except Exception as e:
524
+ return None, f"{case_id}: {e}"
525
+
526
+ def load_done_case_ids():
527
+ """Resume support: cases already present in the CSV are skipped."""
528
+ if not FEATURES_CSV_PATH.exists():
529
+ return set()
530
+ existing = pd.read_csv(FEATURES_CSV_PATH, usecols=['case_id'])
531
+ return set(existing['case_id'].tolist())
532
+
533
+ def append_row_to_csv(row):
534
+ file_exists = FEATURES_CSV_PATH.exists()
535
+ with open(FEATURES_CSV_PATH, "a", newline="") as f:
536
+ writer = csv.DictWriter(f, fieldnames=CSV_FIELDNAMES)
537
+ if not file_exists:
538
+ writer.writeheader()
539
+ writer.writerow(row)
540
+
541
+
542
+ if __name__ == "__main__":
543
+ master = pd.read_csv(MASTER)
544
+ usable = master[master["extracted"]].copy()
545
+
546
+ # ── Fresh-start switch: set True to wipe prior outputs and
547
+ # reprocess the full cohort from scratch; set False for normal
548
+ # resumable behavior (skip cases already in the output CSV). ────────
549
+ import shutil
550
+ FRESH_START = False
551
+
552
+ if FRESH_START:
553
+ if FEATURES_CSV_PATH.exists():
554
+ FEATURES_CSV_PATH.unlink()
555
+ print(f"Deleted: {FEATURES_CSV_PATH}")
556
+ if PLOTS_DIR.exists():
557
+ shutil.rmtree(PLOTS_DIR)
558
+ PLOTS_DIR.mkdir(parents=True)
559
+ print(f"Cleared: {PLOTS_DIR}")
560
+ if HEATMAP_DIR.exists():
561
+ shutil.rmtree(HEATMAP_DIR)
562
+ HEATMAP_DIR.mkdir(parents=True)
563
+ print(f"Cleared: {HEATMAP_DIR}")
564
+ print("Fresh start: all prior outputs cleared, full cohort will be reprocessed.\n")
565
+
566
+ done_ids = load_done_case_ids()
567
+ todo = [(row["case_id"], row["dataset"], row.get("hpv_norm"))
568
+ for _, row in usable.iterrows() if row["case_id"] not in done_ids]
569
+
570
+ print(f"Total cases: {len(usable)} | Already done (resumed, skipped): {len(done_ids)} "
571
+ f"| To process: {len(todo)}")
572
+
573
+ # Scale to your available CPU cores; each worker builds a
574
+ # (D,H,W,3) float32 heatmap + a matplotlib figure per case, so
575
+ # watch memory usage if you see thrashing.
576
+ N_WORKERS = min(8, os.cpu_count() or 2)
577
+ print(f"Processing with {N_WORKERS} workers...")
578
+
579
+ n_done = 0
580
+ n_failed = 0
581
+ t0 = time.time()
582
+
583
+ if todo:
584
+ try:
585
+ with ProcessPoolExecutor(max_workers=N_WORKERS) as ex:
586
+ futs = {ex.submit(_process_one, t): t for t in todo}
587
+ for fut in as_completed(futs):
588
+ row, err = fut.result()
589
+ if err:
590
+ n_failed += 1
591
+ print(f"βœ— {err}")
592
+ elif row:
593
+ append_row_to_csv(row) # written immediately -- resumable
594
+ n_done += 1
595
+ print(f"βœ“ Saved {row['case_id']} β†’ {Path(row['plot_path']).name} "
596
+ f"| heatmaps β†’ {Path(row['heatmap_path']).name}")
597
+ if (n_done + n_failed) % 50 == 0:
598
+ elapsed = time.time() - t0
599
+ rate = (n_done + n_failed) / elapsed
600
+ eta = (len(todo) - n_done - n_failed) / rate if rate > 0 else float('nan')
601
+ print(f" {n_done + n_failed}/{len(todo)} {elapsed:.0f}s elapsed "
602
+ f"ETA {eta/60:.1f}min")
603
+ except Exception as mp_err:
604
+ print(f"Multiprocessing failed ({mp_err}), falling back to serial...")
605
+ for case_id, dataset, hpv_label in todo:
606
+ row, err = _process_one((case_id, dataset, hpv_label))
607
+ if err:
608
+ n_failed += 1; print(f"βœ— {err}")
609
+ elif row:
610
+ append_row_to_csv(row); n_done += 1
611
+ print(f"βœ“ Saved {row['case_id']}")
612
+
613
+ elapsed = time.time() - t0
614
+ print(f"\nβœ… DONE in {elapsed/60:.1f}min. "
615
+ f"{n_done} newly processed, {n_failed} failed, {len(done_ids)} skipped (already done).")
616
+ print(f"Features CSV: {FEATURES_CSV_PATH}")
617
+ print(f"Plots: {PLOTS_DIR}/")
618
+ print(f"Heatmaps: {HEATMAP_DIR}/<dataset>/<case_id>_heatmaps.nii.gz "
619
+ f"(4D NIfTI, D x H x W x 3, channels = [airway_lesion, nodal, necrosis])")