AsymmetryNet / AsymmetryNet_extraction.py
jdmayfield's picture
Rename AAM v3 extraction to AsymmetryNet_extraction.py
5cd05ad verified
Raw History Blame Contribute Delete
28.8 kB
"""
AsymmetryNet: Asymmetry Attention Mechanism (AAM) β€” segmentation-free
feature extraction for radiologist-inspired left-right asymmetry in
head and neck CT, developed for HPV status prediction in oropharyngeal
squamous cell carcinoma (OPSCC).
Given a cropped neck CT volume (NIfTI), this pipeline computes three
per-slice asymmetry channels across the full volume in a single pass:
Ch1 (airway lesion): seeded region growing from the airway lumen,
bounded by Sobel-detected soft-tissue edges,
constrained to the side of airway deviation
from a spine-anchored midline.
Ch2 (nodal/soft-tissue asymmetry): mirrored left-right comparison of
fat-plane and soft-tissue density about the
same midline.
Ch3 (necrosis): per-case, histogram-derived HU-windowed fluid
detection, restricted to the union of the
Ch1 and Ch2 regions.
Midline is detected per case via spine-anchored bone-mass/compactness
scoring, not image center or a fixed landmark.
Outputs per case: a QC plot (5 representative slices), a 4D heatmap
NIfTI (D x H x W x 3, channels = [airway_lesion, nodal, necrosis]) for
downstream radiomic feature extraction restricted to these regions,
and a row of per-case rich features appended to a CSV.
Parallelized via ProcessPoolExecutor and resumable: interrupting and
restarting skips cases already written to the output CSV.
"""
import numpy as np
import nibabel as nib
import pandas as pd
import matplotlib
matplotlib.use('Agg') # required for plotting inside worker subprocesses
import matplotlib.pyplot as plt
from pathlib import Path
from scipy import ndimage as ndi
from scipy.signal import find_peaks
import os, time, csv
from concurrent.futures import ProcessPoolExecutor, as_completed
# ── Paths & Output ───────────────────────────────────────────────────
# >>> CHANGE THESE PATHS for your own environment before running <<<
BASE = Path("/path/to/project")
CROP_DIR = BASE / "crops" # expects {CROP_DIR}/{dataset}/{case_id}/crop224.nii.gz
MASTER = BASE / "master_index.csv" # expects columns: case_id, dataset, extracted, [hpv_norm]
OUTPUT_DIR = BASE / "AAM_results"
PLOTS_DIR = OUTPUT_DIR / "plots"
HEATMAP_DIR = OUTPUT_DIR / "heatmaps"
FEATURES_CSV_PATH = OUTPUT_DIR / "aam_features_rich.csv"
PLOTS_DIR.mkdir(parents=True, exist_ok=True)
HEATMAP_DIR.mkdir(parents=True, exist_ok=True)
print(f"Results will be saved to: {OUTPUT_DIR}")
# ── Constants ──────────────────────────────────────────────────────────
HU_MIN, HU_MAX = -200.0, 300.0
# z-gating already happened during crop generation (z_ctr/thick_mm
# centering + shift/flip/reshape corrections) -- gating further here
# would discard real oropharyngeal content across the full crop depth
Z_GATE_FRAC = 0.0
AIR_THRESH = 0.08
BONE_THRESH = 0.88
SOFT_LO = 0.34
# Fraction-of-width is the right approach after all: it automatically
# scales with per-case zoom/FOV (a physically-zoomed-in crop makes the
# same vertebra occupy more pixels of the same W, which a percentage
# tracks correctly), and vertebral size itself scales somewhat with
# overall neck/body size (sex, habitus) -- both handled better by a
# relative measure than a fixed pixel constant.
SPINE_RADIUS_FRAC = 0.1
MIN_FLUID_PX = 15
CENTROID_HEAT_THRESH = 0.3 # consistent with the 0.3 heat-significance convention used throughout
PARAMS_CONTRAST = {'sobel_sigma': 0.5, 'sobel_edge_thresh': 0.22, 'grow_hu_tol': 0.10, 'grow_iters': 15, 'nodal_thresh': 0.18}
PARAMS_NONCONTRAST = {'sobel_sigma': 1.0, 'sobel_edge_thresh': 0.20, 'grow_hu_tol': 0.12, 'grow_iters': 12, 'nodal_thresh': 0.12}
CSV_FIELDNAMES = [
'case_id', 'dataset', 'hpv', 'is_contrast', 'contrast_conf', 'midline',
'n_valid_slices',
'max_nodal_ratio', 'mean_nodal_ratio', 'total_nodal_area',
'nodal_laterality_bias', 'n_slices_with_nodal',
'max_node_short_axis_px', 'bilateral_nodal_slices', 'bilateral_nodal_frac',
'mean_node_centroid_y_vs_airway',
'max_necrosis_area', 'mean_necrosis_frac', 'max_necrosis_frac',
'n_slices_with_necrosis',
'max_airway_eff', 'mean_airway_eff', 'n_slices_with_airway_eff',
'ch1_max_area_px', 'ch1_mean_area_px', 'ch1_max_extent_px', 'n_slices_with_ch1',
'ch1_airway_z', 'ch1_airway_y', 'ch1_airway_x', 'ch1_airway_side',
'ch1_airway_volume_vox', 'ch1_airway_peak_intensity',
'ch2_nodal_z', 'ch2_nodal_y', 'ch2_nodal_x', 'ch2_nodal_side',
'ch2_nodal_volume_vox', 'ch2_nodal_peak_intensity',
'plot_path', 'heatmap_path',
]
# ── Helper functions (identical detection logic throughout) ───────────
def hu_to_norm(hu):
return (hu - HU_MIN) / (HU_MAX - HU_MIN)
def detect_contrast(vol, z_gate):
oropharynx = vol[z_gate:]
soft = oropharynx[(oropharynx > 0.15) & (oropharynx < BONE_THRESH)]
if soft.size == 0: return True, 0.5
enhance_frac = (soft > 0.65).sum() / soft.size
is_contrast = enhance_frac > 0.03
confidence = min(abs(enhance_frac - 0.03) / 0.03, 1.0)
return is_contrast, confidence
def get_params(is_contrast):
return PARAMS_CONTRAST if is_contrast else PARAMS_NONCONTRAST
def find_spine_center(sl, H, W):
bone = ndi.binary_fill_holes(ndi.binary_closing(sl >= BONE_THRESH, iterations=2))
lower_half = sl[H//2:, :]
is_air_heavy = (lower_half < 0.1).mean() > 0.6
start_y = H // 2
best_y, best_x = int(0.75*H), W//2
best_score = -1.0
directions = [range(start_y, H-20)]
if is_air_heavy:
directions.append(range(start_y, max(30, int(H*0.35)), -1))
for y_range in directions:
for y_start in y_range:
posterior = np.zeros((H, W), bool)
posterior[y_start:, :] = True
spine_bone = bone & posterior
if spine_bone.sum() < 25: continue
ys, xs = np.where(spine_bone)
cy, cx = int(ys.mean()), int(xs.mean())
local = spine_bone[max(0,cy-35):cy+35, max(0,cx-30):cx+30]
if local.size == 0: continue
compactness = local.sum() / local.size
score = sl[spine_bone].mean() * compactness * np.sqrt(spine_bone.sum())
if score > best_score:
best_score = score; best_y, best_x = cy, cx
return best_y, best_x
def body_mask(sl):
return ndi.binary_erosion(ndi.binary_fill_holes(sl >= AIR_THRESH), iterations=2)
def bone_mask(sl):
bone = ndi.binary_closing(sl >= BONE_THRESH, iterations=2)
return ndi.binary_dilation(ndi.binary_fill_holes(bone), iterations=3)
def airway_mask(sl, H, W):
rl, rh = int(0.15*H), int(0.72*H)
lbl, nf = ndi.label(sl < AIR_THRESH)
edge = set(lbl[0,:])|set(lbl[-1,:])|set(lbl[:,0])|set(lbl[:,-1])
aw = np.zeros(sl.shape, bool)
for i in range(1, nf+1):
if i in edge: continue
c = (lbl==i)
if c.sum() < 20: continue
ys, xs = np.where(c)
if ys.mean()<rl or ys.mean()>rh: continue
if (xs.max()-xs.min()+1)/W > 0.45: continue
aw |= c
return aw
def global_midline(vol, z_gate):
D, H, W = vol.shape
cx_list = []
for z in range(z_gate, D):
sl = vol[z]
if is_skull_base(sl): continue
_, spx = find_spine_center(sl, H, W)
cx_list.append(spx)
return int(np.median(cx_list)) if cx_list else W//2
def sobel_clean(sl, body, bone_m, spine_m, sigma=0.8):
clean = sl * body * (~bone_m) * (~spine_m)
clean = np.clip(clean, 0, hu_to_norm(80))
sm = ndi.gaussian_filter(clean, sigma)
gy = ndi.sobel(sm, axis=0); gx = ndi.sobel(sm, axis=1)
mag = np.sqrt(gy**2 + gx**2)
mag *= body * (~bone_m) * (~spine_m)
return mag
def mirror_about(arr, midline, W):
cw = min(midline, W-midline)
if cw < 5: return arr.copy()
out = np.zeros_like(arr)
left = arr[:, midline-cw:midline]; right = arr[:, midline:midline+cw]
out[:, midline-cw:midline] = np.flip(right, axis=1)
out[:, midline:midline+cw] = np.flip(left, axis=1)
return out
def anterior_mask(H, W, spine_y):
mask = np.zeros((H, W), bool)
mask[:spine_y-2, :] = True
return mask
def spine_cylinder_mask(H, W, spine_y, spine_x, radius):
yy, xx = np.ogrid[:H, :W]
return np.sqrt((yy-spine_y)**2 + (xx-spine_x)**2) < radius
def is_skull_base(sl, thresh=0.25):
return (sl >= BONE_THRESH).sum() / sl.size > thresh
def ch1_airway_lesion(sl, sobel_map, midline, work_mask, H, W, params):
aw = airway_mask(sl, H, W)
if aw.sum() < 20:
return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
aw_cx = np.where(aw)[1].mean()
shift_px = aw_cx - midline
if abs(shift_px) < 2:
lesion_side = 'left' if aw[:,:midline].sum() < aw[:,midline:].sum() else 'right'
else:
lesion_side = 'left' if shift_px > 0 else 'right'
aw_mirror = mirror_about(aw.astype(np.float32), midline, W) > 0.5
tissue = (sl >= SOFT_LO) & work_mask
seed = tissue & aw_mirror & (~aw)
if lesion_side == 'left': seed[:, midline:] = False
else: seed[:, :midline] = False
if seed.sum() < 3:
return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
sob_norm = sobel_map / (sobel_map.max() + 1e-6)
barrier = sob_norm > params['sobel_edge_thresh']
grown = seed.copy()
for _ in range(params['grow_iters']):
ref_hu = sl[grown].mean()
expanded = ndi.binary_dilation(grown, iterations=1)
candidates = expanded & (~grown) & work_mask & (~barrier)
candidates &= np.abs(sl - ref_hu) < params['grow_hu_tol']
if lesion_side == 'left': candidates[:, midline:] = False
else: candidates[:, :midline] = False
if candidates.sum() == 0: break
grown |= candidates
heat = ndi.gaussian_filter(grown.astype(np.float32), 2.0)
return (heat/(heat.max()+1e-6) if heat.max()>0 else heat), grown.copy()
def ch2_nodal(sl, sobel_map, midline, work_mask, H, W, params):
body = ndi.binary_erosion(
ndi.binary_fill_holes(sl >= AIR_THRESH),
iterations=max(4, int(0.04*max(H,W))))
bone_int = bone_mask(sl)
aw = airway_mask(sl, H, W)
central = ndi.binary_dilation(aw, iterations=10)
spy, _ = np.where(bone_int)
spine_y = int(np.percentile(spy, 75)) if len(spy)>0 else int(0.7*H)
deep_roi = body & (~bone_int) & (~central)
deep_roi[spine_y:, :] = False
if deep_roi.sum() < 50:
return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
FAT_LO = hu_to_norm(-150); FAT_HI = hu_to_norm(-30)
soft = (sl >= SOFT_LO) & (sl < BONE_THRESH) & deep_roi
fat = (sl >= FAT_LO) & (sl <= FAT_HI) & deep_roi
cw = min(midline, W-midline)
if cw < 10: return np.zeros((H,W), np.float32), np.zeros((H,W), bool)
mir_fat = mirror_about(fat.astype(np.float32), midline, W) > 0.5
A = (soft & mir_fat).astype(np.float32)
xor = fat ^ (mirror_about(fat.astype(np.float32), midline, W) > 0.5)
softness = np.clip((sl - SOFT_LO) / (0.64 - SOFT_LO), 0, 1)
B = xor.astype(np.float32) * softness * soft
heat = ndi.gaussian_filter(np.maximum(A, B), 4.0) * deep_roi
mask = (heat > heat.max()*0.3) if heat.max()>0 else np.zeros((H,W), bool)
return (heat/(heat.max()+1e-6) if heat.max()>0 else heat), mask
def compute_fluid_window(vol, z_gate):
oropharynx = vol[z_gate:]
body_vox = oropharynx[(oropharynx > 0.05) & (oropharynx < 0.85)]
counts, edges = np.histogram(body_vox, bins=200)
centers = 0.5*(edges[:-1]+edges[1:])
smooth = ndi.gaussian_filter1d(counts.astype(float), sigma=3)
peaks, _ = find_peaks(smooth, height=smooth.max()*0.05, distance=15)
soft_cands = [(i, smooth[p]) for i,p in enumerate(peaks) if centers[p]>0.30]
soft_peak = centers[peaks[max(soft_cands, key=lambda x:x[1])[0]]] if soft_cands else 0.48
return soft_peak-0.08, soft_peak-0.02, soft_peak
def ch3_necrosis(sl, midline, work_mask, H, W, fluid_lo, fluid_hi, m1, m2):
lesion_roi = ndi.binary_dilation(m1|m2, iterations=1)
if lesion_roi.sum() < 5: return np.zeros((H,W), np.float32)
aw_dil = ndi.binary_dilation(airway_mask(sl,H,W), iterations=3)
fluid = (sl>=fluid_lo)&(sl<=fluid_hi)&work_mask&(~aw_dil)&lesion_roi
lbl, n = ndi.label(fluid)
heat = np.zeros((H,W), np.float32)
for i in range(1,n+1):
c=(lbl==i)
if c.sum()>=MIN_FLUID_PX: heat[c]=1.0
heat = ndi.gaussian_filter(heat, 3.0)
return heat/(heat.max()+1e-6) if heat.max()>0 else heat
def volume_weighted_centroid(heat_vol, midline, thresh=CENTROID_HEAT_THRESH):
mask = heat_vol > thresh
n_vox = int(mask.sum())
if n_vox == 0:
return {"z": np.nan, "y": np.nan, "x": np.nan,
"volume_vox": 0, "side": None, "peak_intensity": 0.0}
zs, ys, xs = np.where(mask)
weights = heat_vol[mask]
z_c = float(np.average(zs, weights=weights))
y_c = float(np.average(ys, weights=weights))
x_c = float(np.average(xs, weights=weights))
side = "left" if x_c > midline else "right"
return {"z": z_c, "y": y_c, "x": x_c,
"volume_vox": n_vox, "side": side,
"peak_intensity": float(heat_vol.max())}
# ══════════════════════════════════════════════════════════════════════
# SINGLE consolidated per-case worker: one pass over all slices produces
# the plot, the heatmap NIfTI, AND the rich feature row together.
# ══════════════════════════════════════════════════════════════════════
def process_case_full(case_id, dataset, hpv_label=None):
path = CROP_DIR / dataset / case_id / "crop224.nii.gz"
if not path.exists():
return None, "missing_crop"
vol = nib.load(str(path)).get_fdata().astype(np.float32)
if vol.max() > 1.5:
vol = np.clip(vol, HU_MIN, HU_MAX)
vol = (vol - HU_MIN) / (HU_MAX - HU_MIN)
vol = np.clip(vol, 0, 1).astype(np.float32)
D, H, W = vol.shape
z_gate = int(D * Z_GATE_FRAC)
midline = global_midline(vol, z_gate)
is_contrast, conf = detect_contrast(vol, z_gate)
params = get_params(is_contrast)
fluid_lo, fluid_hi, _ = compute_fluid_window(vol, z_gate)
# which slices get cached for the QC plot (5 representative slices)
plot_zs = set(np.linspace(0, D - 1, 5).round().astype(int).tolist())
plot_cache = {}
# rich-feature accumulators
nodal_ratios, airway_effs = [], []
total_nodal = left_total = right_total = 0
max_necrosis = 0
n_valid_slices = 0
max_node_short_axis_px = 0
necrosis_fracs = []
node_centroid_ys = []
bilateral_slices = 0
ch1_areas = []
ch1_max_extent_px = 0
# full-volume heatmap accumulators
ch1_vol = np.zeros((D, H, W), dtype=np.float32)
ch2_vol = np.zeros((D, H, W), dtype=np.float32)
ch3_vol = np.zeros((D, H, W), dtype=np.float32)
for z in range(D):
sl = vol[z]
if is_skull_base(sl):
if z in plot_zs:
plot_cache[z] = None # signal "skull base, plain image only"
continue
body = body_mask(sl)
bone_m = bone_mask(sl)
spine_y, spine_x = find_spine_center(sl, H, W)
spine_m = spine_cylinder_mask(H, W, spine_y, spine_x, int(W * SPINE_RADIUS_FRAC))
work_mask = body & (~bone_m) & (~spine_m)
ant_work_mask = work_mask & anterior_mask(H, W, spine_y)
sob = sobel_clean(sl, body, bone_m, spine_m, sigma=params['sobel_sigma'])
h1, m1 = ch1_airway_lesion(sl, sob, midline, work_mask, H, W, params)
h2, m2 = ch2_nodal(sl, sob, midline, ant_work_mask, H, W, params)
h3 = ch3_necrosis(sl, midline, ant_work_mask, H, W, fluid_lo, fluid_hi, m1, m2)
# store raw per-slice heat into the full-volume accumulators
ch1_vol[z] = h1
ch2_vol[z] = h2
ch3_vol[z] = h3
# cache for the QC plot -- avoids any recompute later
if z in plot_zs:
plot_cache[z] = dict(sob=sob, h1=h1, h2=h2, h3=h3,
spine_y=spine_y, spine_x=spine_x)
n_valid_slices += 1
# ── rich per-slice stats (unchanged logic) ──────────────────────
if m2.sum() > 30:
l = int(m2[:, :midline].sum())
r = int(m2[:, midline:].sum())
total_nodal += l + r; left_total += l; right_total += r
if l + r > 0:
nodal_ratios.append(abs(l-r) / (l+r))
if l > 20 and r > 20:
bilateral_slices += 1
lbl_n, n_comp = ndi.label(m2)
if n_comp > 0:
comp_sizes = [(i, (lbl_n==i).sum()) for i in range(1, n_comp+1)]
best_comp = (lbl_n == max(comp_sizes, key=lambda x:x[1])[0])
ys_c, xs_c = np.where(best_comp)
h_comp = ys_c.max() - ys_c.min() + 1
w_comp = xs_c.max() - xs_c.min() + 1
max_node_short_axis_px = max(max_node_short_axis_px, min(h_comp, w_comp))
aw_here = airway_mask(sl, H, W)
if aw_here.sum() > 0:
aw_y = float(np.where(aw_here)[0].mean())
node_centroid_ys.append(float(ys_c.mean()) - aw_y)
raw_fluid = (ndi.binary_dilation(m2, iterations=1)
& (sl >= fluid_lo) & (sl <= fluid_hi) & work_mask)
necrosis_fracs.append(float(raw_fluid.sum() / max(m2.sum(), 1)))
raw_fluid_all = ndi.binary_dilation((m1 | m2), iterations=1) & (sl >= fluid_lo) & (sl <= fluid_hi)
max_necrosis = max(max_necrosis, int(raw_fluid_all.sum()))
aw = airway_mask(sl, H, W)
if aw.sum() > 20:
l_aw = int(aw[:, :midline].sum()); r_aw = int(aw[:, midline:].sum())
if l_aw + r_aw > 0:
airway_effs.append(abs(l_aw - r_aw) / (l_aw + r_aw))
if m1.sum() > 10:
ch1_areas.append(int(m1.sum()))
ys_m1, xs_m1 = np.where(m1)
ch1_max_extent_px = max(ch1_max_extent_px, int(np.max(np.abs(xs_m1 - midline))))
# ── 3D coordinates ──────────────────────────────────────────────────
ch1_coords = volume_weighted_centroid(ch1_vol, midline)
ch2_coords = volume_weighted_centroid(ch2_vol, midline)
# ── save heatmap channels ────────────────────────────────────────────
heatmap_stack = np.stack([ch1_vol, ch2_vol, ch3_vol], axis=-1)
case_heatmap_dir = HEATMAP_DIR / dataset
case_heatmap_dir.mkdir(parents=True, exist_ok=True)
heatmap_path = case_heatmap_dir / f"{case_id}_heatmaps.nii.gz"
nib.save(nib.Nifti1Image(heatmap_stack, np.eye(4)), str(heatmap_path))
# ── QC plot, using cached slices (no recompute) ─────────────────────
zs_sorted = sorted(plot_zs)
fig, axes = plt.subplots(len(zs_sorted), 5, figsize=(18, 4*len(zs_sorted)))
if len(zs_sorted) == 1:
axes = axes[np.newaxis, :]
titles = ["CT + midline + spine", "Sobel (clean)", "Ch1: Airway lesion",
"Ch2: Nodal disease", "Ch3: Necrosis"]
for row_idx, z in enumerate(zs_sorted):
sl = vol[z]
cached = plot_cache.get(z)
if cached is None:
for c in range(5):
ax = axes[row_idx, c]
ax.imshow(sl, cmap='gray', vmin=0, vmax=1)
ax.axis('off')
continue
panels = [None, cached['sob'], cached['h1'], cached['h2'], cached['h3']]
cmaps = [None, "hot", "Reds", "YlOrRd", "Blues"]
for c in range(5):
ax = axes[row_idx, c]
ax.imshow(sl, cmap='gray', vmin=0, vmax=1)
if c == 0:
ax.axvline(midline, color='cyan', lw=2, ls='--')
theta = np.linspace(0, 2*np.pi, 60)
r = int(W * SPINE_RADIUS_FRAC)
ax.plot(cached['spine_x'] + r*np.cos(theta),
cached['spine_y'] + r*np.sin(theta), 'r-', lw=1.5, alpha=0.7)
ax.axhline(cached['spine_y'] - 2, color='yellow', lw=1.5, ls='--', alpha=0.6)
if panels[c] is not None:
pmax = panels[c].max()
if pmax > 0:
disp = panels[c] / pmax
masked = np.ma.masked_where(disp < 0.25, disp)
ax.imshow(masked, cmap=cmaps[c], alpha=0.65, vmin=0, vmax=1)
if row_idx == 0:
ax.set_title(titles[c], fontsize=11)
ax.axis('off')
title = (f"{case_id} [{dataset}] | {'CONTRAST' if is_contrast else 'NON-CONTRAST'} "
f"| midline={midline} | conf={conf:.2f}")
fig.suptitle(title, fontsize=14, fontweight='bold')
plt.tight_layout()
plot_path = PLOTS_DIR / f"{case_id}.png"
plt.savefig(plot_path, dpi=160, bbox_inches='tight')
plt.close(fig)
row = {
'case_id': case_id, 'dataset': dataset,
'hpv': hpv_label, 'is_contrast': is_contrast, 'contrast_conf': round(conf, 3),
'midline': midline, 'n_valid_slices': n_valid_slices,
'max_nodal_ratio': round(max(nodal_ratios), 4) if nodal_ratios else 0.0,
'mean_nodal_ratio': round(float(np.mean(nodal_ratios)), 4) if nodal_ratios else 0.0,
'total_nodal_area': total_nodal,
'nodal_laterality_bias': round(abs(left_total-right_total) / max(left_total+right_total, 1), 4),
'n_slices_with_nodal': len(nodal_ratios),
'max_node_short_axis_px': max_node_short_axis_px,
'bilateral_nodal_slices': bilateral_slices,
'bilateral_nodal_frac': round(bilateral_slices / max(n_valid_slices, 1), 4),
'mean_node_centroid_y_vs_airway': round(float(np.mean(node_centroid_ys)), 2) if node_centroid_ys else 0.0,
'max_necrosis_area': max_necrosis,
'mean_necrosis_frac': round(float(np.mean(necrosis_fracs)), 4) if necrosis_fracs else 0.0,
'max_necrosis_frac': round(float(np.max(necrosis_fracs)), 4) if necrosis_fracs else 0.0,
'n_slices_with_necrosis': sum(1 for f in necrosis_fracs if f > 0.05),
'max_airway_eff': round(max(airway_effs), 4) if airway_effs else 0.0,
'mean_airway_eff': round(float(np.mean(airway_effs)), 4) if airway_effs else 0.0,
'n_slices_with_airway_eff': len(airway_effs),
'ch1_max_area_px': max(ch1_areas) if ch1_areas else 0,
'ch1_mean_area_px': round(float(np.mean(ch1_areas)), 1) if ch1_areas else 0.0,
'ch1_max_extent_px': ch1_max_extent_px,
'n_slices_with_ch1': len(ch1_areas),
'ch1_airway_z': ch1_coords['z'], 'ch1_airway_y': ch1_coords['y'], 'ch1_airway_x': ch1_coords['x'],
'ch1_airway_side': ch1_coords['side'], 'ch1_airway_volume_vox': ch1_coords['volume_vox'],
'ch1_airway_peak_intensity': ch1_coords['peak_intensity'],
'ch2_nodal_z': ch2_coords['z'], 'ch2_nodal_y': ch2_coords['y'], 'ch2_nodal_x': ch2_coords['x'],
'ch2_nodal_side': ch2_coords['side'], 'ch2_nodal_volume_vox': ch2_coords['volume_vox'],
'ch2_nodal_peak_intensity': ch2_coords['peak_intensity'],
'plot_path': str(plot_path.relative_to(BASE)),
'heatmap_path': str(heatmap_path.relative_to(BASE)),
}
return row, None
# ══════════════════════════════════════════════════════════════════════
# Parallel + resumable driver
# ══════════════════════════════════════════════════════════════════════
def _process_one(args):
"""Top-level, picklable wrapper required by ProcessPoolExecutor."""
case_id, dataset, hpv_label = args
try:
row, err = process_case_full(case_id, dataset, hpv_label)
return row, err
except Exception as e:
return None, f"{case_id}: {e}"
def load_done_case_ids():
"""Resume support: cases already present in the CSV are skipped."""
if not FEATURES_CSV_PATH.exists():
return set()
existing = pd.read_csv(FEATURES_CSV_PATH, usecols=['case_id'])
return set(existing['case_id'].tolist())
def append_row_to_csv(row):
file_exists = FEATURES_CSV_PATH.exists()
with open(FEATURES_CSV_PATH, "a", newline="") as f:
writer = csv.DictWriter(f, fieldnames=CSV_FIELDNAMES)
if not file_exists:
writer.writeheader()
writer.writerow(row)
if __name__ == "__main__":
master = pd.read_csv(MASTER)
usable = master[master["extracted"]].copy()
# ── Fresh-start switch: set True to wipe prior outputs and
# reprocess the full cohort from scratch; set False for normal
# resumable behavior (skip cases already in the output CSV). ────────
import shutil
FRESH_START = False
if FRESH_START:
if FEATURES_CSV_PATH.exists():
FEATURES_CSV_PATH.unlink()
print(f"Deleted: {FEATURES_CSV_PATH}")
if PLOTS_DIR.exists():
shutil.rmtree(PLOTS_DIR)
PLOTS_DIR.mkdir(parents=True)
print(f"Cleared: {PLOTS_DIR}")
if HEATMAP_DIR.exists():
shutil.rmtree(HEATMAP_DIR)
HEATMAP_DIR.mkdir(parents=True)
print(f"Cleared: {HEATMAP_DIR}")
print("Fresh start: all prior outputs cleared, full cohort will be reprocessed.\n")
done_ids = load_done_case_ids()
todo = [(row["case_id"], row["dataset"], row.get("hpv_norm"))
for _, row in usable.iterrows() if row["case_id"] not in done_ids]
print(f"Total cases: {len(usable)} | Already done (resumed, skipped): {len(done_ids)} "
f"| To process: {len(todo)}")
# Scale to your available CPU cores; each worker builds a
# (D,H,W,3) float32 heatmap + a matplotlib figure per case, so
# watch memory usage if you see thrashing.
N_WORKERS = min(8, os.cpu_count() or 2)
print(f"Processing with {N_WORKERS} workers...")
n_done = 0
n_failed = 0
t0 = time.time()
if todo:
try:
with ProcessPoolExecutor(max_workers=N_WORKERS) as ex:
futs = {ex.submit(_process_one, t): t for t in todo}
for fut in as_completed(futs):
row, err = fut.result()
if err:
n_failed += 1
print(f"βœ— {err}")
elif row:
append_row_to_csv(row) # written immediately -- resumable
n_done += 1
print(f"βœ“ Saved {row['case_id']} β†’ {Path(row['plot_path']).name} "
f"| heatmaps β†’ {Path(row['heatmap_path']).name}")
if (n_done + n_failed) % 50 == 0:
elapsed = time.time() - t0
rate = (n_done + n_failed) / elapsed
eta = (len(todo) - n_done - n_failed) / rate if rate > 0 else float('nan')
print(f" {n_done + n_failed}/{len(todo)} {elapsed:.0f}s elapsed "
f"ETA {eta/60:.1f}min")
except Exception as mp_err:
print(f"Multiprocessing failed ({mp_err}), falling back to serial...")
for case_id, dataset, hpv_label in todo:
row, err = _process_one((case_id, dataset, hpv_label))
if err:
n_failed += 1; print(f"βœ— {err}")
elif row:
append_row_to_csv(row); n_done += 1
print(f"βœ“ Saved {row['case_id']}")
elapsed = time.time() - t0
print(f"\nβœ… DONE in {elapsed/60:.1f}min. "
f"{n_done} newly processed, {n_failed} failed, {len(done_ids)} skipped (already done).")
print(f"Features CSV: {FEATURES_CSV_PATH}")
print(f"Plots: {PLOTS_DIR}/")
print(f"Heatmaps: {HEATMAP_DIR}/<dataset>/<case_id>_heatmaps.nii.gz "
f"(4D NIfTI, D x H x W x 3, channels = [airway_lesion, nodal, necrosis])")