diff --git a/npx_utils/__init__.py b/npx_utils/__init__.py index 13f7fea..69b6bbc 100644 --- a/npx_utils/__init__.py +++ b/npx_utils/__init__.py @@ -1,5 +1,6 @@ +from . import noise, plotting_helpers, stability from .data_helpers import * from .ks_helpers import * from .metrics import * from .other_helpers import * -from .sglx_helpers import * +from .sglx import sglx_helpers diff --git a/npx_utils/data_helpers.py b/npx_utils/data_helpers.py index 656467e..4559d0c 100644 --- a/npx_utils/data_helpers.py +++ b/npx_utils/data_helpers.py @@ -3,10 +3,17 @@ import cupy as cp import numpy as np +from joblib import Parallel, delayed from numpy.typing import NDArray from tqdm import tqdm -from npx_utils.sglx_helpers import get_bits_to_uV, get_data_memmap, read_meta +from npx_utils.ks_helpers import get_binary_path, get_meta_path +from npx_utils.sglx.sglx_helpers import ( + get_bits_to_uV, + get_channel_counts, + get_data_memmap, + read_meta, +) def extract_spikes( @@ -36,8 +43,8 @@ def extract_spikes( spikes are extracted. Defaults to -1. Returns: - spikes (NDArray): Array of extracted spike waveforms with shape - (# of spikes, # of channels, # of timepoints). + spikes: Array of extracted spike waveforms + with shape (# of spikes, # of channels, # of timepoints). """ times = times_multi[clust_id].astype("int64") # spikes cut off by the ends of the recording is handled in times_multi @@ -76,13 +83,18 @@ def extract_all_spikes( pre_samples: int, post_samples: int, max_spikes: int, + n_jobs: int = -1, ): - spikes = {} - for clust_id in tqdm(clust_ids, "Extracting spikes..."): - spikes[clust_id] = extract_spikes( + """ + Extracts spikes for all clusters in parallel. + """ + results = Parallel(n_jobs=n_jobs, backend="threading")( + delayed(extract_spikes)( data, times_multi, clust_id, pre_samples, post_samples, max_spikes ) - return spikes + for clust_id in tqdm(clust_ids, "Extracting spikes...") + ) + return dict(zip(clust_ids, results)) def calc_mean_wf( @@ -91,6 +103,7 @@ def calc_mean_wf( cluster_ids: list[int], times_multi: dict[NDArray[np.int_]], data: NDArray[np.int_], + all_spikes: dict[int, NDArray] = None, ) -> NDArray: """ Calculate mean waveform and std waveform for each cluster. Need to have loaded some metrics. If return_spikes is True, also returns the spike waveforms. @@ -102,11 +115,10 @@ def calc_mean_wf( cluster_ids (list): List of cluster ids to calculate waveforms for. times_multi (dict): Dictionary of spike times indexed by cluster id. data (NDArray): Ephys data with shape (n_timepoints, n_channels). + all_spikes (dict): Optional dictionary of spike waveforms for each cluster. Returns: NDArray: Mean waveforms for each cluster (uV). Shape (n_clusters, n_channels, pre_samples + post_samples) dtype float32 - NDArray: Std waveforms for each cluster (uV). Shape (n_clusters, n_channels, pre_samples + post_samples) dtype float32 - dict[int, NDArray]: Spike waveforms for each cluster (bits). NDArray shape (n_spikes, n_channels, pre_samples + post_samples) dtype int16 """ mean_wf_path = os.path.join(params["KS_folder"], "mean_waveforms.npy") @@ -128,32 +140,38 @@ def calc_mean_wf( params["pre_samples"] + params["post_samples"], ) ) - for i in tqdm(cluster_ids, desc="Calculating mean waveforms"): - spikes = extract_spikes( + + if all_spikes is None: + all_spikes = extract_all_spikes( data, times_multi, - i, + cluster_ids, params["pre_samples"], params["post_samples"], params["max_spikes"], ) - if len(spikes) > 0: # edge case - spikes_cp = cp.array(spikes, dtype=cp.float32) - mean_wf[i, :, :] = cp.mean(spikes_cp, axis=0) + for i in tqdm(cluster_ids, desc="Calculating mean waveforms"): + spikes = all_spikes[i] + if len(spikes) > 0: # edge case + spikes_cp = cp.array(spikes, dtype=cp.float32) + mean_wf[i, :, :] = cp.mean(spikes_cp, axis=0) # convert mean_wf uV meta = read_meta(params["meta_path"]) - bits_to_uV = get_bits_to_uV(meta) # convert from bits to uV - bits_to_uV = cp.float32(bits_to_uV) # convert to cupy float32 - mean_wf *= bits_to_uV + nAP = get_channel_counts(meta)[0] + bits_to_uV = get_bits_to_uV(np.arange(nAP), meta) + bits_to_uV = cp.array(bits_to_uV) + # kept sync on current mean_wf so just add 1 to bits_to_uV + bits_to_uV = cp.append(bits_to_uV, 1) + mean_wf_uV = mean_wf * bits_to_uV[cp.newaxis, :, cp.newaxis] tqdm.write("Saving mean waveforms...") - cp.save(mean_wf_path, mean_wf) + cp.save(mean_wf_path, mean_wf_uV) # Convert back to numpy arrays for compatibility - mean_wf = cp.asnumpy(mean_wf) + mean_wf_uV = cp.asnumpy(mean_wf_uV) - return mean_wf + return mean_wf_uV def calc_mean_wf_split( @@ -163,6 +181,7 @@ def calc_mean_wf_split( times_multi: dict[NDArray[np.int_]], data: NDArray[np.int_], n_splits: int = 2, + all_spikes: dict[int, NDArray] = None, ): if n_splits < 2: raise ValueError("n_splits must be at least 2. Otherwise use calc_mean_wf.") @@ -192,15 +211,17 @@ def calc_mean_wf_split( n_splits, ) ) - for i in tqdm(cluster_ids, desc="Calculating mean waveforms"): - spikes = extract_spikes( + if all_spikes is None: + all_spikes = extract_all_spikes( data, times_multi, - i, + cluster_ids, params["pre_samples"], params["post_samples"], params["max_spikes"], ) + for i in tqdm(cluster_ids, desc="Calculating mean waveforms"): + spikes = all_spikes[i] if len(spikes) > 0: # edge case spikes_cp = cp.array(spikes, dtype=cp.float32) for split in range(n_splits): @@ -212,17 +233,20 @@ def calc_mean_wf_split( # convert mean_wf uV meta = read_meta(params["meta_path"]) - bits_to_uV = get_bits_to_uV(meta) # convert from bits to uV - bits_to_uV = cp.float32(bits_to_uV) # convert to cupy float32 - mean_wf *= bits_to_uV + nAP = get_channel_counts(meta)[0] + bits_to_uV = get_bits_to_uV(np.arange(nAP), meta) + bits_to_uV = cp.array(bits_to_uV) + # kept sync on current mean_wf so just add 1 to bits_to_uV + bits_to_uV = cp.append(bits_to_uV, 1) + mean_wf_uV = mean_wf * bits_to_uV[cp.newaxis, :, cp.newaxis, cp.newaxis] tqdm.write("Saving mean waveforms...") - cp.save(mean_wf_path, mean_wf) + cp.save(mean_wf_path, mean_wf_uV) # Convert back to numpy arrays for compatibility - mean_wf = cp.asnumpy(mean_wf) + mean_wf_uV = cp.asnumpy(mean_wf_uV) - return mean_wf + return mean_wf_uV def find_times_multi_ks( @@ -233,7 +257,10 @@ def find_times_multi_ks( ): sp_times = np.load(os.path.join(ks_folder, "spike_times.npy")) sp_clust = np.load(os.path.join(ks_folder, "spike_clusters.npy")) - data = get_data_memmap(ks_folder) + meta_path = get_meta_path(ks_folder) + meta = read_meta(meta_path) + bin_path = get_binary_path(ks_folder) + data = get_data_memmap(bin_path, meta) if clust_ids is None: clust_ids = np.arange(np.max(sp_clust) + 1) diff --git a/npx_utils/footprint.py b/npx_utils/footprint.py new file mode 100644 index 0000000..9e1d1b2 --- /dev/null +++ b/npx_utils/footprint.py @@ -0,0 +1,136 @@ +import math +from pathlib import Path + +import numpy as np +from scipy import interpolate +from tqdm import tqdm + + +def get_sign(mean_wf): + troughs = np.min(mean_wf, axis=1) + peaks = np.max(mean_wf, axis=1) + sign = abs(np.min(troughs)) > abs(np.max(peaks)) + return sign + + +def get_amp_map(shank_pos, shank_amplitudes): + max_amp_idx = np.argmax(shank_amplitudes) + max_y = shank_pos[max_amp_idx, 1] + # find channels within 500 um + nearby_mask = np.abs(shank_pos[:, 1] - max_y) <= 500 + xcoords = shank_pos[nearby_mask, 0] + ycoords = shank_pos[nearby_mask, 1] + xmin, xmax = xcoords.min(), xcoords.max() + ymin, ymax = ycoords.min(), ycoords.max() + width = np.ceil(xmax - xmin + 1).astype(int) + height = np.ceil(ymax - ymin + 1).astype(int) + xcoords2 = np.round(xcoords - xmin).astype(int) + ycoords2 = np.round(ycoords - ymin).astype(int) + width1 = np.arange(0, width) + height1 = np.arange(0, height) + xx, yy = np.meshgrid(width1, height1) + if len(np.unique(xcoords)) > 1: + gd1 = interpolate.griddata( + (ycoords2, xcoords2), shank_amplitudes, (yy, xx), method="cubic" + ) + else: + spl = interpolate.interp1d( + ycoords2, shank_amplitudes, kind="cubic", fill_value="extrapolate" + ) + gd1 = spl(yy) + return gd1 + + +def polar_to_cartesian(theta, r): + xall = [] + yall = [] + for i, r1 in enumerate(r): + x = r1 * math.cos(theta[i]) + y = r1 * math.sin(theta[i]) + xall.append(x) + yall.append(y) + + return xall, yall + + +def create_sampling_matrix(pointN): + # sample amp at different radius from the peak site + # first design the sampling position matrix + rho = np.arange(0, pointN) + # Initialize theta array and set up parameters + theta = np.zeros_like(rho) + count1 = 0 + # Define beta range from -pi to pi with steps of pi/6 + beta_values = np.arange(-np.pi, np.pi, np.pi / 6) + # Initialize xq and yq as empty lists + xq = np.empty((0, pointN)) + yq = np.empty((0, pointN)) + for beta in beta_values: + x, y = polar_to_cartesian(theta + beta, rho) # Use numpy's pol2cart equivalent + xq = np.append(xq, np.array(x).reshape(1, -1), axis=0) + yq = np.append(yq, np.array(y).reshape(1, -1), axis=0) + + return xq, yq + + +def get_spatial_footprint(ks_folder, threshold=30): + "https://github.com/zhiwen10/Neuropixels-footprint/blob/main/python/get_footprint_new3.py#L74" + # make ks_folder a Path object if not already + ks_folder = Path(ks_folder) + mean_wfs = np.load(ks_folder / "mean_waveforms.npy") + channel_positions = np.load(ks_folder / "channel_positions.npy") + footprint = [] + for mean_wf in tqdm( + mean_wfs, total=len(mean_wfs), desc="Calculating spatial footprints" + ): + sign = get_sign(mean_wf) + if sign == 0: + mean_wf = -mean_wf + + trough_idx = np.argmin(mean_wf, axis=1) + peak_idx = np.argmax(mean_wf[:, trough_idx:], axis=1) + trough_idx + amplitudes = mean_wf[:, peak_idx] - mean_wf[:, trough_idx] + # loop through each shank + for shank in range(4): + shank_mask = (channel_positions[:, 0] >= shank * 250) & ( + channel_positions[:, 0] < (shank + 1) * 250 + ) + shank_pos = channel_positions[shank_mask] + shank_amplitudes = amplitudes[shank_mask] + if len(shank_amplitudes) == 0: + continue + gd1 = get_amp_map(shank_pos, shank_amplitudes) + flat_index = np.nanargmax(np.abs(gd1)) + row, col = np.unravel_index(flat_index, gd1.shape) + if len(np.unique(shank_pos[:, 0])) > 1: + pointN = 401 + xq, yq = create_sampling_matrix(pointN) + xq2 = np.round(xq + col).astype(int) + yq2 = np.round(yq + row).astype(int) + vq = np.full(xq2.shape, np.nan) + for i in range(xq2.shape[0]): + for j in range(xq2.shape[1]): + if (1 <= yq2[i, j] < gd1.shape[0]) and ( + 1 <= xq2[i, j] < gd1.shape[1] + ): + vq[i, j] = gd1[yq2[i, j], xq2[i, j]] + else: + rho = np.arange(0, 401) + yq2 = np.empty((0, 401)) + a1 = np.round(rho + row).astype(int) + yq2 = np.append(yq2, np.array(a1).reshape(1, -1), axis=0) + a2 = np.round(-rho + row).astype(int) + yq2 = np.append(yq2, np.array(a2).reshape(1, -1), axis=0) + yq2 = yq2.astype(int) + # Initialize vq with NaN values + vq = np.full(yq2.shape, np.nan) + # Iterate through each element in xq2 and yq2 + for i in range(0, 2): + for j in range(yq2.shape[1]): + if 1 <= yq2[i, j] < gd1.shape[0]: + vq[i, j] = gd1[yq2[i, j]] + vq_mean = np.nanmean(vq, axis=0) + indices = np.where(vq_mean <= threshold)[0] + ft = indices[0] if indices.size > 0 else 100 + footprint.append(ft) + return np.array(footprint) \ No newline at end of file diff --git a/npx_utils/metrics.py b/npx_utils/metrics.py index db06b89..f19fe14 100644 --- a/npx_utils/metrics.py +++ b/npx_utils/metrics.py @@ -36,51 +36,6 @@ def calc_sliding_RP_viol( return RP_viol -def _sliding_RP_viol( - correlogram, - bin_size: float = 0.25, - acceptThresh: float = 0.1, -) -> float: - """ - Calculate the sliding refractory period violation confidence for a cluster. - Args: - correlogram (NDArray): The auto-correlogram of the cluster. - bin_size (float, optional): The size of each bin in ms. Defaults to 0.25. - acceptThresh (float, optional): The threshold for accepting refractory period violations. Defaults to 0.1. - Returns: - float: The refractory period violation confidence for the cluster. - """ - # create various refractory periods sizes to test (between 0 and 20x bin size) - b = np.arange(0, 21 * bin_size, bin_size) / 1000 - bTestIdx = np.array([1, 2, 4, 6, 8, 12, 16, 20], dtype="int8") - bTest = [b[i] for i in bTestIdx] - - # calculate and avg halves of acg to ensure symmetry - # keep only second half of acg, refractory period violations are compared from the center of acg - half_len = int(correlogram.shape[0] / 2) - correlogram = (correlogram[half_len:] + correlogram[:half_len][::-1]) / 2 - - acg_cumsum = np.cumsum(correlogram) - sum_res = acg_cumsum[bTestIdx - 1] # -1 bc 0th bin corresponds to 0-bin_size ms - - # low-pass filter acg and use max as baseline event rate - order = 4 # Hz - cutoff_freq = 250 # Hz - fs = 1 / bin_size * 1000 - nyqist = fs / 2 - cutoff = cutoff_freq / nyqist - sos = butter(order, cutoff, btype="low", output="sos") - smoothed_acg = sosfiltfilt(sos, correlogram) - - bin_rate_max = np.max(smoothed_acg) - max_conts_max = np.array(bTest) / bin_size * 1000 * (bin_rate_max * acceptThresh) - # compute confidence of less than acceptThresh contamination at each refractory period - confs = 1 - stats.poisson.cdf(sum_res, max_conts_max) - rp_viol = 1 - confs.max() - - return rp_viol - - def auto_correlogram( c1_times: NDArray[np.float64], window_size: float, @@ -133,6 +88,51 @@ def x_correlogram( return _correlogram(c1_times, c2_times, window_size, bin_width, overlap_tol) +def _sliding_RP_viol( + correlogram, + bin_size: float = 0.25, + acceptThresh: float = 0.1, +) -> float: + """ + Calculate the sliding refractory period violation confidence for a cluster. + Args: + correlogram (NDArray): The auto-correlogram of the cluster. + bin_size (float, optional): The size of each bin in ms. Defaults to 0.25. + acceptThresh (float, optional): The threshold for accepting refractory period violations. Defaults to 0.1. + Returns: + float: The refractory period violation confidence for the cluster. + """ + # create various refractory periods sizes to test (between 0 and 20x bin size) + b = np.arange(0, 21 * bin_size, bin_size) / 1000 + bTestIdx = np.array([1, 2, 4, 6, 8, 12, 16, 20], dtype="int8") + bTest = [b[i] for i in bTestIdx] + + # calculate and avg halves of acg to ensure symmetry + # keep only second half of acg, refractory period violations are compared from the center of acg + half_len = int(correlogram.shape[0] / 2) + correlogram = (correlogram[half_len:] + correlogram[:half_len][::-1]) / 2 + + acg_cumsum = np.cumsum(correlogram) + sum_res = acg_cumsum[bTestIdx - 1] # -1 bc 0th bin corresponds to 0-bin_size ms + + # low-pass filter acg and use max as baseline event rate + order = 4 # Hz + cutoff_freq = 250 # Hz + fs = 1 / bin_size * 1000 + nyqist = fs / 2 + cutoff = cutoff_freq / nyqist + sos = butter(order, cutoff, btype="low", output="sos") + smoothed_acg = sosfiltfilt(sos, correlogram) + + bin_rate_max = np.max(smoothed_acg) + max_conts_max = np.array(bTest) / bin_size * 1000 * (bin_rate_max * acceptThresh) + # compute confidence of less than acceptThresh contamination at each refractory period + confs = 1 - stats.poisson.cdf(sum_res, max_conts_max) + rp_viol = 1 - confs.max() + + return rp_viol + + def _correlogram( c1_times: NDArray[np.float64], c2_times: NDArray[np.float64], diff --git a/npx_utils/noise/lfpBandPower.m b/npx_utils/noise/lfpBandPower.m index 46aae78..2b40c22 100644 --- a/npx_utils/noise/lfpBandPower.m +++ b/npx_utils/noise/lfpBandPower.m @@ -82,7 +82,7 @@ i = i + bStart; % to index into the original array % fprintf( 'Chan 0 peak index: %d, peak freq: %.2f Hz \n', i(1), F(i(1))); -[nChan,nPSDBin] = size(allPowerEst); +% [nChan,nPSDBin] = size(allPowerEst); if measChan > 3 @@ -115,8 +115,8 @@ peakInt = 0; bStart = i(currChan)-intWind; bEnd = i(currChan)+intWind; - if (bStart < 1) bStart = 1; end - if (bEnd > nPowerEst) bEnd = nPowerEst; end + if (bStart < 1), bStart = 1; end + if (bEnd > nPowerEst), bEnd = nPowerEst; end for f = bStart:bEnd peakInt = peakInt + allPowerEst(f,currChan) - backEst; end @@ -127,8 +127,8 @@ bStart = round(noiseRange(1)/hzSpan)+1; bEnd = round(noiseRange(2)/hzSpan)+1; - backSum = 0; - msgStr = sprintf( 'minPeakStart, maxPeakEnd: %d, %d\n', minPeakStart, maxPeakEnd ); + % backSum = 0; + % msgStr = sprintf( 'minPeakStart, maxPeakEnd: %d, %d\n', minPeakStart, maxPeakEnd ); %disp(msgStr) if (backSkipPeak == 1) if( bStart < minPeakStart ) && ( bEnd > maxPeakEnd ) @@ -150,7 +150,7 @@ else if( currChan == 1) - [pxx_nf, pxx_nchan] = size(allPowerEst); + [pxx_nf, ~] = size(allPowerEst); fprintf('number of bins in power spectrum: %d\n ', pxx_nf); fprintf('range for backSum: %d, %d\n ', bStart, bEnd); @@ -177,7 +177,7 @@ % end normPP = peakToPeakEst/mean(peakToPeakEst); - [filePath, currTitle, dumExt] = fileparts(lfpFilename); + [filePath, currTitle, ~] = fileparts(lfpFilename); currTitle = sprintf('%s_sh%d_b%d_t%d', currTitle, shank_index, bank_index, round(tStart,0)); titleStr = sprintf( 'Data from: %s',currTitle); diff --git a/npx_utils/noise/noise_channels.py b/npx_utils/noise/noise_channels.py new file mode 100644 index 0000000..ac5831e --- /dev/null +++ b/npx_utils/noise/noise_channels.py @@ -0,0 +1,212 @@ +import npx_utils as npx +import numpy as np +import scipy.signal + + +def detect_noise_channels( + data, + meta, + method="rms", + time_s=20, + # parameters for rms method + rms_threshold_uv=40, + # parameters for rms no spikes method + rms_no_spikes_threshold_uv=30, + spike_times=None, + # parameters for psd coherence method + psd_hf_threshold=0.2, + channel_positions=None, +): + """ + Detects noisy channels in the given data based. + + Returns: + Channels labels: 0: good, 1: dead low coherence / amplitude, 2: noisy, 3: outside of the brain + """ + method_list = ["rms", "rms_no_spikes", "psd_coherence"] + assert ( + method in method_list + ), f"Method {method} not recognized. Choose from {method_list}" + + # get random subset of data for estimation + n_samples, n_channels = data.shape + fs = npx.sglx_helpers.get_sample_rate(meta) + assert time_s * fs < n_samples, "Data too short for noise estimation" + # random start in middle 3/4 of data + start = np.random.randint(n_samples // 4, n_samples // 4 * 3 - time_s * fs) + data_sub = data[start : int(start + time_s * fs), :] + data_uv = npx.sglx_helpers.convert_data_to_uV( + data_sub, np.arange(n_channels - 1), meta + ) + + channel_labels = np.zeros(n_channels, dtype=int) + if method == "rms": + rms = np.sqrt(np.mean(data_uv**2, axis=0)) + noisy_channels = np.where(rms > rms_threshold_uv)[0] + channel_labels[noisy_channels] = 2 + + elif method == "rms_no_spikes": + assert ( + spike_times is not None + ), "spike_times must be provided for rms_no_spikes method" + # adjust spike times to subset + spike_times = ( + spike_times[(spike_times >= start) & (spike_times < start + time_s * fs)] + - start + ) + noise_data = extract_noise( + data_uv, spike_times, pre_samples=20, post_samples=62 + ) + rms = np.sqrt(np.mean(noise_data**2, axis=0)) + noisy_channels = np.where(rms > rms_no_spikes_threshold_uv)[0] + channel_labels[noisy_channels] = 2 + elif method == "psd_coherence": + assert ( + channel_positions is not None + ), "channel_positions must be provided for psd_coherence method" + # sort channels by x position then y position + sort_idx = np.lexsort((channel_positions[:, 1], channel_positions[:, 0])) + data_sorted = data_uv[:, sort_idx] + + channel_labels = psd_coherence_detection( + data_sorted, + fs, + psd_hf_threshold, + ) + # unsort channel labels + unsort_idx = np.argsort(sort_idx) + channel_labels = channel_labels[unsort_idx] + return channel_labels + + +def extract_noise(data, times, pre_samples=20, post_samples=62): + """ + Extract snippets of noise from the data. + Args: + data (NDArray): The data to extract noise from. + times (NDArray): The spike times. + post_samples (int): The number of samples after the spike time to include. + pre_samples (int): The number of samples before the spike time to include. + Returns: + NDArray: The extracted noise snippets. Shape (max_snippets, n_channels). + """ + total_samples = len(data) + signal_mask = np.zeros(total_samples, dtype=bool) + for time in times: + start_idx = max(0, time - pre_samples) + end_idx = min(total_samples, time + post_samples + 1) + signal_mask[start_idx:end_idx] = True + noise_indices = np.where(~signal_mask)[0] + noise_samples = data[noise_indices, :] + return noise_samples + + +# taken from spikeinterface +def psd_coherence_detection( + raw, + fs, + psd_hf_threshold, + dead_channel_thr=-0.5, + noisy_channel_thr=1.0, + outside_channel_thr=-0.75, + n_neighbors=11, + nyquist_threshold=0.8, + welch_window_ms=0.3, + outside_channels_location="top", +): + """ + Bad channels detection for Neuropixel probes developed by IBL + + Parameters + ---------- + raw : traces + (num_samples, n_channels) raw traces + fs : float + sampling frequency + psd_hf_threshold : float + Threshold for high frequency PSD. If mean PSD above `nyquist_threshold` * fn is greater than this + value, channels are flagged as noisy (together with channel coherence condition). + dead_channel_thr : float, default: -0.5 + Threshold for channel coherence below which channels are labeled as dead + noisy_channel_thr : float, default: 1 + Threshold for channel coherence above which channels are labeled as noisy (together with psd condition) + outside_channel_thr : float, default: -0.75 + Threshold for channel coherence above which channels + n_neighbors : int, default: 11 + Number of neighbors to compute median fitler + nyquist_threshold : float, default: 0.8 + Threshold on Nyquist frequency to calculate HF noise band + welch_window_ms : float, default: 0.3 + Window size for the scipy.signal.welch that will be converted to nperseg + outside_channels_location : "top" | "bottom" | "both", default: "top" + Location of the outside channels. If "top", only the channels at the top of the probe can be + marked as outside channels. If "bottom", only the channels at the bottom of the probe can be + marked as outside channels. If "both", both the channels at the top and bottom of the probe can be + marked as outside channels + + Returns + ------- + 1d array + Channels labels: 0: good, 1: dead low coherence / amplitude, 2: noisy, 3: outside of the brain + """ + _, nc = raw.shape + raw = raw - np.mean(raw, axis=0)[np.newaxis, :] + nperseg = int(welch_window_ms * fs / 1000) + fscale, psd = scipy.signal.welch(raw, fs=fs, axis=0, window="hann", nperseg=nperseg) + + # compute similarities + ref = np.median(raw, axis=1) + xcorr = np.sum(raw * ref[:, np.newaxis], axis=0) / np.sum(ref**2) + + # compute coherence + xcorr_neighbors = detrend(xcorr, n_neighbors) + xcorr_distant = xcorr - detrend(xcorr, n_neighbors) - 1 + + # make recommendation + psd_hf = np.mean(psd[fscale > (fs / 2 * nyquist_threshold), :], axis=0) + + ichannels = np.zeros(nc, dtype=int) + idead = np.where(xcorr_neighbors < dead_channel_thr)[0] + inoisy = np.where( + np.logical_or(psd_hf > psd_hf_threshold, xcorr_neighbors > noisy_channel_thr) + )[0] + + ichannels[idead] = 1 + ichannels[inoisy] = 2 + + # the channels outside of the brains are the contiguous channels below the threshold on the trend coherency + # the chanels outside need to be at the extreme of the probe + (ioutside,) = np.where(xcorr_distant < outside_channel_thr) + a = np.cumsum(np.r_[0, np.diff(ioutside) - 1]) + if ioutside.size > 0: + if outside_channels_location == "top": + # channels are sorted bottom to top, so the last channel needs to be (nc - 1) + if ioutside[-1] == (nc - 1): + ioutside = ioutside[(a == np.max(a)) & (a > 0)] + ichannels[ioutside] = 3 + elif outside_channels_location == "bottom": + # outside channels are at the bottom of the probe, so the first channel needs to be 0 + if ioutside[0] == 0: + ioutside = ioutside[(a == np.min(a)) & (a < np.max(a))] + ichannels[ioutside] = 3 + else: # both extremes are considered + if ioutside[-1] == (nc - 1) or ioutside[0] == 0: + ioutside = ioutside[(a == np.max(a)) | (a == np.min(a))] + ichannels[ioutside] = 3 + + return ichannels + + +def detrend(x, nmed): + """ + Subtract the trend from a vector + The trend is a median filtered version of the said vector with tapering + + x: input vector + nmed : number of points of the median filter + """ + ntap = int(np.ceil(nmed / 2)) + xf = np.r_[np.zeros(ntap) + x[0], x, np.zeros(ntap) + x[-1]] + + xf = scipy.signal.medfilt(xf, nmed)[ntap:-ntap] + return x - xf diff --git a/npx_utils/noise/noise_spectra.py b/npx_utils/noise/noise_spectra.py index 494020c..95a85d0 100644 --- a/npx_utils/noise/noise_spectra.py +++ b/npx_utils/noise/noise_spectra.py @@ -3,15 +3,16 @@ from tkinter import Tk, filedialog import matplotlib.pyplot as plt -import npx_utils as npx import numpy as np from scipy.signal import welch +import npx_utils as npx + def _my_time_power_spectrum_mat(x, fs): """ - Helper function to calculate the power spectrum using Welch's method. - \ """ + Helper function to calculate the power spectrum using Welch's method. + """ L = x.shape[0] # get the next power of 2 nfft = 1 << (L - 1).bit_length() @@ -292,13 +293,13 @@ def lfp_band_power(run_dir=None, stream="ap", prb_ind=0, peak_pos_hz=1000, plot= print(f"Meta file {meta_file} does not exist. Exiting.") return meta = npx.read_meta(meta_file) - bank_times = npx.get_svy_bank_times(meta) + bank_times = npx.sglx_helpers.get_svy_bank_times(meta) - fs = npx.get_sample_rate(meta) + fs = npx.sglx_helpers.get_sample_rate(meta) n_chans_in_file = int(meta["nSavedChans"]) if meta["typeThis"] == "imec": - [ap_gain, lf_gain, _, _] = npx.get_chan_gains_imec(meta) - [ap, _, _] = npx.get_chan_counts_imec(meta) + [ap_gain, lf_gain, _, _] = npx.sglx.sglx_helpers.ChanGainsIM(meta) + [ap, _, _] = npx.sglx_helpers.get_channel_counts(meta) n_chans = ap if stream == "ap": gain = ap_gain[0] @@ -314,8 +315,8 @@ def lfp_band_power(run_dir=None, stream="ap", prb_ind=0, peak_pos_hz=1000, plot= freq_bands = [[0.5, 1000]] noise_range = (0.5, 1000) elif meta["typeThis"] == "nidq": - gain = npx.get_chan_gains_ni(meta) - [_, _, n_chans, _] = npx.get_chan_counts_ni(meta) + gain = npx.sglx_helpers.ChanGainNI(meta) + [_, _, n_chans, _] = npx.sglx_helpers.get_channel_counts(meta) skip_chan = [] disp_range = (0.5, 15000) marginal_chans = [0] @@ -323,7 +324,7 @@ def lfp_band_power(run_dir=None, stream="ap", prb_ind=0, peak_pos_hz=1000, plot= noise_range = (0.001, 500) elif meta["typeThis"] == "obx": gain = 1 - [_, n_chans, _] = npx.get_chan_counts_obx(meta) + [_, n_chans, _] = npx.sglx_helpers.get_channel_counts(meta) skip_chan = [] disp_range = (0.5, 15000) marginal_chans = [0] @@ -332,14 +333,12 @@ def lfp_band_power(run_dir=None, stream="ap", prb_ind=0, peak_pos_hz=1000, plot= else: raise ValueError("Unknown stream type in metadata.") - tomV = 1e3 * (npx.int2volts(meta) / gain) + tomV = 1e3 * (npx.sglx_helpers.Int2Volts(meta) / gain) # Get disabled channels from the metadata skip_chan = [] if meta["typeThis"] == "imec": - # Use regex to parse the shank map string, similar to MATLAB's textscan shank_map_str = meta.get("snsShankMap", meta.get("snsGeomMap")) - # Matches tuples like (c:h:a:n) and captures the 4th number (enabled flag) matches = re.findall(r"\(\d+:\d+:\d+:(\d)\)", shank_map_str) if matches: enabled = np.array([int(m) for m in matches]) @@ -349,7 +348,6 @@ def lfp_band_power(run_dir=None, stream="ap", prb_ind=0, peak_pos_hz=1000, plot= nclips = 3 n_banks = 1 if type(bank_times) == int else len(bank_times) - # --- 4. Run Analysis --- if n_banks > 1: for k in range(n_banks): t_start = bank_times[k, 3] + 200 @@ -445,20 +443,20 @@ def lfp_band_power(run_dir=None, stream="ap", prb_ind=0, peak_pos_hz=1000, plot= # ax2.set_xlim(disp_range) # plt.show() + else: + fig = None return ptp_mean, ptp_std, noise_mean, noise_std, fig -import os -import re - -from lfp_power import lfp_band_power - def find_match_dir(search_path, pattern): """ Finds subdirectories within a given path that match a regex pattern. """ pattern = re.compile(pattern) + # check if search_path matches the pattern itself + if pattern.search(search_path): + return [search_path] subdirs = [] for dirpath, dirnames, filenames in os.walk(search_path): for dirname in dirnames: @@ -472,21 +470,21 @@ def run_noise(base_path, stream="ap", peak_pos_hz=1000, overwrite=False, plot=Tr Main function to run the batch analysis. """ # --- Start Processing --- - # Find all run folders (e.g., matching '_g' followed by digits) + # Find all run folders (e.g., matching '_g' followed by digits), including the base path itself run_folders = find_match_dir(base_path, r"_g\d+$") for run_name in run_folders: run_dir = os.path.join(base_path, run_name) print(f"Processing: {run_dir}") - # Find all probe folders within the run folder (e.g., matching '_imec' followed by digits) + # Find all probe folders probe_folders = find_match_dir(run_dir, r"_imec\d+$") for probe_dir_name in probe_folders: - # Extract the probe number using a regex group + # Extract the probe number match = re.search(r"_imec(\d+)$", probe_dir_name) if not match: - continue # Skip if the folder name is malformed + continue probe_num = int(match.group(1)) print(f"\tProbe {probe_num}") @@ -499,24 +497,24 @@ def run_noise(base_path, stream="ap", peak_pos_hz=1000, overwrite=False, plot=Tr print(f"\t\tOutput exists. Skipping.") continue - # Call the main analysis function - # try: (_, _, noise_mean, noise_std, fig1) = lfp_band_power( - run_dir, stream, probe_num, peak_pos_hz, True + run_dir, stream, probe_num, peak_pos_hz, plot ) # Save the numerical results to a text file - # The 'with' statement automatically handles closing the file with open(output_path, "w") as f: result_str = f"noise mean & std: {noise_mean:.2f}, {noise_std:.2f}\n" print(f"\t\t{result_str.strip()}") f.write(result_str) - # Save the figure generated by the analysis function - fig_path = os.path.join(run_dir, f"{probe_dir_name}_spectra.png") + if plot: + fig_path = os.path.join(run_dir, f"{probe_dir_name}_spectra.png") + fig1.savefig(fig_path, dpi=300) + print(f"\t\tSaved results to {output_name} and plot {fig_path}.") + else: + print(f"\t\tSaved results to {output_name}.") - fig1.savefig(fig_path, dpi=300) - print(f"\t\tSaved results to {output_name} and spectra PNG.") - # except Exception as e: - # print(f"\t\tERROR processing {probe_dir_name}: {e}") +if __name__ == "__main__": + base_path = r"Z:\SalineTests\23107805413\20250818_23107805413_saline_ext_g0" + run_noise(base_path, stream="ap", peak_pos_hz=1000, overwrite=True, plot=True) diff --git a/npx_utils/other_helpers.py b/npx_utils/other_helpers.py index 6366be9..092048c 100644 --- a/npx_utils/other_helpers.py +++ b/npx_utils/other_helpers.py @@ -4,8 +4,10 @@ from tqdm import tqdm +import npx_utils as npx + from .ks_helpers import get_meta_path, get_probe_id -from .sglx_helpers import read_meta +from .sglx.sglx_helpers import read_meta def copy_folder_with_progress(src, dest, overwrite=False): @@ -51,7 +53,8 @@ def get_probe_folders(ks_folders): if probe_num not in probe_folders: probe_folders[probe_num] = [] probe_folders[probe_num].append(ks_folder) - return probe_folders + sorted_dict = dict(sorted(probe_folders.items())) + return sorted_dict def get_details(ks_folder, drug_dict=None): @@ -80,3 +83,60 @@ def get_details(ks_folder, drug_dict=None): "drug": drug, } return details + + +def get_run_folders(subject_folder, day_folders=None): + if day_folders is None: + day_folders = [ + os.path.join(subject_folder, folder) + for folder in os.listdir(subject_folder) + if os.path.isdir(os.path.join(subject_folder, folder)) + and ("SvyPrb" not in folder) + and ("old" not in folder) + ] + else: + day_folders = [ + folder if os.path.isabs(folder) else os.path.join(subject_folder, folder) + for folder in day_folders + ] + + run_folders = [] + for day_folder in day_folders: + possible_run_folders = [ + os.path.join(day_folder, folder) + for folder in os.listdir(day_folder) + if npx.is_run_folder(folder) + ] + # find one with supercat + supercat_folders = [ + folder for folder in possible_run_folders if "supercat" in folder + ] + if len(supercat_folders) > 1: + raise ValueError( + f"Multiple supercat folders found in {day_folder}: {supercat_folders}" + ) + if len(supercat_folders) == 1: + run_folders.append(supercat_folders[0]) + else: + # find catgt folders + catgt_folders = [ + folder for folder in possible_run_folders if "catgt" in folder + ] + if len(catgt_folders) > 1: + raise ValueError( + f"Multiple catgt folders found in {day_folder}: {catgt_folders}" + ) + if len(catgt_folders) == 1: + run_folders.append(catgt_folders[0]) + else: + if len(possible_run_folders) == 1: + run_folders.append(possible_run_folders[0]) + elif len(possible_run_folders) == 0: + continue + else: + raise ValueError( + f"Multiple run folders found in {day_folder}: {possible_run_folders}" + ) + run_folders.sort() + return run_folders + diff --git a/npx_utils/plotting_helpers.py b/npx_utils/plotting_helpers.py new file mode 100644 index 0000000..f8fb4f2 --- /dev/null +++ b/npx_utils/plotting_helpers.py @@ -0,0 +1,216 @@ +import matplotlib.patches as patches +import matplotlib.pyplot as plt +import numpy as np + + +def plot_np2(channel_positions=None): + num_shanks = 4 + tot_sites = 5120 + shank_pitch = 250 + shank_width = 70 + shank_length = 10000 + tip_length = 175 + site_size = (12, 12) + horizontal_site_pitch = 32 + vertical_site_pitch = 15 + bank_size = 2700 + + num_sites_per_shank = tot_sites // num_shanks + if channel_positions is not None: + # Round coordinates and convert to a set of tuples for very fast searching + highlight_coords = set(map(tuple, np.round(channel_positions).astype(int))) + else: + highlight_coords = set() + + fig, ax = plt.subplots(figsize=(8, 12)) + for i in range(num_shanks): + # DRAW SHANK + shank_start_x = i * shank_pitch + shank_color = "#bdc3c7" + + body = patches.Rectangle( + (shank_start_x, tip_length), + shank_width, + shank_length - tip_length, + facecolor=shank_color, + edgecolor="black", + ) + ax.add_patch(body) + + tip = patches.Polygon( + [ + (shank_start_x, tip_length), + (shank_start_x + shank_width, tip_length), + (shank_start_x + shank_width / 2, 0), + ], + facecolor=shank_color, + edgecolor="black", + ) + ax.add_patch(tip) + + # DRAW RECORDING SITES + shank_center_x = shank_start_x + shank_width / 2 + col_x_coords = [ + shank_center_x - horizontal_site_pitch / 2 + 8, + shank_center_x + horizontal_site_pitch / 2 + 8, + ] + + num_sites_per_column = num_sites_per_shank // 2 + + for j in range(num_sites_per_column): + # Calculate the y-position, same for both columns in a given row + site_y = tip_length + (j * vertical_site_pitch) + + # Stop drawing if sites go beyond the shank's physical length + if (site_y + site_size[1]) > shank_length: + break + + # Draw site in the first column + site1_x = col_x_coords[0] - site_size[0] / 2 + site1_center_x = col_x_coords[0] + site1_center_y = site_y - tip_length + site1_coords = (int(round(site1_center_x)), int(round(site1_center_y))) + facecolor1 = "lime" if site1_coords in highlight_coords else "none" + + site1 = patches.Rectangle( + (site1_x, site_y), + site_size[0], + site_size[1], + facecolor=facecolor1, + edgecolor="k", + linewidth=0.2, + ) + ax.add_patch(site1) + + # Draw site in the second column + site2_x = col_x_coords[1] - site_size[0] / 2 + site2_center_x = col_x_coords[1] + site2_center_y = site_y - tip_length + site2_coords = (int(round(site2_center_x)), int(round(site2_center_y))) + facecolor2 = "lime" if site2_coords in highlight_coords else "none" + site2 = patches.Rectangle( + (site2_x, site_y), + site_size[0], + site_size[1], + facecolor=facecolor2, + edgecolor="k", + linewidth=0.2, + ) + ax.add_patch(site2) + + # DRAW BANKS + # y_start_sites = tip_length + # num_banks = int(np.ceil((shank_length - y_start_sites) / bank_size)) + + # for i in range(num_banks): + # bank_top_y = y_start_sites + ((i + 1) * bank_size) + + # boundary_color = "#ff9393" + # if bank_top_y < shank_length: + # ax.axhline( + # y=bank_top_y, color=boundary_color, linestyle="--", linewidth=0.5 + # ) + + # 5. Finalize and style the plot + ax.set_xlabel("Distance (µm)", fontsize=12) + ax.set_ylabel("Distance from Tip (µm)", fontsize=12) + ax.set_xlim(-shank_pitch / 2, (num_shanks - 0.5) * shank_pitch) + ax.set_ylim(-tip_length, shank_length * 1.05) + + fig.tight_layout() + # plt.show() + + +def plot_peak_heatmap(channel_positions, peak_channels): + """ + Plots the NP 2.0 probe geometry with heatmap showing locations of neurons based on peak channel + + Args: + channel_positions (np.ndarray): Array of shape (num_channels, 2) with [x, z] coordinates of channels. + peak_channels (np.ndarray): 1D array of channel indices for each neuron's peak. + """ + unique_channels, counts = np.unique(peak_channels, return_counts=True) + peak_counts = dict(zip(unique_channels, counts)) + max_count = max(peak_counts.values()) + cmap = plt.get_cmap("viridis") + norm = plt.Normalize(vmin=0, vmax=max_count) + + num_shanks = 4 + shank_pitch = 250 + shank_width = 70 + shank_length = 10000 + tip_length = 175 + site_size = (12, 12) + + # 4. Setup plot + fig, ax = plt.subplots(figsize=(8, 50)) + ax.set_title("Neuron Peak Channel Density", fontsize=16, pad=20) + + # 5. Draw shanks + for i in range(num_shanks): + shank_start_x = i * shank_pitch + shank_color = "#bdc3c7" + # Draw body and tip + body = patches.Rectangle( + (shank_start_x, tip_length), + shank_width, + shank_length - tip_length, + facecolor=shank_color, + edgecolor="black", + zorder=1, + ) + ax.add_patch(body) + tip = patches.Polygon( + [ + (shank_start_x, tip_length), + (shank_start_x + shank_width, tip_length), + (shank_start_x + shank_width / 2, 0), + ], + facecolor=shank_color, + edgecolor="black", + zorder=1, + ) + ax.add_patch(tip) + + # 6. Draw sites as a heatmap + for channel_id, (x_pos, z_pos) in enumerate(channel_positions): + count = peak_counts.get(channel_id, 0) + + # Determine color and style based on count + if count > 0: + face_color = cmap(norm(count)) + edge_color = "k" + line_width = 0.2 + else: + face_color = "none" + edge_color = "#d3d3d3" # Faint grey for unused sites + line_width = 0.1 + + # The channel_positions gives the center, so calculate the bottom-left corner + bottom_left_x = x_pos - site_size[0] / 2 + bottom_left_y = z_pos + tip_length + + site = patches.Rectangle( + (bottom_left_x, bottom_left_y), + site_size[0], + site_size[1], + facecolor=face_color, + edgecolor=edge_color, + linewidth=line_width, + zorder=2, # Draw sites on top of the shank + ) + ax.add_patch(site) + + # 7. Add a colorbar + sm = plt.cm.ScalarMappable(cmap=cmap, norm=norm) + sm.set_array([]) + cbar = fig.colorbar(sm, ax=ax, pad=0.02, aspect=30, shrink=0.8) + cbar.set_label("Number of Neurons", rotation=270, labelpad=15, fontsize=12) + + # 8. Finalize and style the plot + ax.set_xlabel("Distance (µm)", fontsize=12) + ax.set_ylabel("Depth (µm)", fontsize=12) + ax.set_xlim(-shank_pitch / 2, (num_shanks - 0.5) * shank_pitch) + ax.set_ylim(-tip_length, shank_length * 1.05) + fig.tight_layout() + plt.show() diff --git a/npx_utils/sglx/_SGLXMetaToCoords.py b/npx_utils/sglx/_SGLXMetaToCoords.py new file mode 100644 index 0000000..dc14449 --- /dev/null +++ b/npx_utils/sglx/_SGLXMetaToCoords.py @@ -0,0 +1,770 @@ +# -*- coding: utf-8 -*- +""" +Requires python 3 + +The main() function at the bottom of this file can run from an +interpreter, or, the helper functions can be imported into a +new module or Jupyter notebook. + +Standalone program to generate a coordinate file from a SpikeGLX +metadata file for 3A, NP1.0, or NP2 single shank or multishank probes. + +The output is set by the outType parameter: + 0 for text coordinate file; + 1 for Kilosort or Kilosort2 channel map file; + 2 for strings to paste into JRClust .prm file + 3 add fields to metadata that were missing prior to SpikeGLX version 032623-phase30 + 4 npy file with (nchan,2) matrix of xy coordinates, e.g. for YASS and related sorters + + +@author: Jennifer Colonell, Janelia Research Campus + +""" + +import shutil +from pathlib import Path +from tkinter import Tk, filedialog + +import matplotlib.pyplot as plt +import numpy as np +import scipy.io + + +# ========================================================= +# Parse ini file returning a dictionary whose keys are the metadata +# left-hand-side-tags, and values are string versions of the right-hand-side +# metadata values. We remove any leading '~' characters in the tags to match +# the MATLAB version of readMeta. +# +# The string values are converted to numbers using the "int" and "float" +# fucntions. Note that python 3 has no size limit for integers. +# +def readMeta(metaPath): + metaDict = {} + # if metaPath is not pathlib, convert + if type(metaPath) == str: + metaPath = Path(metaPath) + if metaPath.exists(): + # print("meta file present") + with metaPath.open() as f: + mdatList = f.read().splitlines() + # convert the list entries into key value pairs + for m in mdatList: + csList = m.split(sep="=") + if csList[0][0] == "~": + currKey = csList[0][1 : len(csList[0])] + else: + currKey = csList[0] + metaDict.update({currKey: csList[1]}) + else: + print("no meta file") + + return metaDict + + +# ========================================================= +# Return counts of each imec channel type that composes the timepoints +# stored in the binary files. +# +def ChannelCountsIM(meta): + chanCountList = meta["snsApLfSy"].split(sep=",") + AP = int(chanCountList[0]) + LF = int(chanCountList[1]) + SY = int(chanCountList[2]) + + return (AP, LF, SY) + + +# ========================================================= +# Return geometry paramters for supported probe types +# These are used to calculate positions from metadata +# that includes only ~snsShankMap +# +def getGeomParams(meta): + # many part numbers have the same geometry parameters ; + # define those sets in lists + # [nShank, shankWidth, shankPitch, even_xOff, odd_xOff, horizPitch, vertPitch, rowsPerShank, elecPerShank] + # offset and pitch values in um + np1_stag_70um = [1, 70, 0, 27, 11, 32, 20, 480, 960] + nhp_lin_70um = [1, 70, 0, 27, 27, 32, 20, 480, 960] + nhp_stag_125um_med = [1, 125, 0, 27, 11, 87, 20, 1368, 2496] + nhp_stag_125um_long = [1, 125, 0, 27, 11, 87, 20, 2208, 4416] + nhp_lin_125um_med = [1, 125, 0, 11, 11, 103, 20, 1368, 2496] + nhp_lin_125um_long = [1, 125, 0, 11, 11, 103, 20, 2208, 4416] + uhd_8col_1bank = [1, 70, 0, 14, 14, 6, 6, 48, 384] + uhd_8col_16bank = [1, 70, 0, 14, 14, 6, 6, 768, 6144] + np2_ss = [1, 70, 0, 27, 27, 32, 15, 640, 1280] + np2_4s = [4, 70, 250, 27, 27, 32, 15, 640, 1280] + NP1120 = [1, 70, 0, 6.75, 6.75, 4.5, 4.5, 192, 384] + NP1121 = [1, 70, 0, 6.25, 6.25, 3, 3, 384, 384] + NP1122 = [1, 70, 0, 6.75, 6.75, 4.5, 4.5, 24, 384] + NP1123 = [1, 70, 0, 10.25, 10.25, 3, 3, 32, 384] + NP1300 = [1, 70, 0, 11, 11, 48, 20, 480, 960] + NP1200 = [1, 70, 0, 27, 11, 32, 20, 64, 128] + NXT3000 = [1, 70, 0, 53, 53, 0, 15, 128, 128] + + M = dict( + [ + ("3A", np1_stag_70um), + ("PRB_1_4_0480_1", np1_stag_70um), + ("PRB_1_4_0480_1_C", np1_stag_70um), + ("NP1010", np1_stag_70um), + ("NP1011", np1_stag_70um), + ("NP1012", np1_stag_70um), + ("NP1013", np1_stag_70um), + ("NP1015", nhp_lin_70um), + ("NP1015", nhp_lin_70um), + ("NP1016", nhp_lin_70um), + ("NP1017", nhp_lin_70um), + ("NP1020", nhp_stag_125um_med), + ("NP1021", nhp_stag_125um_med), + ("NP1030", nhp_stag_125um_long), + ("NP1031", nhp_stag_125um_long), + ("NP1022", nhp_lin_125um_med), + ("NP1032", nhp_lin_125um_long), + ("NP1100", uhd_8col_1bank), + ("NP1110", uhd_8col_16bank), + ("PRB2_1_4_0480_1", np2_ss), + ("PRB2_1_2_0640_0", np2_ss), + ("NP2000", np2_ss), + ("NP2003", np2_ss), + ("NP2004", np2_ss), + ("PRB2_4_2_0640_0", np2_4s), + ("PRB2_4_4_0480_1", np2_4s), + ("NP2010", np2_4s), + ("NP2013", np2_4s), + ("NP2014", np2_4s), + ("NP1120", NP1120), + ("NP1121", NP1121), + ("NP1122", NP1122), + ("NP1123", NP1123), + ("NP1300", NP1300), + ("NP1200", NP1200), + ("NXT3000", NXT3000), + ] + ) + + # get probe part number; if absent, this is a 3A + if "imDatPrb_pn" in meta: + pn = meta["imDatPrb_pn"] + else: + pn = "3A" + + if pn in M: + geomList = M[pn] + else: + print("unsupported probe part number\n") + geomList = [] + + return geomList + + +# ========================================================= +# Return full MUX table string to append to metadata +# +def getMuxTable(meta): + # Read probe part number from meta + # Return full MUX table string to append to metadata + # As of 032923, there are 4 mux tables + + np1 = r"~muxTbl=(32,12)(0 1 24 25 48 49 72 73 96 97 120 121 144 145 168 169 192 193 216 217 240 241 264 265 288 289 312 313 336 337 360 361)(2 3 26 27 50 51 74 75 98 99 122 123 146 147 170 171 194 195 218 219 242 243 266 267 290 291 314 315 338 339 362 363)(4 5 28 29 52 53 76 77 100 101 124 125 148 149 172 173 196 197 220 221 244 245 268 269 292 293 316 317 340 341 364 365)(6 7 30 31 54 55 78 79 102 103 126 127 150 151 174 175 198 199 222 223 246 247 270 271 294 295 318 319 342 343 366 367)(8 9 32 33 56 57 80 81 104 105 128 129 152 153 176 177 200 201 224 225 248 249 272 273 296 297 320 321 344 345 368 369)(10 11 34 35 58 59 82 83 106 107 130 131 154 155 178 179 202 203 226 227 250 251 274 275 298 299 322 323 346 347 370 371)(12 13 36 37 60 61 84 85 108 109 132 133 156 157 180 181 204 205 228 229 252 253 276 277 300 301 324 325 348 349 372 373)(14 15 38 39 62 63 86 87 110 111 134 135 158 159 182 183 206 207 230 231 254 255 278 279 302 303 326 327 350 351 374 375)(16 17 40 41 64 65 88 89 112 113 136 137 160 161 184 185 208 209 232 233 256 257 280 281 304 305 328 329 352 353 376 377)(18 19 42 43 66 67 90 91 114 115 138 139 162 163 186 187 210 211 234 235 258 259 282 283 306 307 330 331 354 355 378 379)(20 21 44 45 68 69 92 93 116 117 140 141 164 165 188 189 212 213 236 237 260 261 284 285 308 309 332 333 356 357 380 381)(22 23 46 47 70 71 94 95 118 119 142 143 166 167 190 191 214 215 238 239 262 263 286 287 310 311 334 335 358 359 382 383)" + np2 = r"~muxTbl=(24,16)(0 1 32 33 64 65 96 97 128 129 160 161 192 193 224 225 256 257 288 289 320 321 352 353)(2 3 34 35 66 67 98 99 130 131 162 163 194 195 226 227 258 259 290 291 322 323 354 355)(4 5 36 37 68 69 100 101 132 133 164 165 196 197 228 229 260 261 292 293 324 325 356 357)(6 7 38 39 70 71 102 103 134 135 166 167 198 199 230 231 262 263 294 295 326 327 358 359)(8 9 40 41 72 73 104 105 136 137 168 169 200 201 232 233 264 265 296 297 328 329 360 361)(10 11 42 43 74 75 106 107 138 139 170 171 202 203 234 235 266 267 298 299 330 331 362 363)(12 13 44 45 76 77 108 109 140 141 172 173 204 205 236 237 268 269 300 301 332 333 364 365)(14 15 46 47 78 79 110 111 142 143 174 175 206 207 238 239 270 271 302 303 334 335 366 367)(16 17 48 49 80 81 112 113 144 145 176 177 208 209 240 241 272 273 304 305 336 337 368 369)(18 19 50 51 82 83 114 115 146 147 178 179 210 211 242 243 274 275 306 307 338 339 370 371)(20 21 52 53 84 85 116 117 148 149 180 181 212 213 244 245 276 277 308 309 340 341 372 373)(22 23 54 55 86 87 118 119 150 151 182 183 214 215 246 247 278 279 310 311 342 343 374 375)(24 25 56 57 88 89 120 121 152 153 184 185 216 217 248 249 280 281 312 313 344 345 376 377)(26 27 58 59 90 91 122 123 154 155 186 187 218 219 250 251 282 283 314 315 346 347 378 379)(28 29 60 61 92 93 124 125 156 157 188 189 220 221 252 253 284 285 316 317 348 349 380 381)(30 31 62 63 94 95 126 127 158 159 190 191 222 223 254 255 286 287 318 319 350 351 382 383)" + np1100 = r"~muxTbl=(32,12)(0 1 24 25 48 49 72 73 96 97 120 121 144 145 168 169 192 193 216 217 240 241 264 265 288 289 312 313 336 337 360 361)(2 3 26 27 50 51 74 75 98 99 122 123 146 147 170 171 194 195 218 219 242 243 266 267 290 291 314 315 338 339 362 363)(4 5 28 29 52 53 76 77 100 101 124 125 148 149 172 173 196 197 220 221 244 245 268 269 292 293 316 317 340 341 364 365)(6 7 30 31 54 55 78 79 102 103 126 127 150 151 174 175 198 199 222 223 246 247 270 271 294 295 318 319 342 343 366 367)(8 9 32 33 56 57 80 81 104 105 128 129 152 153 176 177 200 201 224 225 248 249 272 273 296 297 320 321 344 345 368 369)(10 11 34 35 58 59 82 83 106 107 130 131 154 155 178 179 202 203 226 227 250 251 274 275 298 299 322 323 346 347 370 371)(12 13 36 37 60 61 84 85 108 109 132 133 156 157 180 181 204 205 228 229 252 253 276 277 300 301 324 325 348 349 372 373)(14 15 38 39 62 63 86 87 110 111 134 135 158 159 182 183 206 207 230 231 254 255 278 279 302 303 326 327 350 351 374 375)(16 17 40 41 64 65 88 89 112 113 136 137 160 161 184 185 208 209 232 233 256 257 280 281 304 305 328 329 352 353 376 377)(18 19 42 43 66 67 90 91 114 115 138 139 162 163 186 187 210 211 234 235 258 259 282 283 306 307 330 331 354 355 378 379)(20 21 44 45 68 69 92 93 116 117 140 141 164 165 188 189 212 213 236 237 260 261 284 285 308 309 332 333 356 357 380 381)(22 23 46 47 70 71 94 95 118 119 142 143 166 167 190 191 214 215 238 239 262 263 286 287 310 311 334 335 358 359 382 383)" + np128ch = r"~muxTbl=(12,12)(84 11 85 5 74 10 56 112 46 121 39 127)(100 26 110 33 69 24 63 109 45 93 25 99)(87 0 82 6 71 15 53 117 43 122 42 116)(102 28 81 34 70 18 60 103 17 94 27 101)(73 1 86 7 68 16 50 106 40 123 128 128)(105 29 75 35 67 12 54 89 20 95 128 128)(76 2 83 8 65 13 47 118 49 124 128 128)(108 30 78 36 64 14 51 90 23 96 128 128)(79 3 80 9 62 114 44 119 52 125 128 128)(104 31 72 37 61 113 57 91 19 97 128 128)(88 4 77 21 59 111 41 120 55 126 128 128)(107 32 66 38 58 115 48 92 22 98 128 128)" + + M = dict( + [ + ("3A", np1), + ("PRB_1_4_0480_1", np1), + ("PRB_1_4_0480_1_C", np1), + ("NP1010", np1), + ("NP1011", np1), + ("NP1012", np1), + ("NP1013", np1), + ("NP1015", np1), + ("NP1015", np1), + ("NP1016", np1), + ("NP1017", np1), + ("NP1020", np1), + ("NP1021", np1), + ("NP1030", np1), + ("NP1031", np1), + ("NP1022", np1), + ("NP1032", np1), + ("NP1100", np1), + ("NP1110", np1100), + ("PRB2_1_4_0480_1", np2), + ("PRB2_1_2_0640_0", np2), + ("NP2000", np2), + ("NP2003", np2), + ("NP2004", np2), + ("PRB2_4_2_0640_0", np2), + ("PRB2_4_4_0480_1", np2), + ("NP2010", np2), + ("NP2013", np2), + ("NP2014", np2), + ("NP1120", np1), + ("NP1121", np1), + ("NP1122", np1), + ("NP1123", np1), + ("NP1300", np1), + ("NP1200", np128ch), + ("NXT3000", np128ch), + ] + ) + + # get probe part number; if absent, this is a 3A + if "imDatPrb_pn" in meta: + pn = meta["imDatPrb_pn"] + else: + pn = "3A" + + if pn in M: + muxTableStr = M[pn] + "\n" + else: + print("unsupported probe part number\n") + muxTableStr = [] + + return muxTableStr + + +# ========================================================= +# Parse imro table to extract 'new' metadata items +# (SpikeGLX 032623-phase 30 and later) +# Returns strings for: +# imChan0apGain +# imChan0lfGain +# imAnyChanFullBand +# +# +def imroMetaItems(meta): + + # read in the imro table + imroTbl = meta["imroTbl"].split(sep=")") + + # there is an entry in the map for each saved channel + # number of entries in map -- subtract 1 for header, one for trailing ')' + nEntry = len(imroTbl) - 2 + + # read header to get what type of imro this is + headEntry = imroTbl[0] + headEntry = headEntry[1 : len(headEntry)] # remove leading '(' + headList = headEntry.split(sep=",") + currType = headList[0] + + if int(currType) > 50000: + # this is a 3A probe + currType = "0" + + if currType == "24" or currType == "21": + imChan0apGainStr = "imChan0apGain=80" + imChan0lfGainStr = "imChan0lfGain=80" + imAnyChanFullBandStr = "imAnyChanFullBand=true" + + elif currType == "0": + # parse first entry of the table for lf and apgain + currEntry = imroTbl[1] + currEntry = currEntry[1 : len(currEntry)] # remove leading '(' + currList = currEntry.split(sep=" ") + imChan0apGainStr = "imChan0apGain=" + currList[3] + imChan0lfGainStr = "imChan0lfGain=" + currList[4] + + imAnyChanFullBandStr = "imAnyChanFullBand=false" + if len(currList) == 6: # indicates this is not a 3A imro table + # check for any full band + for i in range(nEntry): + # check for any channels where the AP filter was not used + currEntry = imroTbl[i + 1] + currEntry = currEntry[1 : len(currEntry)] + currList = currEntry.split(sep=" ") + if currList[5] == "0": + imAnyChanFullBandStr = "imAnyChanFullBand=true" + break + + elif currType == "1110": + # ap and lf gain and filter option are in the header + imChan0apGainStr = "imChan0apGain=" + headList[3] + imChan0lfGainStr = "imChan0lfGain=" + headList[4] + imAnyChanFullBandStr = "imAnyChanFullBand=false" + if headList[5] == 0: + imAnyChanFullBandStr = "imAnyChanFullBand=true" + + # add newline characters to each + imChan0apGainStr = imChan0apGainStr + "\n" + imChan0lfGainStr = imChan0lfGainStr + "\n" + imAnyChanFullBandStr = imAnyChanFullBandStr + "\n" + + return imChan0apGainStr, imChan0lfGainStr, imAnyChanFullBandStr + + +# ========================================================= +# Parse snsGeomMap for XY coordinates +# +def geomMapToGeom(meta): + + # read in the shank map + geomMap = meta["snsGeomMap"].split(sep=")") + + # there is an entry in the map for each saved channel + # number of entries in map -- subtract 1 for header, one for trailing ')' + nEntry = len(geomMap) - 2 + + shankInd = np.zeros((nEntry,)) + xCoord = np.zeros((nEntry,)) + yCoord = np.zeros((nEntry,)) + connected = np.zeros((nEntry,)) + + for i in range(nEntry): + # get parameter list from this entry, skipping first header entry + currEntry = geomMap[i + 1] + currEntry = currEntry[1 : len(currEntry)] + currList = currEntry.split(sep=":") + shankInd[i] = int(currList[0]) + xCoord[i] = float(currList[1]) + yCoord[i] = float(currList[2]) + connected[i] = int(currList[3]) + + # parse header for number of shanks + currList = geomMap[0].split(",") + nShank = int(currList[1]) + shankPitch = float(currList[2]) + shankWidth = float(currList[3]) + + return nShank, shankWidth, shankPitch, shankInd, xCoord, yCoord, connected + + +# ========================================================= +# Build snsGeomMap from xy coordinates +# +def snsGeom(meta, shankInd, xCoord, yCoord, use): + # header + # get probe part number; if absent, this is a 3A + if "imDatPrb_pn" in meta: + pn = meta["imDatPrb_pn"] + else: + pn = "3A" + + geomList = getGeomParams(meta) + # geomList = + # [nShank, shankWidth, shankPitch, even_xOff, odd_xOff, horizPitch, vertPitch, rowsPerShank, elecPerShank] + + snsGeomStr = "~snsGeomMap=(" + pn + snsGeomStr = snsGeomStr + ",{:d},{:g},{:g})".format( + geomList[0], geomList[2], geomList[1] + ) + nEntry = shankInd.shape[0] + + for i in range(0, nEntry): + snsGeomStr = snsGeomStr + "({:g}:{:g}:{:g}:{:g})".format( + shankInd[i], xCoord[i], yCoord[i], use[i] + ) + + snsGeomStr = snsGeomStr + "\n" + + return snsGeomStr + + +# ========================================================= +# Get XY coordinates from snsShankMap plus hard coded geom values +# +def shankMapToGeom(meta): + # get number of saved AP channels (some early metadata files have a + # SYNC entry in the snsChanMap) + AP, LF, SY = ChannelCountsIM(meta) + + shankMap = meta["snsShankMap"].split(sep=")") + + shankInd = np.zeros((AP,)) + colInd = np.zeros((AP,)) + rowInd = np.zeros((AP,)) + connected = np.zeros((AP,)) + xCoord = np.zeros((AP,)) + yCoord = np.zeros((AP,)) + + for i in range(AP): + # get parameter list from this entry, skipping first header entry + currEntry = shankMap[i + 1] + currEntry = currEntry[1 : len(currEntry)] + currList = currEntry.split(sep=":") + shankInd[i] = int(currList[0]) + colInd[i] = int(currList[1]) + rowInd[i] = int(currList[2]) + connected[i] = int(currList[3]) + + geomList = getGeomParams(meta) + # geomList = + # [nShank, shankWidth, shankPitch, even_xOff, odd_xOff, horizPitch, vertPitch, rowsPerShank, elecPerShank] + + oddRows = np.bool_(rowInd % 2) + evenRows = ~oddRows + xCoord = colInd * float(geomList[5]) + xCoord[evenRows] = xCoord[evenRows] + geomList[3] + xCoord[oddRows] = xCoord[oddRows] + geomList[4] + yCoord = rowInd * float(geomList[6]) + + nShank = geomList[0] + shankWidth = geomList[1] + shankPitch = geomList[2] + + return nShank, shankWidth, shankPitch, shankInd, xCoord, yCoord, connected + + +# ========================================================= +# Plot x z positions of all electrodes and saved channels +# +def plotSaved(xCoord, yCoord, shankInd, meta): + + geomList = getGeomParams(meta) + # geomList = + # [nShank, shankWidth, shankPitch, even_xOff, odd_xOff, horizPitch, vertPitch, rowsPerShank, elecPerShank] + + # calculate positions on one shank + nCol = geomList[8] / geomList[7] + rowInd = np.arange(geomList[8]) + rowInd = np.floor(rowInd / nCol) + oddRows = np.bool_(rowInd % 2) + evenRows = ~oddRows + + colInd = np.arange(geomList[8]) + colInd = colInd % nCol + + xall = colInd * geomList[5] + xall[evenRows] = xall[evenRows] + geomList[3] + xall[oddRows] = xall[oddRows] + geomList[4] + + yall = rowInd * geomList[6] + + fig = plt.figure(figsize=(2, 12)) + + shankSep = geomList[2] + + # loop over shanks + for sI in range(geomList[0]): + + # plot all positions + marker_style = dict(c="w", edgecolor="k", linestyle="None", marker="s", s=5) + plt.scatter(shankSep * sI + xall, yall, **marker_style) + + # plot selected positions + currInd = np.argwhere(shankInd == sI) + marker_style = dict(c="b", edgecolor="g", linestyle="None", marker="s", s=15) + plt.scatter(shankSep * sI + xCoord[currInd], yCoord[currInd], **marker_style) + + # after looping over all shanks, show plot + plt.show() + + return + + +# ========================================================= +# CoordsTo... functions to write output in different +# formats +# +def CoordsToText( + meta, + chans, + xCoord, + yCoord, + connected, + shankInd, + shankSep, + baseName, + savePath, + buildPath, +): + + if buildPath: + newName = baseName + "_siteCoords.txt" + saveFullPath = Path(savePath / newName) + else: + saveFullPath = savePath + + # Note that the channel index written is the index of that channel in the saved file + + with open(saveFullPath, "w") as outFile: + for i in range(0, chans.size): + currX = shankInd[i] * shankSep + xCoord[i] + currLine = "{:d}\t{:g}\t{:g}\t{:g}\n".format( + i, currX, yCoord[i], shankInd[i] + ) + outFile.write(currLine) + + +def CoordsToNPY( + meta, + chans, + xCoord, + yCoord, + connected, + shankInd, + shankSep, + baseName, + savePath, + buildPath, +): + + if buildPath: + newName = baseName + "_siteCoords.npy" + saveFullPath = Path(savePath / newName) + else: + saveFullPath = savePath + + # write an npy file of nChanx2 + nchan = xCoord.shape[0] + geom = np.zeros((nchan, 2)) + geom[:, 0] = xCoord + shankInd * shankSep + geom[:, 1] = yCoord + + np.save(saveFullPath, geom) + + +def CoordsToJRCString( + meta, + chans, + xCoord, + yCoord, + connected, + shankInd, + shankSep, + baseName, + savePath, + buildPath, +): + + if buildPath: + newName = baseName + "_forJRCprm.txt" + saveFullPath = Path(savePath / newName) + else: + saveFullPath = savePath + + # siteMap, equivalent of chanMap in KS, is the order of channels in the saved file + # rather than original channel indicies. + nChan = chans.size + siteMap = np.arange(0, nChan, dtype="int") + siteMap = siteMap + 1 # convert to 1-based for MATLAB + + shankInd = shankInd + 1 # conver to 1-based for MATLAB + + shankStr = "shankMap = [" + coordStr = "siteLoc = [" + siteMapStr = "siteMap = [" + + xCoord = shankInd * shankSep + xCoord + + for i in range(0, chans.size - 1): + shankStr = shankStr + "{:g},".format( + shankInd[i] + ) # convert to 1-based for MATLAB + coordStr = coordStr + "{:g},{:g};".format(xCoord[i], yCoord[i]) + siteMapStr = siteMapStr + "{:d},".format( + siteMap[i] + ) # convert to 1-based for MATLAB + + # final entries + shankStr = shankStr + "{:g}];\n".format(shankInd[nChan - 1]) + coordStr = coordStr + "{:g},{:g}];\n".format(xCoord[nChan - 1], yCoord[nChan - 1]) + siteMapStr = siteMapStr + "{:d}];\n".format(siteMap[nChan - 1]) + + with open(saveFullPath, "w") as outFile: + outFile.write(shankStr) + outFile.write(coordStr) + outFile.write(siteMapStr) + + +def CoordsToKSChanMap( + meta, + chans, + xCoord, + yCoord, + connected, + shankInd, + shankSep, + baseName, + savePath, + buildPath, +): + + if buildPath: + newName = baseName + "_kilosortChanMap.mat" + saveFullPath = Path(savePath / newName) + else: + saveFullPath = savePath + + nChan = chans.size + # channel map is the order of channels in the file, rather than the + # original indicies of the channels + chanMap0ind = np.arange(0, nChan, dtype="float64") + chanMap0ind = chanMap0ind.reshape((nChan, 1)) + chanMap = chanMap0ind + 1 + + connected = connected == 1 + connected = connected.reshape((nChan, 1)) + + xCoord = shankInd * shankSep + xCoord + xCoord = xCoord.reshape((nChan, 1)) + yCoord = yCoord.reshape((nChan, 1)) + + kcoords = shankInd + 1 + kcoords = kcoords.reshape((nChan, 1)) + kcoords = kcoords.astype("float64") + + name = baseName + + mdict = { + "chanMap": chanMap, + "chanMap0ind": chanMap0ind, + "connected": connected, + "name": name, + "xcoords": xCoord, + "ycoords": yCoord, + "kcoords": kcoords, + } + scipy.io.savemat(saveFullPath, mdict) + + +def CoordsToGeomMap( + meta, + chans, + xCoord, + yCoord, + connected, + shankInd, + shankSep, + baseName, + savePath, + buildPath, +): + + if buildPath: + newName = baseName + "_orig.meta" + copyFullPath = Path(savePath / newName) + else: + print("Can only make new metadata in same directory as original.") + return + + origPath = Path(savePath / (baseName + ".meta")) + shutil.move(origPath, copyFullPath) + shutil.copy(copyFullPath, origPath) + + # check 'new' fields; add if not present, add everything (most common case) + if "imChan0apGain" not in meta: + imChan0apGainStr, imChan0lfGainStr, imAnyChanFullBandStr = imroMetaItems(meta) + muxTableStr = getMuxTable(meta) + snsGeomStr = snsGeom(meta, shankInd, xCoord, yCoord, connected) + + with open(origPath, "a") as outFile: + outFile.write(imChan0apGainStr) + outFile.write(imChan0lfGainStr) + outFile.write(imAnyChanFullBandStr) + outFile.write(muxTableStr) + outFile.write(snsGeomStr) + + +# ========================================================= +# Given a path to a SpikeGLX metadata file, write out coordinates +# in formats for analysis software to consume +# Input params: +# metaFullPath: full path, including the file name +# outType: format for the output +# badChan: channels other than reference channels to exclude +# destFullPath: +# +def MetaToCoords( + metaFullPath, + outType, + badChan=np.zeros((0), dtype="int"), + destFullPath="", + showPlot=False, +): + + # Read in metadata; returns a dictionary with string for values + meta = readMeta(metaFullPath) + + # Get coordinates for saved channels from snsGeomMap, if present, + # otherwise from snsShankMap + if "snsGeomMap" in meta: + [nShank, shankWidth, shankPitch, shankInd, xCoord, yCoord, connected] = ( + geomMapToGeom(meta) + ) + else: + [nShank, shankWidth, shankPitch, shankInd, xCoord, yCoord, connected] = ( + shankMapToGeom(meta) + ) + + if showPlot: + plotSaved(xCoord, yCoord, shankInd, meta) + + [AP, LF, SY] = ChannelCountsIM(meta) + chans = np.arange(AP) + + # Channels identified as noisy by kilosort helper indexed + # according to position in the file + # since these can include the SYNC channel, remove any from + # list that are outside the range of AP channels + badChan = badChan[badChan < AP] + connected[badChan] = 0 + + baseName = metaFullPath.stem + NchanTOT = meta["nSavedChans"] + + if outType >= 0: + if len(destFullPath) == 0: + savePath = metaFullPath.parent + buildPath = True + else: + buildPath = False + savePath = destFullPath + outputSwitch = { + 0: CoordsToText, + 1: CoordsToKSChanMap, + 2: CoordsToJRCString, + 3: CoordsToGeomMap, + 4: CoordsToNPY, + } + + writeFunc = outputSwitch.get(outType) + writeFunc( + meta, + chans, + xCoord, + yCoord, + connected, + shankInd, + shankPitch, + baseName, + savePath, + buildPath, + ) + + return xCoord, yCoord, shankInd, connected, NchanTOT + + +# ========================================================= +# Sample calling program to get a metadata file from the user, +# output a file set by outType +# 0 = tab delimited text file of coordinates in um, index, x, y, shank index +# 1 = KS2 chan map .mat file +# 2 = strings of channel map, shank index, and x,y pairs for JRClust +# 3 = create a new metadata file with snsGeomMap appended (for converting old metadata to new) +# 4 = npy file with (nchan,2) matrix of xy coordinates, e.g. for YASS and related sorters +# file is saved to the same path as the metadata file. +# +def main(): + + outType = 1 + + # Get file from user + root = Tk() # create the Tkinter widget + root.withdraw() # hide the Tkinter root window + + # Windows specific; forces the window to appear in front + root.attributes("-topmost", True) + + metaFullPath = Path(filedialog.askopenfilename(title="Select meta file")) + root.destroy() # destroy the Tkinter widget + + MetaToCoords(metaFullPath=metaFullPath, outType=outType, showPlot=True) + + +if __name__ == "__main__": + main() diff --git a/npx_utils/sglx/__init__,py b/npx_utils/sglx/__init__,py new file mode 100644 index 0000000..e69de29 diff --git a/npx_utils/sglx/sglx_helpers.py b/npx_utils/sglx/sglx_helpers.py new file mode 100644 index 0000000..50bf176 --- /dev/null +++ b/npx_utils/sglx/sglx_helpers.py @@ -0,0 +1,359 @@ +import os +import re + +import numpy as np +from scipy.stats import mode + +from . import _SGLXMetaToCoords + + +def read_meta(meta_path): + return _SGLXMetaToCoords.readMeta(meta_path) + + +def get_channel_counts(meta): + return _SGLXMetaToCoords.ChannelCountsIM(meta) + + +def get_sample_rate(meta): + if meta["typeThis"] == "imec": + srate = float(meta["imSampRate"]) + elif meta["typeThis"] == "nidq": + srate = float(meta["niSampRate"]) + elif meta["typeThis"] == "obx": + srate = float(meta["obSampRate"]) + else: + raise ValueError("Unknown stream type") + return srate + + +def convert_data_to_uV(data, chan_list, meta): + if meta["typeThis"] == "imec": + data_V = GainCorrectIM(data, chan_list, meta) + elif meta["typeThis"] == "nidq": + data_V = GainCorrectNI(data, chan_list, meta) + elif meta["typeThis"] == "obx": + data_V = GainCorrectOBX(data, chan_list, meta) + else: + raise ValueError("Unknown stream type") + data_uV = data_V * 1e6 + return data_uV + + +def get_bits_to_uV(chan_list, meta): + if meta["typeThis"] == "imec": + bits_to_V = get_gain_correction_im(chan_list, meta) + elif meta["typeThis"] == "nidq": + # bits_to_V = GainCorrectNI(data, chan_list, meta) + pass + elif meta["typeThis"] == "obx": + # bits_to_V = GainCorrectOBX(data, chan_list, meta) + pass + else: + raise ValueError("Unknown stream type") + bits_to_uV = bits_to_V * 1e6 + return bits_to_uV + + +def get_data_memmap(bin_path, meta): + nChan = int(meta["nSavedChans"]) + nFileSamp = int(int(meta["fileSizeBytes"]) / (2 * nChan)) + rawData = np.memmap( + bin_path, dtype="int16", mode="r", shape=(nFileSamp, nChan), offset=0, order="F" + ) + return rawData + + +def get_same_channel_positions(ks_folders): + """ + Get the kilosort folders with the same channel positions. + """ + channel_positions = [] + channel_positions_tuples = [] + for ks_folder in ks_folders: + channel_position = np.load(os.path.join(ks_folder, "channel_positions.npy")) + # convert to immutable tuple + channel_position_tuple = tuple(map(tuple, channel_position)) + channel_positions.append(channel_position) + channel_positions_tuples.append(channel_position_tuple) + + most_common = mode(channel_positions_tuples, axis=0) + indices = [ + i + for i in range(len(channel_positions)) + if np.array_equal(channel_positions[i], most_common.mode) + ] + return [ks_folders[i] for i in indices] + + +def get_svy_bank_times(meta): + if meta.get("svySBTT") is None: + return 0 + pattern = r"(\d+)\s+(\d+)\s+(\d+)\s+(\d+)" + parsed_data = re.findall(pattern, meta["svySBTT"]) + if not parsed_data: + return ValueError("No valid bank times found in svySBTT") + parsed_data = np.array(parsed_data, dtype=float) + n_bank = len(parsed_data) + 1 + bank_times = np.zeros((n_bank, 4), dtype=float) + sample_rate = get_sample_rate(meta) + bank_times[1:, 0:2] = parsed_data[:, 0:2] # Shank and Bank IDs + bank_times[1:, 2:4] = parsed_data[:, 2:4] / sample_rate + return bank_times + + +### Helpers +def OriginalChans(meta): + """ + Return array of original channel IDs. As an example, suppose we want the + imec gain for the ith channel stored in the binary data. A gain array + can be obtained using ChanGainsIM(), but we need an original channel + index to do the lookup. Because you can selectively save channels, the + ith channel in the file isn't necessarily the ith acquired channel. + Use this function to convert from ith stored to original index. + Note that the SpikeGLX channels are 0 based. + """ + if meta["snsSaveChanSubset"] == "all": + # output = int32, 0 to nSavedChans - 1 + chans = np.arange(0, int(meta["nSavedChans"])) + else: + # parse the snsSaveChanSubset string + # split at commas + chStrList = meta["snsSaveChanSubset"].split(sep=",") + chans = np.arange(0, 0) # creates an empty array of int32 + for sL in chStrList: + currList = sL.split(sep=":") + if len(currList) > 1: + # each set of contiguous channels specified by + # chan1:chan2 inclusive + newChans = np.arange(int(currList[0]), int(currList[1]) + 1) + else: + newChans = np.arange(int(currList[0]), int(currList[0]) + 1) + chans = np.append(chans, newChans) + return chans + + +def ChannelCountsNI(meta): + """ + Return counts of each nidq channel type that composes the timepoints stored in the binary file. + """ + chanCountList = meta["snsMnMaXaDw"].split(sep=",") + MN = int(chanCountList[0]) + MA = int(chanCountList[1]) + XA = int(chanCountList[2]) + DW = int(chanCountList[3]) + return (MN, MA, XA, DW) + + +def ChannelCountsIM(meta): + """ + Return counts of each imec channel type that composes the timepoints stored in the binary files. + """ + chanCountList = meta["snsApLfSy"].split(sep=",") + AP = int(chanCountList[0]) + LF = int(chanCountList[1]) + SY = int(chanCountList[2]) + return (AP, LF, SY) + + +def ChannelCountsOBX(meta): + """ + Return counts of each obx channel type that composes the timepoints stored in the binary files. + """ + chanCountList = meta["snsXaDwSy"].split(sep=",") + XA = int(chanCountList[0]) + DW = int(chanCountList[1]) + SY = int(chanCountList[2]) + return (XA, DW, SY) + + +def ChanGainNI(ichan, savedMN, savedMA, meta): + """ + Return gain for ith channel stored in nidq file. + ichan is a saved channel index, rather than the original (acquired) index. + """ + if ichan < savedMN: + gain = float(meta["niMNGain"]) + elif ichan < (savedMN + savedMA): + gain = float(meta["niMAGain"]) + else: + gain = 1 # non multiplexed channels have no extra gain + return gain + + +def ChanGainsIM(meta): + """ + Return gain for imec channels. + Index into these with the original (acquired) channel IDs. + """ + # list of probe types with NP 1.0 imro format + np1_imro = [0, 1020, 1030, 1200, 1100, 1120, 1121, 1122, 1123, 1300] + # number of channels acquired + acqCountList = meta["acqApLfSy"].split(sep=",") + APgain = np.zeros(int(acqCountList[0])) # default type = float64 + LFgain = np.zeros(int(acqCountList[1])) # empty array for 2.0 + + if "imDatPrb_type" in meta: + probeType = int(meta["imDatPrb_type"]) + else: + probeType = 0 + + if sum(np.isin(np1_imro, probeType)): + # imro + probe allows setting gain independently for each channel + imroList = meta["imroTbl"].split(sep=")") + # One entry for each channel plus header entry, + # plus a final empty entry following the last ')' + for i in range(0, int(acqCountList[0])): + currList = imroList[i + 1].split(sep=" ") + APgain[i] = float(currList[3]) + LFgain[i] = float(currList[4]) + else: + # get gain from imChan0apGain + if "imChan0apGain" in meta: + APgain = APgain + float(meta["imChan0apGain"]) + if int(acqCountList[1]) > 0: + LFgain = LFgain + float(meta["imChan0lfGain"]) + elif probeType == 1110: + # active UHD, for metadata lacking imChan0apGain, get gain from + # imro table header + imroList = meta["imroTbl"].split(sep=")") + currList = imroList[0].split(sep=",") + APgain = APgain + float(currList[3]) + LFgain = LFgain + float(currList[4]) + elif (probeType == 21) or (probeType == 24): + # development NP 2.0; APGain = 80 for all AP + # return 0 for LFgain (no LF channels) + APgain = APgain + 80 + elif probeType == 2013: + # commercial NP 2.0; APGain = 100 for all AP + APgain = APgain + 100 + else: + print("unknown gain, setting APgain to 1") + APgain = APgain + 1 + fI2V = Int2Volts(meta) + APChan0_to_uV = 1e6 * fI2V / APgain[0] + if LFgain.size > 0: + LFChan0_to_uV = 1e6 * fI2V / LFgain[0] + else: + LFChan0_to_uV = 0 + return (APgain, LFgain, APChan0_to_uV, LFChan0_to_uV) + + +def Int2Volts(meta): + """ + Return a multiplicative factor for converting 16-bit file data + to voltage. This does not take gain into account. The full + conversion with gain is: + dataVolts = dataInt * fI2V / gain + Note that each channel may have its own gain. + """ + if meta["typeThis"] == "imec": + if "imMaxInt" in meta: + maxInt = int(meta["imMaxInt"]) + else: + maxInt = 512 + fI2V = float(meta["imAiRangeMax"]) / maxInt + elif meta["typeThis"] == "nidq": + maxInt = int(meta["niMaxInt"]) + fI2V = float(meta["niAiRangeMax"]) / maxInt + elif meta["typeThis"] == "obx": + maxInt = int(meta["obMaxInt"]) + fI2V = float(meta["obAiRangeMax"]) / maxInt + else: + print("Error: unknown stream type") + fI2V = 1 + + return fI2V + + +def GainCorrectNI(dataArray, chanList, meta): + """ + Having accessed a block of raw nidq data using makeMemMapRaw, convert + values to gain-corrected voltage. The conversion is only applied to the + saved-channel indices in chanList + """ + MN, MA, XA, DW = ChannelCountsNI(meta) + fI2V = Int2Volts(meta) + # print statements used for testing... + # print("NI fI2V: %.3e" % (fI2V)) + # print("NI ChanGainNI: %.3f" % (ChanGainNI(0, MN, MA, meta))) + + # make array of floats to return. dataArray contains only the channels + # in chanList, so output matches that shape + convArray = np.zeros(dataArray.shape, dtype=float) + for i in range(0, len(chanList)): + j = chanList[i] # index in saved data + conv = fI2V / ChanGainNI(j, MN, MA, meta) + # dataArray contains only the channels in chanList + convArray[:, i] = dataArray[:, i] * conv + return convArray + + +def GainCorrectOBX(dataArray, chanList, meta): + """ + Having accessed a block of raw obx data using makeMemMapRaw, convert + values to volts. The conversion is only applied to the + saved-channel + """ + fI2V = Int2Volts(meta) + + # make array of floats to return. dataArray contains only the channels + # in chanList, so output matches that shape + convArray = np.zeros(dataArray.shape, dtype=float) + for i in range(0, len(chanList)): + # dataArray contains only the channels in chanList + convArray[:, i] = dataArray[:, i] * fI2V + return convArray + + +def get_gain_correction_im(chanList, meta): + chans = OriginalChans(meta) + APgain, LFgain, _, _ = ChanGainsIM(meta) + nAP = len(APgain) + nNu = nAP * 2 + fI2V = Int2Volts(meta) + conversion = np.zeros(len(chanList)) + for i in range(0, len(chanList)): + j = chanList[i] # index into timepoint + k = chans[j] # acquisition index + if k < nAP: + conv = fI2V / APgain[k] + elif k < nNu: + conv = fI2V / LFgain[k - nAP] + else: + conv = 1 + # The dataArray contains only the channels in chanList + conversion[i] = conv + return conversion + + +def GainCorrectIM(dataArray, chanList, meta): + """ + Having accessed a block of raw imec data using makeMemMapRaw, convert + values to gain corrected voltages. The conversion is only applied to + the saved-channel indices in chanList. + """ + # Look up gain with acquired channel ID + chans = OriginalChans(meta) + APgain, LFgain, _, _ = ChanGainsIM(meta) + nAP = len(APgain) + nNu = nAP * 2 + + # Common conversion factor + fI2V = Int2Volts(meta) + + # make array of floats to return. dataArray contains only the channels + # in chanList, so output matches that shape + convArray = np.zeros(dataArray.shape, dtype="float") + for i in range(0, len(chanList)): + j = chanList[i] # index into timepoint + k = chans[j] # acquisition index + if k < nAP: + conv = fI2V / APgain[k] + elif k < nNu: + conv = fI2V / LFgain[k - nAP] + else: + conv = 1 + # The dataArray contains only the channels in chanList + convArray[:, i] = dataArray[:, i] * conv + return convArray diff --git a/npx_utils/sglx_helpers.py b/npx_utils/sglx_helpers.py deleted file mode 100644 index cf470d3..0000000 --- a/npx_utils/sglx_helpers.py +++ /dev/null @@ -1,284 +0,0 @@ -import os -import re -from pathlib import Path - -import numpy as np -from scipy.stats import mode - -from .ks_helpers import get_lfp_meta_path, load_params - - -def read_meta(meta_path): - meta_dict = {} - if os.path.exists(meta_path): - with open(meta_path) as f: - mdatList = f.read() - mdatList = mdatList.splitlines() - # convert the list entries into key value pairs - for m in mdatList: - csList = m.split(sep="=") - if csList[0][0] == "~": - currKey = csList[0][1 : len(csList[0])] - else: - currKey = csList[0] - meta_dict.update({currKey: csList[1]}) - else: - print("no meta file") - - return meta_dict - - -def get_all_channel_counts(meta): - chanCountList = meta["snsApLfSy"].split(sep=",") - n_channel_ap = int(chanCountList[0]) - n_channel_lf = int(chanCountList[1]) - n_channel_sync = int(chanCountList[2]) - - return n_channel_ap, n_channel_lf, n_channel_sync - - -def get_ap_data_channel_count(meta): - return get_all_channel_counts(meta)[0] - - -def get_bits_to_uV(meta): - if "imDatPrb_type" in meta: - pType = meta["imDatPrb_type"] - if pType == "0": - probe_type = "NP1" - else: - probe_type = "NP" + pType - else: - probe_type = "3A" # 3A probe is default - - # first check if metadata includes the imChan0apGain key - if "uVPerBit" in meta: - return float(meta["uVPerBit"]) - - if "imChan0apGain" in meta: - APgain = float(meta["imChan0apGain"]) - voltage_range = float(meta["imAiRangeMax"]) - float(meta["imAiRangeMin"]) - maxInt = float(meta["imMaxInt"]) - uVPerBit = (1e6) * (voltage_range / APgain) / (2 * maxInt) - - else: - imroList = meta["imroTbl"].split(sep=")") - # One entry for each channel plus header entry, - # plus a final empty entry following the last ')' - # channel zero is the 2nd element in the list - - if probe_type == "NP21" or probe_type == "NP24": - # NP 2.0; APGain = 80 for all channels - # voltage range = 1V - # 14 bit ADC - uVPerBit = (1e6) * (1.0 / 80) / pow(2, 14) - elif probe_type == "NP1110": - # UHD2 with switches, special imro table with gain in header - currList = imroList[0].split(sep=",") - APgain = float(currList[3]) - uVPerBit = (1e6) * (1.2 / APgain) / pow(2, 10) - else: - # 3A, 3B1, 3B2 (NP 1.0), or other NP 1.0-like probes - # voltage range = 1.2V - # 10 bit ADC - currList = imroList[1].split( - sep=" " - ) # 2nd element in list, skipping header - APgain = float(currList[3]) - uVPerBit = (1e6) * (1.2 / APgain) / pow(2, 10) - - # # save this value in meta - # with open(meta_path, "ab") as f: - # f.write(f"uVPerBit={uVPerBit}\n".encode("utf-8")) - - return uVPerBit - - -def get_same_channel_positions(ks_folders): - """ - Get the kilosort folders with the same channel positions. - """ - channel_positions = [] - channel_positions_tuples = [] - for ks_folder in ks_folders: - channel_position = np.load(os.path.join(ks_folder, "channel_positions.npy")) - # convert to immutable tuple - channel_position_tuple = tuple(map(tuple, channel_position)) - channel_positions.append(channel_position) - channel_positions_tuples.append(channel_position_tuple) - - most_common = mode(channel_positions_tuples, axis=0) - indices = [ - i - for i in range(len(channel_positions)) - if np.array_equal(channel_positions[i], most_common.mode) - ] - return [ks_folders[i] for i in indices] - - -def get_data_memmap(ks_folder): - """ - Load the data from the binary file as a memory-mapped array. - """ - params = load_params(ks_folder) - data = np.memmap(params["dat_path"], dtype="int16", mode="r") - data = np.reshape(data, (-1, params["n_channels_dat"])) - return data - - -def get_lfp_memmap(ks_folder): - """ - Load the data from the binary file as a memory-mapped array. - """ - params = load_params(ks_folder) - lfp_path = params["dat_path"].replace(".ap.bin", ".lf.bin") - data = np.memmap(lfp_path, dtype="int16", mode="r") - data = np.reshape(data, (-1, params["n_channels_dat"])) - return data - - -def get_lfp_sample_rate(ks_folder): - meta_path = get_lfp_meta_path(ks_folder) - meta = read_meta(meta_path) - return float(meta["imSampRate"]) - - -def get_sample_rate(meta): - """ - Get the sample rate from the metadata. - """ - if meta["typeThis"] == "imec": - sample_rate = float(meta["imSampRate"]) - elif meta["typeThis"] == "nidq": - sample_rate = float(meta["niSampRate"]) - elif meta["typeThis"] == "obx": - sample_rate = float(meta["obSampRate"]) - else: - print("Error: unknown stream type") - sample_rate = 1 - return sample_rate - - -def get_svy_bank_times(meta): - if meta.get("svySBTT") is None: - return 0 - pattern = r"(\d+)\s+(\d+)\s+(\d+)\s+(\d+)" - parsed_data = re.findall(pattern, meta["svySBTT"]) - if not parsed_data: - return ValueError("No valid bank times found in svySBTT") - parsed_data = np.array(parsed_data, dtype=float) - n_bank = len(parsed_data) + 1 - bank_times = np.zeros((n_bank, 4), dtype=float) - sample_rate = get_sample_rate(meta) - bank_times[1:, 0:2] = parsed_data[:, 0:2] # Shank and Bank IDs - bank_times[1:, 2:4] = parsed_data[:, 2:4] / sample_rate - return bank_times - - -def get_chan_gains_imec(meta): - # list of probe types with NP 1.0 imro format - np1_imro = [0, 1020, 1030, 1200, 1100, 1120, 1121, 1122, 1123, 1300] - # number of channels acquired - acqCountList = meta["acqApLfSy"].split(sep=",") - APgain = np.zeros(int(acqCountList[0])) # default type = float64 - LFgain = np.zeros(int(acqCountList[1])) # empty array for 2.0 - - if "imDatPrb_type" in meta: - probeType = int(meta["imDatPrb_type"]) - else: - probeType = 0 - - if sum(np.isin(np1_imro, probeType)): - # imro + probe allows setting gain independently for each channel - imroList = meta["imroTbl"].split(sep=")") - # One entry for each channel plus header entry, - # plus a final empty entry following the last ')' - for i in range(0, int(acqCountList[0])): - currList = imroList[i + 1].split(sep=" ") - APgain[i] = float(currList[3]) - LFgain[i] = float(currList[4]) - else: - # get gain from imChan0apGain - if "imChan0apGain" in meta: - APgain = APgain + float(meta["imChan0apGain"]) - if int(acqCountList[1]) > 0: - LFgain = LFgain + float(meta["imChan0lfGain"]) - elif probeType == 1110: - # active UHD, for metadata lacking imChan0apGain, get gain from - # imro table header - imroList = meta["imroTbl"].split(sep=")") - currList = imroList[0].split(sep=",") - APgain = APgain + float(currList[3]) - LFgain = LFgain + float(currList[4]) - elif (probeType == 21) or (probeType == 24): - # development NP 2.0; APGain = 80 for all AP - # return 0 for LFgain (no LF channels) - APgain = APgain + 80 - elif probeType == 2013: - # commercial NP 2.0; APGain = 100 for all AP - APgain = APgain + 100 - else: - print("unknown gain, setting APgain to 1") - APgain = APgain + 1 - fI2V = int2volts(meta) - APChan0_to_uV = 1e6 * fI2V / APgain[0] - if LFgain.size > 0: - LFChan0_to_uV = 1e6 * fI2V / LFgain[0] - else: - LFChan0_to_uV = 0 - return (APgain, LFgain, APChan0_to_uV, LFChan0_to_uV) - - -def int2volts(meta): - if meta["typeThis"] == "imec": - if "imMaxInt" in meta: - maxInt = int(meta["imMaxInt"]) - else: - maxInt = 512 - fI2V = float(meta["imAiRangeMax"]) / maxInt - elif meta["typeThis"] == "nidq": - maxInt = int(meta["niMaxInt"]) - fI2V = float(meta["niAiRangeMax"]) / maxInt - elif meta["typeThis"] == "obx": - maxInt = int(meta["obMaxInt"]) - fI2V = float(meta["obAiRangeMax"]) / maxInt - else: - print("Error: unknown stream type") - fI2V = 1 - - return fI2V - - -def get_chan_gains_ni(ichan, savedMN, savedMA, meta): - if ichan < savedMN: - gain = float(meta["niMNGain"]) - elif ichan < (savedMN + savedMA): - gain = float(meta["niMAGain"]) - else: - gain = 1 # non multiplexed channels have no extra gain - return gain - - -def get_chan_counts_imec(meta): - chanCountList = meta["snsApLfSy"].split(sep=",") - AP = int(chanCountList[0]) - LF = int(chanCountList[1]) - SY = int(chanCountList[2]) - return (AP, LF, SY) - - -def get_chan_counts_ni(meta): - chanCountList = meta["snsMnMaXaDw"].split(sep=",") - MN = int(chanCountList[0]) - MA = int(chanCountList[1]) - XA = int(chanCountList[2]) - DW = int(chanCountList[3]) - return (MN, MA, XA, DW) - - -def get_chan_counts_obx(meta): - chanCountList = meta["snsXaDwSy"].split(sep=",") - XA = int(chanCountList[0]) - DW = int(chanCountList[1]) - SY = int(chanCountList[2]) - return (XA, DW, SY) diff --git a/npx_utils/stability/__init__.py b/npx_utils/stability/__init__.py new file mode 100644 index 0000000..0d588bb --- /dev/null +++ b/npx_utils/stability/__init__.py @@ -0,0 +1,2 @@ +from .threshold_event_plots import plot_threshold_events +from .unit_counts import plot_unit_counts diff --git a/npx_utils/stability/threshold_event_plots.py b/npx_utils/stability/threshold_event_plots.py new file mode 100644 index 0000000..69585c4 --- /dev/null +++ b/npx_utils/stability/threshold_event_plots.py @@ -0,0 +1,891 @@ +# -*- coding: utf-8 -*- +# taken from Jennifer Collonel +""" +Created on Sat Mar 1 18:07:21 2025 + +Should only be run on filtered, background subtracted data (e.g. from CatGT). +Threshold the binary, merge spikes that are sufficiently close in time and space +to count as events. Algorithm adapated from JRClust. + +Requires SGLX utilties to read metadata and binary. + +readData creates an npy file of spike properties for a selected shank, named +'{binary_name}_dd_sh{shank index}.npy'. The file is saved in the directory with the +Binary. The plotting routines read these files to make the plots. + +The columns in the drfit data array are: + spike times(sec) + spike z position (from center of mass), in um + amplitude of negative going peak, uV + spike x position (from center of mass), in um + peak channel (of the channels on the shank -- not remapped to original channel in binary) + + +Plot types: + plotOne: drift raster plot for a single shank, single recording + plotMult: plot drift rasters for single shank, multiple recordings (typically, multiple days) + useful for by eye drift assessment + plotMultPDF: plot amplitude prob density function across all shanks, multiple recordings + (figure 2A of Steinmetz, et al. NP 2.0 paper) + plotSpikeRate: plot total spike rate across all shanks + (figure 2C of Steinmetz, et al. NP 2.0 paper) + plotRateVsZ: plot relative spike rate vs. z position for a single shank, multiple recordings + (figure 2B of Steinmetz, et al. NP 2.0 paper) + + +""" +import os +from datetime import datetime +from pathlib import Path + +import matplotlib as mpl +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.colors import ListedColormap +from matplotlib.patches import Rectangle +from npx_utils.ks_helpers import get_binary_path, get_ks_folders +from npx_utils.other_helpers import get_probe_folders, get_run_folders +from npx_utils.sglx._SGLXMetaToCoords import MetaToCoords +from npx_utils.sglx.sglx_helpers import ( + ChanGainsIM, + ChannelCountsIM, + get_data_memmap, + get_same_channel_positions, + get_sample_rate, + read_meta, +) +from tqdm.auto import tqdm + + +def getDriftDataPath(binFullPath, sh_index): + parent_path = binFullPath.parent + bin_name = binFullPath.stem + dd_name = f"{bin_name}_dd_sh{sh_index}.npy" + dd_path = parent_path.joinpath(dd_name) + return dd_path + + +def findPeaks(samplesIn, thresh, xc, zc, excl_chan, fs=30000): + # find peaks with negative amp > threshold, in a batch of + # samples from all channels + # range of points for calculateing a full waveform + n_chan, n_samp = samplesIn.shape + start_time = np.floor(fs * 1.0 / 1000).astype(int) + end_time = n_samp - np.floor(fs * 1.0 / 1000).astype(int) + + # threshold should be in bits + exceeds_thresh = samplesIn < -np.abs(thresh) # boolean (n_chan, n_samp) + + # find local minima by taking only points whose neighbor points are + # more positive. + (peak_chan, peak_ind) = np.where(exceeds_thresh) + + # remove spikes that occur on excluded channels + for ch in excl_chan: + rem_ind = np.where(peak_chan == ch)[0] + if rem_ind.size > 0: + peak_ind = np.delete(peak_ind, rem_ind) + peak_chan = np.delete(peak_chan, rem_ind) + + too_early_ind = np.where(peak_ind < start_time)[0] + if too_early_ind.size > 0: + peak_ind = np.delete(peak_ind, too_early_ind) + peak_chan = np.delete(peak_chan, too_early_ind) + + too_late_ind = np.where(peak_ind > end_time)[0] + if too_late_ind.size > 0: + peak_ind = np.delete(peak_ind, too_late_ind) + peak_chan = np.delete(peak_chan, too_late_ind) + + peak_center = samplesIn[peak_chan, peak_ind] # amplitudes of the peaks + # compare each peak to neighbors in time, earlier and later, keep only + # those that are local minima + loc_min = (samplesIn[peak_chan, peak_ind - 1] > peak_center) & ( + samplesIn[peak_chan, peak_ind + 1] > peak_center + ) + peak_ind = peak_ind[loc_min] + peak_chan = peak_chan[loc_min] + + # valid peaks have 3 points in a row below threshold + pass_inarow = (samplesIn[peak_chan, peak_ind - 1] < thresh) & ( + samplesIn[peak_chan, peak_ind + 1] < thresh + ) + peak_ind = peak_ind[pass_inarow] + peak_chan = peak_chan[pass_inarow] + peak_sig = samplesIn[peak_chan, peak_ind] + # print(f'before merging: {peak_ind.size}') + peak_ind, peak_chan, peak_sig = mergePeaks( + peak_ind, peak_chan, peak_sig, xc, zc, fs + ) + # print(f'after merging: {peak_ind.size}') + xz = spike_pos(samplesIn, peak_ind, peak_chan, xc, zc) + return peak_ind, peak_chan, peak_sig, xz + + +def mergePeaks(peak_ind, peak_chan, peak_sig, xc, zc, fs=30000): + # spikes are detected on multiple sites + # assume that spikes detected within a time threshold + # and physical radius belong to the same spikeing 'event.' + # Merge these into one event, calculate peak channel, + # x and z center of mass based on the negative-going signal + + nLim = np.floor((1 / 1000) * fs).astype(int) # merge spikes within +/- 1 ms + neigh_radius_um = 60 + near_sites = calc_neighbor_sites(xc, zc, neigh_radius_um) + + # sort all spikes in time + sort_order = np.argsort(peak_ind) + peak_ind = peak_ind[sort_order] + peak_chan = peak_chan[sort_order] + peak_sig = peak_sig[sort_order] + + chan_set = np.unique(peak_chan) + num_chan = chan_set.size + + # remove spikes that are within 1 ms in each channels spike train + for i_chan in chan_set: + curr_ind = np.where(peak_chan == i_chan)[0] + curr_times = peak_ind[curr_ind] + curr_amp = peak_sig[curr_ind] + spikes_to_check = np.where(np.diff(curr_times) < nLim)[0] + amp_early = curr_amp[spikes_to_check] + amp_late = curr_amp[spikes_to_check + 1] + keep_early = (amp_early < amp_late).astype( + int + ) # looking for the more negative spike + ind_to_remove = curr_ind[spikes_to_check + keep_early] + peak_ind = np.delete(peak_ind, ind_to_remove) + peak_chan = np.delete(peak_chan, ind_to_remove) + peak_sig = np.delete(peak_sig, ind_to_remove) + + for i_chan in chan_set: + neigh_chan = near_sites[i_chan] + # remove current channel + neigh_chan = neigh_chan[neigh_chan != i_chan] + + for j_chan in neigh_chan: + + i_ind = np.where(peak_chan == i_chan)[0] + j_ind = np.where(peak_chan == j_chan)[0] + + orig_ind = np.concatenate((i_ind, j_ind)) + ij_labels = np.concatenate( + (np.zeros((i_ind.size,), dtype=int), np.ones((j_ind.size,), dtype=int)) + ) + ij_times = np.concatenate((peak_ind[i_ind], peak_ind[j_ind])) + ij_amps = np.concatenate((peak_sig[i_ind], peak_sig[j_ind])) + + order = np.argsort(ij_times) + # reorder everything + orig_ind = orig_ind[order] + ij_labels = ij_labels[order] + ij_times = ij_times[order] + ij_amps = ij_amps[order] + + spikes_to_check = np.where(np.diff(ij_times) < nLim)[0] + amp_early = ij_amps[spikes_to_check] + amp_late = ij_amps[spikes_to_check + 1] + keep_early = (amp_early < amp_late).astype( + int + ) # looking for the more negative spike + relative_ind_to_remove = spikes_to_check + keep_early + ind_to_remove = orig_ind[relative_ind_to_remove] + peak_ind = np.delete(peak_ind, ind_to_remove) + peak_chan = np.delete(peak_chan, ind_to_remove) + peak_sig = np.delete(peak_sig, ind_to_remove) + + return peak_ind, peak_chan, peak_sig + + +def spike_pos(samplesIn, peak_ind, peak_chan, xc, zc): + xz = np.zeros((peak_ind.size, 2)) + near_sites = calc_neighbor_sites(xc, zc, neigh_radius_um=60) + chan_set = np.unique(peak_chan) + nSamp = 15 # before and after peak, 31 total ~ 1 msec + + # calculate these channel-wise because those use the same set of + # neighbor channels + for i_chan in chan_set: + i_ind = np.where(peak_chan == i_chan)[0] + cn = near_sites[i_chan] + for ci in i_ind: + # get section of data + ct = peak_ind[ci] + curr_dat = samplesIn[cn, ct - nSamp : ct + nSamp] + amps = np.abs(np.squeeze(np.min(curr_dat, axis=1))) + norm = np.sum(amps) + cm_x = np.sum(np.multiply(xc[cn], amps)) / norm + cm_z = np.sum(np.multiply(zc[cn], amps)) / norm + xz[ci] = [cm_x, cm_z] + + return xz + + +def calc_neighbor_sites(xc, zc, neigh_radius_um): + # return a list of arrays of site indicies within site_radius_um + n_site = xc.size + near_sites = list() + rad_sq = neigh_radius_um * neigh_radius_um + for i in range(n_site): + dist = np.square(xc - xc[i]) + np.square(zc - zc[i]) + neigh = np.where(dist < rad_sq)[0] + near_sites.append(neigh) + return near_sites + + +def plotOne(drift_data): + fig, ax = plt.subplots(figsize=(6, 2)) + n_spike = drift_data.shape[0] + skip_step = np.floor(n_spike / 50000) + 1 + plot_spikes = np.arange(0, n_spike, skip_step).astype(int) + pd = drift_data[plot_spikes] + + c_lim = np.asarray([np.quantile(pd[:, 2], 0.1), np.quantile(pd[:, 2], 0.9)]) + even_divisor = 50 + c_lim = even_divisor * np.floor(c_lim / even_divisor) + plt.scatter( + pd[:, 0], + pd[:, 1], + s=0.1, + c=pd[:, 2], + cmap="plasma", + vmin=c_lim[0], + vmax=c_lim[1], + ) + c = plt.colorbar() + return fig + + +def plotMult(bin_list, drift_list, day_list, sh_list): + """ + Plots raster plot of spikes for each shank across recording sessions to give estimate of drift. + """ + # build a large sampled array from the n_sets, to look for 'obvious' drift + fig, ax = plt.subplots( + nrows=2, + ncols=2, + figsize=(0.25 + 1 * len(bin_list), 8), + sharex=True, + sharey=False, + ) + fig.suptitle("Spike Drift Across Shanks", fontsize=16) + max_spike = 50000 # total spikes in the output plot + even_divisor = 10 # color and z limits will be integer multiples of this value + + # get global limits for color bar + # get duration of a recording + all_amplitudes = [] + session_durations = [[] for _ in bin_list] + for sh_ind in sh_list: + for n, curr_bin in enumerate(bin_list): + curr_path = getDriftDataPath(curr_bin, sh_ind) + try: + data = np.load(curr_path) + if data.size > 0: + session_durations[n].append(np.max(data[:, 0])) + all_amplitudes.append( + data[:, 2] + ) # Append only the amplitude column + except FileNotFoundError: + print(f"Warning: File not found for shank {sh_ind}, bin {curr_bin}") + global_set_dur = np.array( + [np.max(durs) if durs else 0 for durs in session_durations] + ) + global_segment_endpoints = np.cumsum(global_set_dur) + + if all_amplitudes: + all_amplitudes = np.concatenate(all_amplitudes) + global_c_lim = np.asarray( + [np.quantile(all_amplitudes, 0.05), np.quantile(all_amplitudes, 0.95)] + ) + global_c_lim = even_divisor * np.floor(global_c_lim / even_divisor) + else: + global_c_lim = [0, 100] + + scatter = None + sh_inds = [0, 1, 2, 3] + for sh_ind in sh_inds: + row = sh_ind // 2 + col = sh_ind % 2 + current_ax = ax[row, col] + current_ax.set_title(f"Shank {sh_ind}", fontsize=14) + if sh_ind not in sh_list: + continue + + # load data + n_set = len(bin_list) + total_spike = 0 + set_dur = np.zeros((n_set,)) + dd_list = list() + for n, curr_bin in enumerate(bin_list): + curr_path = getDriftDataPath(curr_bin, sh_ind) + dd_list.append(np.load(curr_path)) + n_spike, n_meas = dd_list[n].shape + total_spike = total_spike + n_spike + if total_spike == 0: + continue + set_dur[n] = np.max(dd_list[n][:, 0]) + skip_step = np.floor(total_spike / max_spike) + 1 + + if total_spike == 0: + continue + + # loop over the arrays, buid sampled array that covers all the datasets + samp_spike = np.zeros((max_spike, n_meas)) + n_samp = 0 + boundaries = np.concatenate(([0], global_segment_endpoints[:-1])) + for n, dd in enumerate(dd_list): + curr_nspike = dd.shape[0] + curr_ind = np.arange(0, curr_nspike, skip_step).astype(int) + curr_samp = dd[curr_ind, :] + curr_samp[:, 0] = curr_samp[:, 0] + boundaries[n] + curr_samp[:, 1] = curr_samp[:, 1] + np.sum(drift_list[0 : n + 1]) + n_curr = curr_ind.size + samp_spike[n_samp : n_samp + n_curr, :] = curr_samp + n_samp = n_samp + n_curr + + scatter = current_ax.scatter( + samp_spike[:, 0], + samp_spike[:, 1], + s=0.1, + c=samp_spike[:, 2], + cmap="plasma", + vmin=global_c_lim[0], + vmax=global_c_lim[1], + ) + + for single_ax in ax.flat: + single_ax.tick_params(axis="both", labelsize=12) + ymin, ymax = single_ax.get_ylim() + single_ax.vlines( + global_segment_endpoints[:-1], + ymin=ymin, + ymax=ymax, + color="black", + linestyle="solid", + alpha=0.6, + zorder=5, # Draw lines behind data points but in front of the grid + ) + if global_segment_endpoints.size > 0: + single_ax.set_xlim(0, global_segment_endpoints[-1]) + bottom_ax = ax[1, 0] + all_boundaries = np.concatenate(([0], global_segment_endpoints)) + label_positions = (all_boundaries[:-1] + all_boundaries[1:]) / 2 + labels = [f"{d}" for d in day_list] + bottom_ax.set_xticks(label_positions) + bottom_ax.set_xticklabels(labels) + # Add a single X-axis label for the whole figure + fig.supxlabel("Recording Day", y=0.02, fontsize=14) + fig.supylabel("Distance from tip (µm)", x=0.08, fontsize=14) + + cbar = fig.colorbar(scatter, ax=ax.ravel().tolist(), pad=0.01, aspect=40) + cbar.set_label("Amplitude (µV)", rotation=270, labelpad=15, fontsize=14) + + return fig + + +def readData( + binFullPath, selected_sh, time_sec=300, excl_chan=[127], thresh=-80, overwrite=False +): + parent_path = binFullPath.parent + dd_path = getDriftDataPath(binFullPath, selected_sh) + save_path = parent_path.joinpath(dd_path) + if save_path.exists() and not overwrite: + return + # Read in metadata; returns a dictionary with string for values + meta_path = str(binFullPath).replace(".bin", ".meta") + meta = read_meta(meta_path) + + # plan to detect peaks in the last 5 minutes of recording + sRate = get_sample_rate(meta) + n_ap, n_lf, n_sync = ChannelCountsIM(meta) + APChan0_to_uV = ChanGainsIM(meta)[2] + thresh_bits = thresh / APChan0_to_uV + x_coord, z_coord, sh_ind, connected, n_chan_tot = MetaToCoords( + binFullPath.with_suffix(".meta"), -1 + ) + rawData = get_data_memmap(binFullPath, meta) + # transpose + rawData = rawData.T + nChan, nFileSamp = rawData.shape + + start_samp = np.floor(nFileSamp - time_sec * sRate).astype(int) + if start_samp < 0: + start_samp = 0 + batch_samp = np.floor(2 * sRate).astype(int) + n_batch = np.floor((nFileSamp - start_samp) / batch_samp).astype(int) + sel_sh_ind = np.where(sh_ind == selected_sh)[0] + # translate excluded channels for this shank + ex_sh = list() + for ch in excl_chan: + if np.sum(sel_sh_ind == ch) > 0: + ex_sh.append(np.where(sel_sh_ind == ch)[0][0]) + + x_sh = x_coord[sel_sh_ind] + z_sh = z_coord[sel_sh_ind] + + if n_batch == 0: + print("Warning: not enough data to process") + for j in tqdm( + range(n_batch), + desc=f"\tExtracting threshold events: sh {selected_sh}", + leave=False, + ): + st = start_samp + j * batch_samp + cb = rawData[sel_sh_ind, st : st + batch_samp] + peak_ind, peak_chan, peak_sig, xz = findPeaks( + cb, thresh_bits, x_sh, z_sh, ex_sh, fs=sRate + ) + # add offset to peak_ind, concatenate onto set + offset = st - start_samp + peak_ind = (peak_ind + offset) / sRate # convert the times to sec + peak_sig = abs(peak_sig * APChan0_to_uV) + if j == 0: + all_spikes = np.vstack( + (peak_ind, xz[:, 1].T, peak_sig, xz[:, 0].T, peak_chan) + ).T + else: + curr_spikes = np.vstack( + (peak_ind, xz[:, 1].T, peak_sig, xz[:, 0].T, peak_chan) + ).T + all_spikes = np.concatenate((all_spikes, curr_spikes)) + + np.save(save_path, all_spikes) + + +def calc_pdf(bin_path, sh_list): + # calculate prob density function across all shanks + for j, sh_ind in enumerate(sh_list): + curr_path = getDriftDataPath(bin_path, sh_ind) + curr_dat = np.load(curr_path) + if len(curr_dat) == 0: + sum_hist = None + continue + if j == 0: + # calculate bin width and bin edgeds + amp_sort = np.sort(curr_dat[:, 2]) + # if no data just do zeros + bin_width = np.unique(np.diff(amp_sort))[1] + bin_edges = (0.5 * bin_width) + np.arange(0, 1000, bin_width) + sum_hist = np.histogram(curr_dat[:, 2], bin_edges)[0] + else: + if sum_hist is None: + amp_sort = np.sort(curr_dat[:, 2]) + # if no data just do zeros + bin_width = np.unique(np.diff(amp_sort))[1] + bin_edges = (0.5 * bin_width) + np.arange(0, 1000, bin_width) + sum_hist = np.histogram(curr_dat[:, 2], bin_edges)[0] + else: + sum_hist = sum_hist + np.histogram(curr_dat[:, 2], bin_edges)[0] + # convert to pdf by normalizing + npts = len(sum_hist) + pdf = np.zeros((npts, 2)) + pdf[:, 0] = bin_edges[0:npts] + pdf[:, 1] = sum_hist / (np.sum(sum_hist) * bin_width) + # plt.plot(pdf[:,0],pdf[:,1]) + return pdf + + +def calcRate(bin_path, sh_list): + # calculate total spike rate for the probe + for j, sh_ind in enumerate(sh_list): + curr_path = getDriftDataPath(bin_path, sh_ind) + curr_dat = np.load(curr_path) + if len(curr_dat) == 0: + min_time = None + max_time = None + total_count = 0 + continue + if j == 0: + # calculate bin width and bin edge + min_time = np.min(curr_dat[:, 0]) + max_time = np.max(curr_dat[:, 0]) + total_count = curr_dat.shape[0] + else: + if min_time is None: + min_time = np.min(curr_dat[:, 0]) + else: + min_time = np.min([min_time, np.min(curr_dat[:, 0])]) + if max_time is None: + max_time = np.max(curr_dat[:, 0]) + else: + max_time = np.max([max_time, np.max(curr_dat[:, 0])]) + total_count = total_count + curr_dat.shape[0] + # convert rate + spike_rate = total_count / (max_time - min_time) + + return spike_rate + + +def adjust_figure_for_legend(figure, legend, bottom_margin=0.05): + """ + Adjust the figure height to make sure the legend fits. + + Args: + figure (matplotlib.figure.Figure): The figure object. + legend (matplotlib.legend.Legend): The legend object. + bottom_margin (float): The desired margin below the legend in inches. + """ + # Draw the canvas to get the final rendered size of the legend + figure.canvas.draw() + + # Get the bounding box of the legend in pixels + legend_bbox = legend.get_window_extent() + + # If the bottom of the legend is below the figure (y=0) + if legend_bbox.y0 < 0: + # Calculate the overflow in pixels + overflow_pixels = -legend_bbox.y0 + + # Get the figure's DPI (dots per inch) + dpi = figure.get_dpi() + + # Calculate the required additional height in inches + margin_pixels = bottom_margin * dpi + required_height_increase = (overflow_pixels + margin_pixels) / dpi + + # Get the current figure size in inches + current_width, current_height = figure.get_size_inches() + + # Set the new figure size + figure.set_size_inches(current_width, current_height + required_height_increase) + + # Optional: Redraw the figure to apply changes + figure.canvas.draw() + + +def plotMultPDF(bin_list, sh_list, day_list): + fig, ax = plt.subplots(figsize=(8, 5)) + n_pdf = len(bin_list) + + cmap = mpl.colormaps["winter"] + colors = cmap(np.linspace(0, 1, n_pdf)) + + for j in range(n_pdf): + curr_pdf = calc_pdf(bin_list[j], sh_list) + ax.plot(curr_pdf[:, 0], curr_pdf[:, 1], color=colors[j], label=f"{day_list[j]}") + + ax.set_xlim(0, 500) + ax.tick_params(axis="both", labelsize=12) + + plt.xlabel("Spike Amplitude (µV)", fontsize=14) + ax.set_ylabel("Probability Density", fontsize=14) + ax.set_title("Spike Amplitude Distribution Over Time", fontsize=16, pad=10) + + # Add a legend to identify the lines + legend = ax.legend(title="Recording Day", fontsize=10) + + # Remove top and right plot borders for a cleaner look + ax.spines["top"].set_visible(False) + ax.spines["right"].set_visible(False) + adjust_figure_for_legend(fig, legend) + fig.tight_layout() + return fig + + +def plotSpikeRate(bin_list, day_list, sh_list): + n_meas = len(bin_list) + spike_rates = np.zeros((n_meas,)) + + for j in range(n_meas): + spike_rates[j] = calcRate(bin_list[j], sh_list) + trend_x = np.asarray([min(day_list), max(day_list)]) + trend_fit = np.poly1d(np.polyfit(np.asarray(day_list), spike_rates, 1)) + trend_y = trend_fit(trend_x) + + fig, ax = plt.subplots(figsize=(7, 5)) + plt.scatter(day_list, spike_rates, s=3) + plt.plot(trend_x, trend_y, marker=None, linewidth=1, linestyle="dashed") + ax.tick_params(axis="both", labelsize=12) + plt.xlabel("Recording Day", fontsize=14) + plt.ylabel("Spike Rate (Hz)", fontsize=14) + ax.spines["top"].set_visible(False) + ax.spines["right"].set_visible(False) + return fig + + +def calcFiringVsZ(binFullPath, selected_sh, bin_width, z_edges): + # plot the relative firing rate vs. channel + curr_path = getDriftDataPath(binFullPath, selected_sh) + curr_dat = np.load(curr_path) + + if len(curr_dat) == 0: + return None, None + + if z_edges is None: + min_z = bin_width * np.floor(np.min(curr_dat[:, 1]) / bin_width) + max_z = bin_width * np.ceil(np.max(curr_dat[:, 1]) / bin_width) + z_edges = np.arange(min_z, max_z, bin_width) + + z_hist = (np.histogram(curr_dat[:, 1], z_edges)[0]).astype("float64") + z_rel = z_hist / np.sum(z_hist) + + return z_rel, z_edges + + +def plotRateVsZ(bin_list, day_list, sh_list): + bin_width = 15 + n_meas = len(bin_list) + + # Pre-calc for scaling + global_max_rate = 0 + z_edges = None + for sh_ind in sh_list: + for b in bin_list: + rates, z_edges = calcFiringVsZ(b, sh_ind, bin_width, z_edges) + global_max_rate = max(global_max_rate, np.max(rates)) + n_bin = len(z_edges) - 1 + bin_center = z_edges[0:n_bin] + bin_width / 2 + + c_lim = [0, global_max_rate] + c_range = c_lim[1] - c_lim[0] + + # Setup plotting + original_cmap = mpl.colormaps["plasma"] + cmap_colors = original_cmap(np.linspace(0, 1, 256)) + cmap_colors[0] = (1, 1, 1, 1) # RGBA for white + custom_cmap = ListedColormap(cmap_colors) + even_divisor = 10 # z limits will be integer multiples of this value + + fig, axes = plt.subplots( + nrows=2, + ncols=2, + figsize=(0.5 + 1 * len(bin_list), 8), + sharex=True, + sharey=False, + ) + fig.suptitle("Firing Rate vs Depth Across Days", fontsize=18) + + for sh_ind in [0, 1, 2, 3]: + row = sh_ind // 2 + col = sh_ind % 2 + ax = axes[row, col] + ax.set_title(f"Shank {sh_ind}", fontsize=14) + if sh_ind not in sh_list: + continue + + rel0, z_edges = calcFiringVsZ(bin_list[0], sh_ind, bin_width, z_edges) + + rel_rates = np.zeros((n_meas * n_bin,)) + x_vals = np.zeros((n_meas * n_bin)) + z_vals = np.zeros((n_meas * n_bin)) + + rel_rates[0:n_bin] = rel0 + z_vals[0:n_bin] = bin_center + + for i in range(1, n_meas): + x_vals[i * n_bin : (i + 1) * n_bin] = i + z_vals[i * n_bin : (i + 1) * n_bin] = bin_center + rel_rates[i * n_bin : (i + 1) * n_bin] = calcFiringVsZ( + bin_list[i], sh_ind, bin_width, z_edges + )[0] + + min_z = even_divisor * np.floor(np.min(z_vals) / even_divisor) + max_z = even_divisor * np.floor(np.max(z_vals) / even_divisor) + ax.set_ylim([min_z - 15, max_z + 15]) + + # these points get covered by the patches; coloring with rel_rates + # creates teh correct colorbar + scatter = ax.scatter( + x_vals, + z_vals, + c=rel_rates, + s=2, + marker="s", + cmap=custom_cmap, + vmin=c_lim[0], + vmax=c_lim[1], + ) + + # Add rectangles + width = 0.75 # in 'recording day' units + height = 15 # in 'um' + + for i in range(len(x_vals)): + color_val = (rel_rates[i] - c_lim[0]) / c_range + ax.add_patch( + Rectangle( + xy=(x_vals[i] - width / 2, z_vals[i] - height / 2), + width=width, + height=height, + edgecolor="None", + facecolor=custom_cmap(color_val), + ) + ) + + xt_range = np.arange(n_meas) + xt_labels = np.asarray(day_list).astype("str") + bottom_ax = axes[1, 0] + bottom_ax.set_xlim([-0.5, n_meas - 0.5]) + bottom_ax.set_xticks(xt_range) + bottom_ax.set_xticklabels(xt_labels) + fig.supxlabel("Recording Day", y=0.02, fontsize=14) + fig.supylabel("Distance from tip (µm)", x=0.08, fontsize=14) + ax.tick_params(axis="both", labelsize=12) + cbar = fig.colorbar(scatter, ax=axes.ravel().tolist(), pad=0.01, aspect=40) + cbar.set_label("Spiking Rate (Hz)", rotation=270, labelpad=15, fontsize=14) + return fig + + +def get_available_shanks(binFullPath): + binFullPath = Path(binFullPath) + meta_path = binFullPath.with_suffix(".meta") + if not meta_path.exists(): + raise Warning("Metadata file not found") + _, _, sh_ind, connected, _ = MetaToCoords(meta_path, -1) + all_shanks = np.unique(sh_ind) + active_shanks = [] + for shank_idx in all_shanks: + channels_on_shank = sh_ind == shank_idx + if np.any(connected[channels_on_shank]): + active_shanks.append(int(shank_idx)) + return sorted(active_shanks) + + +def plot_threshold_events( + subject_folder, + implant_day=None, + recordings=None, + probe_ids=None, + overwrite=False, + save_dir=None, + save_type="png", + ks_version="4", + analysis_time_sec=300, + excl_chan=[127], + threshold=-80, +): + if save_dir is not None: + os.makedirs(save_dir, exist_ok=True) + # only affects plotting in plotMult raster plots + run_folders = get_run_folders(subject_folder, day_folders=recordings) + drift_list = np.zeros(len(run_folders) - 1) + # group folders by probes + ks_folders = [] + for folder in run_folders: + ks_folders.extend(get_ks_folders(folder, ks_version)) + + all_probe_folders = get_probe_folders(ks_folders) + if probe_ids is not None: + all_probe_folders = { + probe_id: all_probe_folders[probe_id] + for probe_id in probe_ids + if probe_id in all_probe_folders + } + + for probe_num in tqdm( + all_probe_folders, "Processing threshold_events...", position=0, unit="probe" + ): + all_probe_ks_folders = all_probe_folders[probe_num] + probe_ks_folders = get_same_channel_positions(all_probe_ks_folders) + days = [ + datetime.strptime( + os.path.basename(os.path.dirname(folder)).split("_")[0], "%Y%m%d" + ) + for folder in probe_ks_folders + ] + if implant_day is None: + start_day = days[0] + else: + start_day = datetime.strptime(implant_day, "%Y%m%d") + day_list = [(day - start_day).days for day in days] + + probe_bins = [ + Path(get_binary_path(ks_folder)) for ks_folder in probe_ks_folders + ] + sh_list = get_available_shanks(probe_bins[0]) + + for binFullPath in tqdm( + probe_bins, desc="\tProcessing recordings", leave=True, position=1 + ): + for sh_ind in sh_list: + readData( + binFullPath, + sh_ind, + analysis_time_sec, + excl_chan, + threshold, + overwrite, + ) + plt_fig1 = True + plt_fig2 = True + plt_fig3 = True + plt_fig4 = True + + if save_dir is not None and not overwrite: + fname1 = os.path.join( + save_dir, f"multi_shank_drift_imec{probe_num}.{save_type}" + ) + if os.path.exists(fname1): + # load saved drift data and plot + fig1 = plt.figure(figsize=(8, 5)) + img = plt.imread(fname1) + plt.imshow(img) + plt.axis("off") + plt_fig1 = False + fname2 = os.path.join( + save_dir, f"multi_spike_rate_depth_imec{probe_num}.{save_type}" + ) + if os.path.exists(fname2): + fig2 = plt.figure(figsize=(7, 5)) + img = plt.imread(fname2) + plt.imshow(img) + plt.axis("off") + plt_fig2 = False + fname3 = os.path.join( + save_dir, f"multi_spike_amplitude_imec{probe_num}.{save_type}" + ) + if os.path.exists(fname3): + fig3 = plt.figure(figsize=(8, 5)) + img = plt.imread(fname3) + plt.imshow(img) + plt.axis("off") + plt_fig3 = False + fname4 = os.path.join( + save_dir, f"multi_spike_rate_imec{probe_num}.{save_type}" + ) + if os.path.exists(fname4): + fig4 = plt.figure(figsize=(7, 5)) + img = plt.imread(fname4) + plt.imshow(img) + plt.axis("off") + plt_fig4 = False + + if plt_fig1: + if len(probe_bins) == 1: + drift_path = getDriftDataPath( + probe_bins[0], sh_list[0] + ) # TODO update function + drift_data = np.load(probe_bins[0].parent.joinpath(drift_path)) + fig1 = plotOne(drift_data) + else: + fig1 = plotMult(probe_bins, drift_list, day_list, sh_list) + if save_dir is not None: + fname1 = os.path.join( + save_dir, f"multi_shank_drift_imec{probe_num}.{save_type}" + ) + fig1.savefig(fname1, dpi=300, format=save_type, bbox_inches="tight") + + if plt_fig2: + fig2 = plotRateVsZ(probe_bins, day_list, sh_list) + if save_dir is not None: + fname2 = os.path.join( + save_dir, f"multi_spike_rate_depth_imec{probe_num}.{save_type}" + ) + fig2.savefig(fname2, dpi=300, format=save_type, bbox_inches="tight") + + if plt_fig3: + fig3 = plotMultPDF(probe_bins, sh_list, day_list) + fname3 = os.path.join( + save_dir, f"multi_spike_amplitude_imec{probe_num}.{save_type}" + ) + if save_dir is not None: + fig3.savefig(fname3, dpi=300, format=save_type, bbox_inches="tight") + + if plt_fig4: + fig4 = plotSpikeRate(probe_bins, day_list, sh_list) + if save_dir is not None: + fname4 = os.path.join( + save_dir, f"multi_spike_rate_imec{probe_num}.{save_type}" + ) + fig4.savefig(fname4, dpi=300, format=save_type, bbox_inches="tight") diff --git a/npx_utils/stability/unit_counts.py b/npx_utils/stability/unit_counts.py new file mode 100644 index 0000000..98ba293 --- /dev/null +++ b/npx_utils/stability/unit_counts.py @@ -0,0 +1,95 @@ +import os +from datetime import datetime + +import matplotlib.pyplot as plt +import npx_utils as npx +import numpy as np +import pandas as pd +from tqdm.autonotebook import tqdm + + +def plot_subject_unit_counts(ks_folders, day_list): + good_counts = np.zeros(len(ks_folders)) + non_noise_counts = np.zeros(len(ks_folders)) + + for i, ks_folder in enumerate(ks_folders): + label_path = os.path.join(ks_folder, "cluster_group.tsv") + labels = pd.read_csv(label_path, sep="\t", index_col="cluster_id") + num_good = (labels["label"] == "good").sum() + num_non_noise = (labels["label"] != "noise").sum() + good_counts[i] = num_good + non_noise_counts[i] = num_non_noise + + fig, ax = plt.subplots(figsize=(6, 4)) + ax.plot(day_list, good_counts, label=f"good") + ax.plot(day_list, non_noise_counts, "--", label=f"non-noise") + ax.set_ylim(bottom=0) + ax.set_xlabel("Time (days)") + ax.set_ylabel("Number of units") + ax.legend() + return fig, ax + + +def plot_unit_counts( + subject_folder, + implant_day, + probe_ids=None, + save_dir=None, + overwrite=False, + save_type="png", + ks_version="4", +): + if save_dir is not None: + os.makedirs(save_dir, exist_ok=True) + # only affects plotting in plotMult raster plots + run_folders = npx.get_run_folders(subject_folder) + + # group folders by probes + ks_folders = [] + for folder in run_folders: + ks_folders.extend(npx.get_ks_folders(folder, ks_version)) + + all_probe_folders = npx.get_probe_folders(ks_folders) + if probe_ids is not None: + all_probe_folders = { + probe_id: all_probe_folders[probe_id] + for probe_id in probe_ids + if probe_id in all_probe_folders + } + + for probe_num in tqdm( + all_probe_folders, "Plotting unit counts...", position=0, unit="probe" + ): + # check if file exists + if save_dir is not None: + fname = os.path.join(save_dir, f"unit_counts_imec{probe_num}.{save_type}") + if os.path.exists(fname) and not overwrite: + # load image and skip + fig = plt.figure(figsize=(6, 4)) + img = plt.imread(fname) + plt.imshow(img) + plt.axis("off") + continue + + all_probe_ks_folders = all_probe_folders[probe_num] + probe_ks_folders = npx.sglx_helpers.get_same_channel_positions( + all_probe_ks_folders + ) + days = [ + datetime.strptime( + os.path.basename(os.path.dirname(folder)).split("_")[0], "%Y%m%d" + ) + for folder in probe_ks_folders + ] + if implant_day is None: + start_day = days[0] + else: + start_day = datetime.strptime(implant_day, "%Y%m%d") + day_list = [(day - start_day).days for day in days] + + fig, ax = plot_subject_unit_counts(probe_ks_folders, day_list) + ax.set_title(f"{os.path.basename(subject_folder)} imec {probe_num} unit counts") + + if save_dir is not None: + fname = os.path.join(save_dir, f"unit_counts_imec{probe_num}.{save_type}") + fig.savefig(fname, dpi=300, format=save_type) diff --git a/setup.py b/setup.py index 0de7d41..a0f4594 100644 --- a/setup.py +++ b/setup.py @@ -3,5 +3,5 @@ setup( name="npx_utils", version="0.1", - packages=find_packages("npx_utils", include=["*"]), + packages=find_packages(), )