""" Filename: part_C.py Description: Third part of processing Author: Tyler de Zeeuw License: GPL-3.0 """ # Built-in imports import os import sys import time import argparse import os.path as op # External library imports import numpy as np import pandas as pd from scipy.spatial.distance import cdist # mne imports from mne import get_config, set_config from mne.channels import make_standard_montage from mne.utils import _check_fname, _validate_type from mne.io import BaseRaw, read_raw_snirf from mne.preprocessing.nirs import optical_density, beer_lambert_law def print_elapsed(label="Timestamp"): """Helper function to print time elapsed since the script started.""" elapsed = time.perf_counter() - APP_START_TIME print(f"[{label}] {elapsed:.3f}s total elapsed time") def _read_fold_xls(fname, atlas="Juelich"): """Read fOLD toolbox xls file. The values are then manipulated in to a tidy dataframe. Note the xls files are not included as no license is provided. Parameters ---------- fname : str Path to xls file. atlas : str Requested atlas. """ page_reference = {"AAL2": 2, "AICHA": 5, "Brodmann": 8, "Juelich": 11, "Loni": 14} tbl = pd.read_excel(fname, sheet_name=page_reference[atlas]) # Remove the spacing between rows empty_rows = np.where(np.isnan(tbl["Specificity"]))[0] tbl = tbl.drop(empty_rows).reset_index(drop=True) # Empty values in the table mean its the same as above for row_idx in range(1, tbl.shape[0]): for col_idx, col in enumerate(tbl.columns): if not isinstance(tbl[col][row_idx], str): if np.isnan(tbl[col][row_idx]): tbl.iloc[row_idx, col_idx] = tbl.iloc[row_idx - 1, col_idx] tbl["Specificity"] = tbl["Specificity"] * 100 tbl["brainSens"] = tbl["brainSens"] * 100 return tbl def _find_closest_standard_location(position, reference, *, out="label"): """Return closest montage label to coordinates. Parameters ---------- position : array, shape (3,) Coordinates. reference : dataframe As generated by _generate_montage_locations. trans_pos : str Apply a transformation to positions to specified frame. Use None for no transformation. """ p0 = np.array(position).reshape(-1, 3) # head_mri_t, _ = _get_trans("fsaverage", "head", "mri") # p0 = apply_trans(head_mri_t, p0) dists = cdist(p0, np.asarray(reference[["x", "y", "z"]], float)) if out == "label": min_idx = np.argmin(dists) return reference["label"][min_idx] else: assert out == "dists" return dists def _source_detector_fold_table(raw, cidx, reference, fold_tbl, interpolate): src = raw.info["chs"][cidx]["loc"][3:6] det = raw.info["chs"][cidx]["loc"][6:9] ref_lab = list(reference["label"]) dists = _find_closest_standard_location([src, det], reference, out="dists") src_min, det_min = np.argmin(dists, axis=1) src_name, det_name = ref_lab[src_min], ref_lab[det_min] tbl = fold_tbl.query("Source == @src_name and Detector == @det_name") dist = np.linalg.norm(dists[[0, 1], [src_min, det_min]]) # Try reversing source and detector if len(tbl) == 0: tbl = fold_tbl.query("Source == @det_name and Detector == @src_name") if len(tbl) == 0 and interpolate: # Try something hopefully not too terrible: pick the one with the # smallest net distance good = np.isin(fold_tbl["Source"], reference["label"]) & np.isin( fold_tbl["Detector"], reference["label"] ) assert good.any() tbl = fold_tbl[good] assert len(tbl) src_idx = [ref_lab.index(src) for src in tbl["Source"]] det_idx = [ref_lab.index(det) for det in tbl["Detector"]] # Original tot_dist = np.linalg.norm([dists[0, src_idx], dists[1, det_idx]], axis=0) assert tot_dist.shape == (len(tbl),) idx = np.argmin(tot_dist) dist_1 = tot_dist[idx] src_1, det_1 = ref_lab[src_idx[idx]], ref_lab[det_idx[idx]] # And the reverse tot_dist = np.linalg.norm([dists[0, det_idx], dists[1, src_idx]], axis=0) idx = np.argmin(tot_dist) dist_2 = tot_dist[idx] src_2, det_2 = ref_lab[det_idx[idx]], ref_lab[src_idx[idx]] if dist_1 < dist_2: new_dist, src_use, det_use = dist_1, src_1, det_1 else: new_dist, src_use, det_use = dist_2, det_2, src_2 tbl = fold_tbl.query("Source == @src_use and Detector == @det_use") tbl = tbl.copy() tbl["BestSource"] = src_name tbl["BestDetector"] = det_name tbl["BestMatchDistance"] = dist tbl["MatchDistance"] = new_dist assert len(tbl) else: tbl = tbl.copy() tbl["BestSource"] = src_name tbl["BestDetector"] = det_name tbl["BestMatchDistance"] = dist tbl["MatchDistance"] = dist tbl = tbl.copy() # don't get warnings about setting values later return tbl def generate_montage_locations(): """Get standard MNI montage locations in dataframe. Data is returned in the same format as the eeg_positions library. """ # standard_1020 and standard_1005 are in MNI (fsaverage) space already, # but we need to undo the scaling that head_scale will do montage = make_standard_montage( "standard_1005", head_size=0.09700884729534559 ) for d in montage.dig: d["coord_frame"] = 2003 montage.dig[:] = montage.dig[3:] montage.add_mni_fiducials() # now in fsaverage space coords = pd.DataFrame.from_dict(montage.get_positions()["ch_pos"]).T coords["label"] = coords.index coords = coords.rename(columns={0: "x", 1: "y", 2: "z"}) return coords.reset_index(drop=True) def _check_load_fold(fold_files, atlas): # _validate_type(fold_files, (list, "path-like", None), "fold_files") if fold_files is None: fold_files = get_config("MNE_NIRS_FOLD_PATH") if fold_files is None: raise ValueError( "MNE_NIRS_FOLD_PATH not set, either set it using " "mne.set_config or pass fold_files as str or list" ) if not isinstance(fold_files, list): # path-like fold_files = _check_fname( fold_files, overwrite="read", must_exist=True, name="fold_files", need_dir=True, ) fold_files = [op.join(fold_files, f"10-{x}.xls") for x in (5, 10)] fold_tbl = pd.DataFrame() for fi, fname in enumerate(fold_files): fname = _check_fname( fname, overwrite="read", must_exist=True, name=f"fold_files[{fi}]" ) fold_tbl = pd.concat( [fold_tbl, _read_fold_xls(fname, atlas=atlas)], ignore_index=True ) return fold_tbl def fold_channel_specificity_normal(raw, fold_files=None, atlas="Juelich", interpolate=False): """Return the landmarks and specificity a channel is sensitive to. Parameters """ # noqa: E501 _validate_type(raw, BaseRaw, "raw") reference_locations = generate_montage_locations() fold_tbl = _check_load_fold(fold_files, atlas) chan_spec = list() for cidx in range(len(raw.ch_names)): tbl = _source_detector_fold_table( raw, cidx, reference_locations, fold_tbl, interpolate ) chan_spec.append(tbl.reset_index(drop=True)) return chan_spec def process_snirf_fold_fast(snirf_path: str, fold_dir_path: str, atlas: str = 'Brodmann') -> dict: """Fast vectorized version loading data structures exactly once.""" set_config('MNE_NIRS_FOLD_PATH', os.path.abspath(fold_dir_path)) print(f"Loading SNIRF file: {snirf_path}...") raw = read_raw_snirf(snirf_path, preload=True) print("Preprocessing data (OD -> BLL)...") raw_od = optical_density(raw) raw_haemo = beer_lambert_law(raw_od) # Isolate HbO channels hbo_raw = raw_haemo.copy().pick(picks='hbo') hbo_channel_names = hbo_raw.ch_names # CRITICAL PERFORMANCE FIX: Load static structures exactly ONCE print("Loading anatomical reference structures...") reference_locations = generate_montage_locations() fold_tbl = _check_load_fold(fold_dir_path, atlas) channel_results = {} print(f"Processing {len(hbo_channel_names)} HbO channels...") # Query directly by index using the cached structures for cidx, channel_name in enumerate(hbo_channel_names): tbl = _source_detector_fold_table( hbo_raw, cidx, reference_locations, fold_tbl, interpolate=True ) print_elapsed() channel_results[channel_name] = [] for _, row in tbl.iterrows(): channel_results[channel_name].append({ 'Landmark': str(row['Landmark']), 'Specificity': float(row['Specificity']) }) return channel_results def main(): #TODO: Likely does NOT work in a packaged build PATH_TO_FOLD_FOLDER = "~/mne_data/fOLD/fOLD-public-master/Supplementary" ATLAS_NAME = "Brodmann" # Options: "AAL2", "AICHA", "Brodmann", "Juelich", "Loni" try: results = process_snirf_fold_fast( snirf_path=SNIRF_FILE_PATH, fold_dir_path=os.path.expanduser(PATH_TO_FOLD_FOLDER), atlas=ATLAS_NAME ) # Pretty-print final results to console print("\n" + "="*50) print("FINAL FOLD SPECIFICITY RESULTS") print("="*50) for channel, mappings in results.items(): print(f"\nChannel: {channel}") if not mappings: print(" No anatomical mappings found.") for mapping in mappings: print(f" Region: {mapping['Landmark']:<25} | Specificity: {mapping['Specificity']:.2f}%") except Exception as e: print(f"\nAn error occurred: {e}", file=sys.stderr) print_elapsed("Complete total final time") if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--start-time", type=float, default=None) parser.add_argument("--snirf", type=str, default=None) args = parser.parse_args() APP_START_TIME = args.start_time if args.start_time is not None else time.perf_counter() SNIRF_FILE_PATH = args.snirf print(SNIRF_FILE_PATH) main()