Download AsymmetryNet_extraction.py from jdmayfield/AsymmetryNet: direct link, hf CLI and curl.
- Browser
- Download file 28.8 kB
-
https://huggingface.co/jdmayfield/AsymmetryNet/resolve/main/AsymmetryNet_extraction.py
- Command line
-
hf download hf://jdmayfield/AsymmetryNet/AsymmetryNet_extraction.py
-
curl -L -o AsymmetryNet_extraction.py https://huggingface.co/jdmayfield/AsymmetryNet/resolve/main/AsymmetryNet_extraction.py
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])") |