From 6516bad3c0be1f34345ed071abeb57f8499c2e3a Mon Sep 17 00:00:00 2001 From: Tate DeWeese <88437490+tatedeweese@users.noreply.github.com> Date: Tue, 9 Sep 2025 11:22:27 -0400 Subject: [PATCH 1/3] add threshold event plots --- npx_utils/__init__.py | 2 +- npx_utils/data_helpers.py | 37 +- npx_utils/metrics.py | 90 +- npx_utils/noise/lfpBandPower.m | 14 +- npx_utils/other_helpers.py | 2 +- npx_utils/sglx/_SGLXMetaToCoords.py | 770 +++++++++++++++++ npx_utils/sglx/__init__,py | 0 npx_utils/sglx/sglx_helpers.py | 360 ++++++++ npx_utils/sglx_helpers.py | 284 ------- npx_utils/stability/__init__.py | 0 npx_utils/stability/threshold_event_plots.py | 848 +++++++++++++++++++ 11 files changed, 2056 insertions(+), 351 deletions(-) create mode 100644 npx_utils/sglx/_SGLXMetaToCoords.py create mode 100644 npx_utils/sglx/__init__,py create mode 100644 npx_utils/sglx/sglx_helpers.py delete mode 100644 npx_utils/sglx_helpers.py create mode 100644 npx_utils/stability/__init__.py create mode 100644 npx_utils/stability/threshold_event_plots.py diff --git a/npx_utils/__init__.py b/npx_utils/__init__.py index 13f7fea..b24bdbf 100644 --- a/npx_utils/__init__.py +++ b/npx_utils/__init__.py @@ -2,4 +2,4 @@ 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..71f60f9 100644 --- a/npx_utils/data_helpers.py +++ b/npx_utils/data_helpers.py @@ -6,7 +6,12 @@ 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.sglx.sglx_helpers import ( + get_bits_to_uV, + get_channel_counts, + get_data_memmap, + read_meta, +) def extract_spikes( @@ -143,17 +148,20 @@ def calc_mean_wf( # 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( @@ -212,17 +220,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( 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/other_helpers.py b/npx_utils/other_helpers.py index 6366be9..c8430c5 100644 --- a/npx_utils/other_helpers.py +++ b/npx_utils/other_helpers.py @@ -5,7 +5,7 @@ from tqdm import tqdm 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): 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..261c456 --- /dev/null +++ b/npx_utils/sglx/sglx_helpers.py @@ -0,0 +1,360 @@ +import os +import re + +import cupy as cp +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=(nChan, nFileSamp), 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..e69de29 diff --git a/npx_utils/stability/threshold_event_plots.py b/npx_utils/stability/threshold_event_plots.py new file mode 100644 index 0000000..0f9518c --- /dev/null +++ b/npx_utils/stability/threshold_event_plots.py @@ -0,0 +1,848 @@ +# -*- 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.cm as cm +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.colors import ListedColormap +from matplotlib.patches import Rectangle +from tqdm import tqdm + +import npx_utils as npx +from npx_utils.sglx._SGLXMetaToCoords import MetaToCoords +from npx_utils.sglx.sglx_helpers import ( + ChanGainsIM, + ChannelCountsIM, + get_data_memmap, + get_sample_rate, + read_meta, +) + + +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): + 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() + + +def plotMult(bin_list, drift_list, day_list): + # build a large sampled array from the n_sets, to look for 'obvious' drift + sh_inds = [0, 1, 2, 3] + fig, ax = plt.subplots(nrows=2, ncols=2, figsize=(13, 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_inds: + 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] + + for i, sh_ind in enumerate(sh_inds): + row = i // 2 + col = i % 2 + current_ax = ax[row, col] + + # 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 + + current_ax.set_title(f"Shank {sh_ind}", fontsize=14) + + 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 + ) + 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.supylabel("Distance from tip (µm)", x=0.06, fontsize=14) + fig.supxlabel("Recording Session", y=0.02, 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) + 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] + + for j in range(n_batch): + 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 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 + 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) + + 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 Session", 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=(12, 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 Session Day", fontsize=14) + fig.supylabel("Distance from tip (µm)", 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_run_folders(subject_folder): + 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) + ] + + 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 + + +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 main(): + # samples for calling the above functions + # note that the file paths, etc are specific, alter to match your data + subject_folder = r"D:\Psilocybin\Cohort_3\T22" + overwrite = False + save_dir = None # os.path.join(subject_folder, "stability") # None if show + save_type = "png" # svg + ks_version = "4" + # TODO add drift later + b_recalc = True + analysis_time_sec = ( + 300 # readData extracts spikes in the last analysis_time_sec of the recording + ) + excl_chan = [127] + threshold = -80 + probe_ids = None # get all + + 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) + drift_list = np.zeros(len(run_folders) - 1) + days = [ + datetime.strptime( + os.path.basename(os.path.dirname(run_folder)).split("_")[0], "%Y%m%d" + ) + for run_folder in run_folders + ] + day_list = [(day - days[0]).days for day in days] + # 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 + } + + pbar1 = tqdm(all_probe_folders, "Processing probes...", position=0) + for probe_num in pbar1: + pbar1.set_description(f"Processing probe {probe_num}") + all_probe_ks_folders = all_probe_folders[probe_num] + probe_ks_folders = npx.sglx.sglx_helpers.get_same_channel_positions( + all_probe_ks_folders + ) + # get indices of removed ks_folders + removed_indices = [ + i + for i, folder in enumerate(all_probe_ks_folders) + if folder not in probe_ks_folders + ] + # remove from day_list + updated_day_list = [ + day for i, day in enumerate(day_list) if i not in removed_indices + ] + + probe_bins = [ + Path(npx.get_binary_path(ks_folder)) for ks_folder in probe_ks_folders + ] + sh_list = get_available_shanks(probe_bins[0]) + + for sh_ind in sh_list: + if b_recalc: + pbar2 = tqdm( + probe_bins, + desc=f"\tProcessing Shank {sh_ind}", + leave=False, + position=1, + ) + for binFullPath in pbar2: + readData( + binFullPath, + sh_ind, + analysis_time_sec, + excl_chan, + threshold, + overwrite, + ) + if len(probe_bins) == 1: + out_name = f"drift_data_sh{sh_ind}.npy" + drift_data = np.load(probe_bins[0].parent.joinpath(out_name)) + plotOne(drift_data) + + else: + # load saved drift data and plot + fig1 = plotMult(probe_bins, drift_list, updated_day_list) + fig2 = plotRateVsZ(probe_bins, updated_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) + + fname2 = os.path.join( + save_dir, f"multi_spike_rate_depth_imec{probe_num}.{save_type}" + ) + fig2.savefig(fname2, dpi=300, format=save_type) + + fig3 = plotMultPDF(probe_bins, sh_list, updated_day_list) + fig4 = plotSpikeRate(probe_bins, updated_day_list, sh_list) + if save_dir is not None: + fname3 = os.path.join( + save_dir, f"multi_spike_amplitude_imec{probe_num}.{save_type}" + ) + fig3.savefig(fname3, dpi=300, format=save_type) + fname4 = os.path.join( + save_dir, f"multi_spike_rate_imec{probe_num}.{save_type}" + ) + fig4.savefig(fname4, dpi=300, format=save_type) + else: + plt.show() + + +if __name__ == "__main__": + main() From ede71ff3609a8cd8e8783a47a8da912a9d33c30c Mon Sep 17 00:00:00 2001 From: Tate DeWeese <88437490+tatedeweese@users.noreply.github.com> Date: Thu, 11 Sep 2025 15:41:39 -0400 Subject: [PATCH 2/3] updated plots to make prettier --- npx_utils/other_helpers.py | 52 ++++++++ npx_utils/stability/threshold_event_plots.py | 132 ++++++++++--------- npx_utils/stability/unit_counts.py | 97 ++++++++++++++ 3 files changed, 217 insertions(+), 64 deletions(-) create mode 100644 npx_utils/stability/unit_counts.py diff --git a/npx_utils/other_helpers.py b/npx_utils/other_helpers.py index c8430c5..582ec4b 100644 --- a/npx_utils/other_helpers.py +++ b/npx_utils/other_helpers.py @@ -4,6 +4,8 @@ from tqdm import tqdm +import npx_utils as npx + from .ks_helpers import get_meta_path, get_probe_id from .sglx.sglx_helpers import read_meta @@ -80,3 +82,53 @@ def get_details(ks_folder, drug_dict=None): "drug": drug, } return details + + +def get_run_folders(subject_folder): + 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) + ] + + 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/stability/threshold_event_plots.py b/npx_utils/stability/threshold_event_plots.py index 0f9518c..f37df3f 100644 --- a/npx_utils/stability/threshold_event_plots.py +++ b/npx_utils/stability/threshold_event_plots.py @@ -39,7 +39,6 @@ from pathlib import Path import matplotlib as mpl -import matplotlib.cm as cm import matplotlib.pyplot as plt import numpy as np from matplotlib.colors import ListedColormap @@ -47,6 +46,7 @@ from tqdm import tqdm import npx_utils as npx +from npx_utils import get_run_folders from npx_utils.sglx._SGLXMetaToCoords import MetaToCoords from npx_utils.sglx.sglx_helpers import ( ChanGainsIM, @@ -256,9 +256,18 @@ def plotOne(drift_data): def plotMult(bin_list, drift_list, day_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 sh_inds = [0, 1, 2, 3] - fig, ax = plt.subplots(nrows=2, ncols=2, figsize=(13, 8), sharex=True, sharey=False) + 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 @@ -354,6 +363,8 @@ def plotMult(bin_list, drift_list, day_list): 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 @@ -361,11 +372,12 @@ def plotMult(bin_list, drift_list, day_list): bottom_ax.set_xticks(label_positions) bottom_ax.set_xticklabels(labels) # Add a single X-axis label for the whole figure - fig.supylabel("Distance from tip (µm)", x=0.06, fontsize=14) + fig.supylabel("Distance from tip (µm)", x=0.08, fontsize=14) fig.supxlabel("Recording Session", y=0.02, 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 @@ -494,6 +506,43 @@ def calcRate(bin_path, sh_list): 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) @@ -513,12 +562,13 @@ def plotMultPDF(bin_list, sh_list, day_list): ax.set_title("Spike Amplitude Distribution Over Time", fontsize=16, pad=10) # Add a legend to identify the lines - ax.legend(title="Recording Day", fontsize=10) + 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 @@ -587,7 +637,11 @@ def plotRateVsZ(bin_list, day_list, sh_list): even_divisor = 10 # z limits will be integer multiples of this value fig, axes = plt.subplots( - nrows=2, ncols=2, figsize=(12, 8), sharex=True, sharey=False + 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) @@ -654,64 +708,14 @@ def plotRateVsZ(bin_list, day_list, sh_list): 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 Session Day", fontsize=14) - fig.supylabel("Distance from tip (µm)", fontsize=14) + fig.supxlabel("Recording Session 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_run_folders(subject_folder): - 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) - ] - - 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 - - def get_available_shanks(binFullPath): binFullPath = Path(binFullPath) meta_path = binFullPath.with_suffix(".meta") @@ -730,9 +734,9 @@ def get_available_shanks(binFullPath): def main(): # samples for calling the above functions # note that the file paths, etc are specific, alter to match your data - subject_folder = r"D:\Psilocybin\Cohort_3\T22" + subject_folder = r"D:\Psilocybin\Cohort_1\T08" overwrite = False - save_dir = None # os.path.join(subject_folder, "stability") # None if show + save_dir = os.path.join(subject_folder, "stability") # None if show save_type = "png" # svg ks_version = "4" # TODO add drift later @@ -822,12 +826,12 @@ def main(): fname1 = os.path.join( save_dir, f"multi_shank_drift_imec{probe_num}.{save_type}" ) - fig1.savefig(fname1, dpi=300, format=save_type) + fig1.savefig(fname1, dpi=300, format=save_type, bbox_inches="tight") fname2 = os.path.join( save_dir, f"multi_spike_rate_depth_imec{probe_num}.{save_type}" ) - fig2.savefig(fname2, dpi=300, format=save_type) + fig2.savefig(fname2, dpi=300, format=save_type, bbox_inches="tight") fig3 = plotMultPDF(probe_bins, sh_list, updated_day_list) fig4 = plotSpikeRate(probe_bins, updated_day_list, sh_list) @@ -835,11 +839,11 @@ def main(): fname3 = os.path.join( save_dir, f"multi_spike_amplitude_imec{probe_num}.{save_type}" ) - fig3.savefig(fname3, dpi=300, format=save_type) + fig3.savefig(fname3, dpi=300, format=save_type, bbox_inches="tight") fname4 = os.path.join( save_dir, f"multi_spike_rate_imec{probe_num}.{save_type}" ) - fig4.savefig(fname4, dpi=300, format=save_type) + fig4.savefig(fname4, dpi=300, format=save_type, bbox_inches="tight") else: plt.show() diff --git a/npx_utils/stability/unit_counts.py b/npx_utils/stability/unit_counts.py new file mode 100644 index 0000000..80d5ec6 --- /dev/null +++ b/npx_utils/stability/unit_counts.py @@ -0,0 +1,97 @@ +import os +from datetime import datetime + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +from tqdm import tqdm + +import npx_utils as npx + + +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 + + +def main(): + # samples for calling the above functions + # note that the file paths, etc are specific, alter to match your data + subject_folder = r"D:\Psilocybin\Cohort_1\T08" + overwrite = False + save_dir = os.path.join(subject_folder, "stability") # None if show + save_type = "png" # svg + ks_version = "4" + probe_ids = None # get all + + 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) + days = [ + datetime.strptime( + os.path.basename(os.path.dirname(run_folder)).split("_")[0], "%Y%m%d" + ) + for run_folder in run_folders + ] + day_list = [(day - days[0]).days for day in days] + + # 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 + } + + pbar1 = tqdm(all_probe_folders, "Processing probes...", position=0) + for probe_num in pbar1: + pbar1.set_description(f"Processing probe {probe_num}") + all_probe_ks_folders = all_probe_folders[probe_num] + probe_ks_folders = npx.sglx.sglx_helpers.get_same_channel_positions( + all_probe_ks_folders + ) + # get indices of removed ks_folders + removed_indices = [ + i + for i, folder in enumerate(all_probe_ks_folders) + if folder not in probe_ks_folders + ] + # remove from day_list + updated_day_list = [ + day for i, day in enumerate(day_list) if i not in removed_indices + ] + + fig = plot_subject_unit_counts(probe_ks_folders, updated_day_list) + + 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) + else: + plt.show() + + +if __name__ == "__main__": + main() From 4b5d8f4606a8c024bbb97a268ee5427f59f5c1f2 Mon Sep 17 00:00:00 2001 From: Tate DeWeese <88437490+tatedeweese@users.noreply.github.com> Date: Wed, 17 Sep 2025 11:07:56 -0400 Subject: [PATCH 3/3] starting to fix plot_one --- npx_utils/other_helpers.py | 22 +++-- npx_utils/stability/threshold_event_plots.py | 87 +++++++++++--------- npx_utils/stability/unit_counts.py | 10 ++- 3 files changed, 67 insertions(+), 52 deletions(-) diff --git a/npx_utils/other_helpers.py b/npx_utils/other_helpers.py index 582ec4b..e97dea3 100644 --- a/npx_utils/other_helpers.py +++ b/npx_utils/other_helpers.py @@ -84,14 +84,20 @@ def get_details(ks_folder, drug_dict=None): return details -def get_run_folders(subject_folder): - 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) - ] +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: diff --git a/npx_utils/stability/threshold_event_plots.py b/npx_utils/stability/threshold_event_plots.py index f37df3f..74d8061 100644 --- a/npx_utils/stability/threshold_event_plots.py +++ b/npx_utils/stability/threshold_event_plots.py @@ -234,7 +234,7 @@ def calc_neighbor_sites(xc, zc, neigh_radius_um): def plotOne(drift_data): - plt.subplots(figsize=(6, 2)) + 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) @@ -253,6 +253,7 @@ def plotOne(drift_data): vmax=c_lim[1], ) c = plt.colorbar() + return fig def plotMult(bin_list, drift_list, day_list): @@ -372,8 +373,8 @@ def plotMult(bin_list, drift_list, 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) - fig.supxlabel("Recording Session", y=0.02, 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) @@ -586,7 +587,7 @@ def plotSpikeRate(bin_list, day_list, sh_list): 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 Session", fontsize=14) + 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) @@ -708,7 +709,7 @@ def plotRateVsZ(bin_list, day_list, sh_list): 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 Session Day", y=0.02, fontsize=14) + 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) @@ -734,13 +735,14 @@ def get_available_shanks(binFullPath): def main(): # samples for calling the above functions # note that the file paths, etc are specific, alter to match your data - subject_folder = r"D:\Psilocybin\Cohort_1\T08" + subject_folder = r"D:\Psilocybin\Cohort_4\T24" + recordings = ["20250912_T24_site_test_halfsites"] # if you dont want all recordings + implant_day = None # if None date will go based on first recording day overwrite = False - save_dir = os.path.join(subject_folder, "stability") # None if show + save_dir = None # os.path.join(subject_folder, "stability") # None if show save_type = "png" # svg ks_version = "4" # TODO add drift later - b_recalc = True analysis_time_sec = ( 300 # readData extracts spikes in the last analysis_time_sec of the recording ) @@ -751,7 +753,7 @@ def main(): 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) + run_folders = get_run_folders(subject_folder, day_folders=recordings) drift_list = np.zeros(len(run_folders) - 1) days = [ datetime.strptime( @@ -759,7 +761,12 @@ def main(): ) for run_folder in run_folders ] - day_list = [(day - days[0]).days for day in days] + 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] + # group folders by probes ks_folders = [] for folder in run_folders: @@ -773,9 +780,7 @@ def main(): if probe_id in all_probe_folders } - pbar1 = tqdm(all_probe_folders, "Processing probes...", position=0) - for probe_num in pbar1: - pbar1.set_description(f"Processing probe {probe_num}") + for probe_num in tqdm(all_probe_folders, "Processing probes...", position=0): all_probe_ks_folders = all_probe_folders[probe_num] probe_ks_folders = npx.sglx.sglx_helpers.get_same_channel_positions( all_probe_ks_folders @@ -797,44 +802,44 @@ def main(): sh_list = get_available_shanks(probe_bins[0]) for sh_ind in sh_list: - if b_recalc: - pbar2 = tqdm( - probe_bins, - desc=f"\tProcessing Shank {sh_ind}", - leave=False, - position=1, + for binFullPath in tqdm( + probe_bins, + desc=f"\tProcessing Shank {sh_ind}", + leave=False, + position=1, + ): + readData( + binFullPath, + sh_ind, + analysis_time_sec, + excl_chan, + threshold, + overwrite, ) - for binFullPath in pbar2: - readData( - binFullPath, - sh_ind, - analysis_time_sec, - excl_chan, - threshold, - overwrite, - ) if len(probe_bins) == 1: - out_name = f"drift_data_sh{sh_ind}.npy" - drift_data = np.load(probe_bins[0].parent.joinpath(out_name)) - plotOne(drift_data) + 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: # load saved drift data and plot fig1 = plotMult(probe_bins, drift_list, updated_day_list) - fig2 = plotRateVsZ(probe_bins, updated_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") + fig2 = plotRateVsZ(probe_bins, updated_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") - 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") + 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") fig3 = plotMultPDF(probe_bins, sh_list, updated_day_list) - fig4 = plotSpikeRate(probe_bins, updated_day_list, sh_list) + # dfig4 = plotSpikeRate(probe_bins, updated_day_list, sh_list) if save_dir is not None: fname3 = os.path.join( save_dir, f"multi_spike_amplitude_imec{probe_num}.{save_type}" diff --git a/npx_utils/stability/unit_counts.py b/npx_utils/stability/unit_counts.py index 80d5ec6..a8f4e5c 100644 --- a/npx_utils/stability/unit_counts.py +++ b/npx_utils/stability/unit_counts.py @@ -34,8 +34,8 @@ def plot_subject_unit_counts(ks_folders, day_list): def main(): # samples for calling the above functions # note that the file paths, etc are specific, alter to match your data - subject_folder = r"D:\Psilocybin\Cohort_1\T08" - overwrite = False + subject_folder = r"D:\Psilocybin\Cohort_3\T22" + implant_day = "20250708" # if None date will go based on first recording day save_dir = os.path.join(subject_folder, "stability") # None if show save_type = "png" # svg ks_version = "4" @@ -51,7 +51,11 @@ def main(): ) for run_folder in run_folders ] - day_list = [(day - days[0]).days for day in days] + 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] # group folders by probes ks_folders = []