""" 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()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}//_heatmaps.nii.gz " f"(4D NIfTI, D x H x W x 3, channels = [airway_lesion, nodal, necrosis])")