OSCR

Cerebellar activity is triggered by reach endpoint during learning of a complex locomotor task.

Code ↔ Paper

23 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 23 matches
  1. [1] § Methods › Paw tracking and swing-stance segementation ↔ extractSwingStancePhases.py, lines 79–158 · score 0.83 · swing stance phases, video recordings, video frame, paw trajectories, paw coordination, wheel speed
  2. [2] § Methods › PSTH analysis ↔ getPsortWalkingActivityAndPawTracesCalculatePSTH.py, lines 151–231 · score 0.82 · 0–20, 40–60, 60–80, 80–100, 20–40, swing duration
  3. [3] § Methods › Clustering of all recorded cells into MLIs and PCs using t-SNE ↔ tools/dataAnalysis_cellClustering.py, lines 74–138 · score 0.82 · early_exaggeration, n_components, design matrix, t-SNE, clustering, neighbor
  4. [4] § Methods › Clustering of all recorded cells into MLIs and PCs using t-SNE ↔ groupAnalysisScripts/clusterPCkernelprofiles.py, lines 41–123 · score 0.82 · early_exaggeration, n_components, design matrix, t-SNE, clustering, neighbor
  5. [5] § Results › Mice refine their step kinematics and paw coordination over learning ↔ tools/createPublicationVisualizations.py, lines 3092–3172 · score 0.80 · Median FR swing, stance onset IQR, single paw, FL paw, stance duration, Swing speed
  6. [6] § Methods › Paw tracking and swing-stance segementation ↔ tools/dataAnalysis.py, lines 1463–1542 · score 0.75 · hind right, hind left, front right, front left, displacement, DLC
  7. [7] § Methods › LocoReach setup ↔ extractSwingStancePhases.py, lines 79–158 · score 0.74 · rotary encoder, behavioral videos, video frames, ACQ4, DAQ, Devices
  8. [8] § Methods › Event kernels and linear models ↔ tools/dataAnalysis.py, lines 2223–2303 · score 0.70 · L2 regularized, Cross validated, full model, L1, coefficient, predicted
  9. [9] § Methods › LocoReach setup ↔ getRawBehaviorImagesSaveVideo.py, lines 56–81 · score 0.69 · rotary encoder, behavioral videos, video frames, Devices, camera, LEDs
  10. [10] § Results › Molecular layer interneurons exhibit rich firing rate dynamics during LocoReach task ↔ tools/createPublicationVisualizations.py, lines 1594–1663 · score 0.67 · baseline CV, walking CV, walking FR, spike firing rates, PC, MLI
  11. [11] § Results › Mice refine their step kinematics and paw coordination over learning ↔ tools/createPublicationVisualizations.py, lines 3092–3172 · score 0.66 · median swing onset, single paw, paw swing, FL paw, swing duration, offset
  12. [12] § Results › LocoReach task allows for finely resolved paw movement analysis in a complex environment ↔ tools/groupAnalysis.py, lines 1158–1247 · score 0.63 · hind right, hind left, front right, front left, HL, HR
  13. [13] § Results › LocoReach task allows for finely resolved paw movement analysis in a complex environment ↔ tools/dataAnalysis.py, lines 1463–1542 · score 0.63 · hind right, hind left, front right, front left, tracked, Video
  14. [14] § Results › LocoReach task allows for finely resolved paw movement analysis in a complex environment ↔ tools/createGroupVisualizations.py, lines 560–667 · score 0.61 · FL speed, FR paw, FL paw, sess, alignments, vertical
  15. [15] § Methods › Event kernels and linear models ↔ tools/dataAnalysis.py, lines 2223–2303 · score 0.60 · scikit-learn, linear models, Lasso, Ridge, L1, L2
  16. [16] § Results › LocoReach task allows for finely resolved paw movement analysis in a complex environment ↔ tools/createPublicationVisualizations.py, lines 3032–3091 · score 0.58 · FL stance onset, correlation coefficient, background, FL paw, wheel speed, swing onset
  17. [17] § Methods › Event kernels and linear models ↔ tools/dataAnalysis.py, lines 2163–2222 · score 0.58 · L2 regularized, Cross validated, L1, model, regressors, variable
  18. [18] § Results › Mice refine their step kinematics and paw coordination over learning ↔ tools/groupAnalysis.py, lines 2383–2489 · score 0.58 · median swing onset, paw cycles, paw swing, offset, variable, FL
  19. [19] § Results › LocoReach task allows for finely resolved paw movement analysis in a complex environment ↔ getPsortWalkingActivityAndPawTraces.py, lines 52–122 · score 0.57 · flat surface, paw trajectory, paw movements, paw speed, Paw position, walking
  20. [20] § Results › LocoReach task allows for finely resolved paw movement analysis in a complex environment ↔ manuscriptFigures/fig_experiment-PawPositionExtraction.py, lines 50–83 · score 0.56 · rotary encoder, paw trajectories, stance duration, paw speed, camera, tracked
  21. [21] § Methods › Statistics ↔ tools/createGroupVisualizations.py, lines 1–88 · score 0.55 · statsmodels, linear models, Scipy, Pearson, ANOVA, squares
  22. [22] § Methods › PSTH analysis ↔ tools/dataAnalysis.py, lines 335–405 · score 0.55 · original spike train, Gaussian, histograms, firing rate, bins, window
  23. [23] § Results › Single MLIs get more entrained during learning and encode step cycle events of multiple paws ↔ tools/dataAnalysis.py, lines 2163–2222 · score 0.53 · L2 regularized, linear regression, sparse, L1, GLM, shifted

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 3,595 lines · 195 KB · gpl · 7 matches

  1. import time
  2. import numpy as np
  3. import sys
  4. import os
  5. import scipy, scipy.io
  6. import matplotlib.pyplot as plt
  7. import matplotlib.patches as patches
  8. import tifffile as tiff
  9. from scipy import io
  10. import pdb
  11. import scipy.ndimage
  12. import itertools
  13. import pandas as pd
  14. from scipy.interpolate import interp1d
  15. from sklearn.decomposition import PCA
  16. import cv2
  17. from scipy import signal
  18. from scipy.signal import find_peaks
  19. import pickle
  20. import random
  21. from statsmodels.stats.anova import anova_single
  22. import scikits.bootstrap as boot
  23. from scipy import ndimage
  24. from matplotlib import rcParams
  25. import matplotlib.pyplot as plt
  26. import matplotlib.gridspec as gridspec
  27. import matplotlib.cm as cm
  28. import matplotlib
  29. import multiprocessing as mp
  30. from joblib import Parallel, delayed
  31. import scipy.stats as stats
  32. from tools.pyqtgraph.Qt import QtGui, QtCore
  33. import tools.pyqtgraph as pg
  34. matplotlib.use('TkAgg') # WxAgg
  35. from array import array
  36. import scipy.interpolate as interpolate
  37. import scipy.optimize as optimize
  38. import array as arr
  39. from numpy import trapz
  40. import numpy as np
  41. import scipy.optimize
  42. import matplotlib.pyplot as plt
  43. from sklearn.cluster import KMeans
  44. numcores = mp.cpu_count()-1
  45. from sklearn import linear_model
  46. from sklearn import preprocessing
  47. from sklearn.model_selection import GroupKFold
  48. from sklearn.model_selection import KFold
  49. from sklearn.model_selection import cross_val_score
  50. from sklearn.model_selection import RepeatedKFold
  51. from scipy.stats import vonmises_line
  52. def getSpeed(angles,times,circumsphere,minSpacing):
  53. angleJumps = angles[np.concatenate((([True]),np.diff(angles)!=0.))] # find angles at the points where the angle value changes
  54. timePoints = times[np.concatenate((([True]),np.diff(angles)!=0.))] # find times at the points where the angle value changes
  55. #
  56. #
  57. dt = np.diff(timePoints) # delta t values of the time array
  58. dtMultipleSpace = dt//minSpacing # how many times does the minSpacing fit in the gaps
  59. dtMultipleSpace = dtMultipleSpace[dtMultipleSpace>1] # gap must be as big as 2 times the spacing to add a point
  60. # note that gap must be larger than 2 times the spacing
  61. startGap = np.arange(len(timePoints))[np.hstack(((dt>2.*minSpacing),np.array(False)))] # index values of the start of the gap
  62. endGap = np.arange(len(timePoints))[np.hstack((np.array(False),(dt>2.*minSpacing)))] # index values of the end of the gap
  63. newSpacingValue = (timePoints[endGap] - timePoints[startGap]) / dtMultipleSpace
  64. newTvalues = []
  65. newAvalues = []
  66. for i in range(len(dtMultipleSpace)):
  67. newTvalues.extend(timePoints[startGap][i] + newSpacingValue[i] * np.arange(1, dtMultipleSpace[i]))
  68. newAvalues.extend(np.repeat(angleJumps[startGap][i], (dtMultipleSpace[i] - 1)))
  69. timesNew = np.hstack((timePoints,np.asarray(newTvalues)))
  70. anglesNew = np.hstack((angleJumps,np.asarray(newAvalues)))
  71. both = np.row_stack((timesNew,anglesNew))
  72. bothSorted = both[:,both[0].argsort()]
  73. angularSpeed = np.diff(bothSorted[1])/np.diff(bothSorted[0])
  74. angularSpeedM = (angularSpeed[1:]+angularSpeed[:-1])/2.
  75. linearSpeed = angularSpeedM*circumsphere/360.
  76. speedTimes = bothSorted[0][1:-1]
  77. #pdb.set_trace()
  78. angularSpeedSmooth = scipy.signal.medfilt(angularSpeed,kernel_size=9)
  79. linearSpeedSmooth = scipy.signal.medfilt(linearSpeed,kernel_size=9)
  80. return (angularSpeedSmooth,linearSpeedSmooth,speedTimes,angularSpeed,linearSpeed)
  81. def crosscorr(deltat, y0, y1, correlationRange=1.5, fast=False):
  82. """
  83. home-written routine to calcualte cross-correlation between two contiuous traces
  84. new version from February 9th, 2011
  85. """
  86. if len(y0) != len(y1):
  87. print('Data to be correlated has different dimensions!')
  88. sys.exit(1)
  89. y0mean = y0.mean()
  90. y1mean = y1.mean()
  91. y0sd = y0.std()
  92. y1sd = y1.std()
  93. if y0sd != 0 and y1sd != 0:
  94. y0norm = (y0 - y0mean) / y0sd
  95. y1norm = (y1 - y1mean) / y1sd
  96. else:
  97. y0norm = y0 - y0mean
  98. y1norm = y1 - y1mean
  99. # defined range calculation of cross-correlation
  100. # value is specified in main routine
  101. # deltat = 0.9
  102. pointnumber1 = len(y0)
  103. ncorrrange = np.ceil(correlationRange / deltat)
  104. corrrange = np.arange(2 * ncorrrange + 1) - ncorrrange
  105. ycorr = np.zeros(len(corrrange))
  106. # print corrrange
  107. if fast:
  108. pass
  109. else:
  110. for n in corrrange:
  111. corrpairs = pointnumber1 - abs(n)
  112. # ccc = arange(corrpairs)
  113. # print n
  114. if n < 0:
  115. y1mod = np.hstack((y1norm[int(-abs(n)):], y1norm[:-int(abs(n))]))
  116. ycorr[int(n + ncorrrange)] = (np.add.reduce(y0norm * y1mod)) / (float(pointnumber1))
  117. # if n > -10 :
  118. # print n, ncorrrange, n+ncorrrange, ycorr[n+ncorrrange], float(pointnumber1)
  119. elif n == 0:
  120. ycorr[int(n + ncorrrange)] = (np.add.reduce(y0norm * y1norm)) / (float(pointnumber1))
  121. # print n, ncorrrange, n+ncorrrange, ycorr[n+ncorrrange], (float(pointnumber1-1))
  122. elif n > 0:
  123. y1mod = np.hstack((y1norm[int(abs(n)):], y1norm[:int(abs(n))]))
  124. ycorr[int(n + ncorrrange)] = (np.add.reduce(y0norm * y1mod)) / (float(pointnumber1))
  125. # if n < 10 :
  126. # print n, ncorrrange, n+ncorrrange, ycorr[n+ncorrrange], float(pointnumber1-1)
  127. else:
  128. print('Problem!')
  129. exit(1)
  130. # print n , ycorr[n+ncorrrange]
  131. float_corrrange = np.array([float(i) for i in corrrange])
  132. xcorr = float_corrrange * deltat
  133. normcorr = np.column_stack((xcorr, ycorr))
  134. return normcorr
  135. ############################################################
  136. ## high-pass filter from http://nullege.com/codes/show/[email hidden]-0.3.3@obspy@[email hidden]
  137. ############################################################
  138. def highpass(data, freq, df=200, corners=4, zerophase=False):
  139. """
  140. Butterworth-Highpass Filter.
  141. Filter data removing data below certain frequency freq using corners.
  142. :param data: Data to filter, type numpy.ndarray.
  143. :param freq: Filter corner frequency.
  144. :param df: Sampling rate in Hz; Default 200.
  145. :param corners: Filter corners. Note: This is twice the value of PITSA's
  146. filter sections
  147. :param zerophase: If True, apply filter once forwards and once backwards.
  148. This results in twice the number of corners but zero phase shift in
  149. the resulting filtered trace.
  150. :return: Filtered data.
  151. """
  152. fe = 0.5 * df
  153. [b, a] = iirfilter(corners, freq / fe, btype='highpass', ftype='butter', output='ba')
  154. if zerophase:
  155. firstpass = lfilter(b, a, data)
  156. return lfilter(b, a, firstpass[::-1])[::-1]
  157. else:
  158. return lfilter(b, a, data)
  159. ############################################################
  160. ## high-pass filter from http://stackoverflow.com/questions/12093594/how-to-implement-band-pass-butterworth-filter-with-scipy-signal-butter
  161. ############################################################
  162. def butter_highpass(interval, sampling_rate, cutoff, order=5):
  163. nyq = sampling_rate * 0.5
  164. stopfreq = float(cutoff)
  165. cornerfreq = 0.4 * stopfreq # (?)
  166. ws = cornerfreq / nyq
  167. wp = stopfreq / nyq
  168. # for bandpass:
  169. # wp = [0.2, 0.5], ws = [0.1, 0.6]
  170. N, wn = scipy.signal.buttord(wp, ws, 3, 16) # (?)
  171. # for hardcoded order:
  172. # N = order
  173. b, a = scipy.signal.butter(N, wn, btype='high') # should 'high' be here for bandpass?
  174. sf = scipy.signal.lfilter(b, a, interval)
  175. return sf
  176. ##################################################################
  177. ## high-pass filters the ephys recording and extracts spikes through thresholding
  178. ##################################################################
  179. def extractSpikes(eData, eTime, stim=False):
  180. highpassfreq = 150. # Hz
  181. spikecountwindow = 0.05 # in sec
  182. binWidth = 1.E-3 # in sec
  183. stimRinging = 0.002
  184. dt = np.mean(eTime[1:] - eTime[:-1])
  185. rate = 1. / dt
  186. # set binned array for convolution
  187. binWidth = 1.E-3 # in sec
  188. tbins = np.linspace(0., len(eData) * dt, int(len(eData) * dt / binWidth) + 1)
  189. nspikecountwindow = spikecountwindow / binWidth
  190. ############################################
  191. # create new group in hdf5 file
  192. #grp_spikes = self.analyzed_data.require_group('spiking_data')
  193. detectSpikes = True
  194. #if ('spikeTreshold' in grp_spikes.keys()) and ('artifactTreshold' in grp_spikes.keys()):
  195. # input_ = raw_input('Spike and Artifact detection thresholds exist already. Do you want to re-detect spikes? (\'y\', or any other key for no) : ')
  196. # if input_ != 'y':
  197. # detectSpikes = False
  198. if detectSpikes:
  199. # get time of the stimuls
  200. if stim:
  201. # in case of external stimulation: exclude period of stimuli
  202. stimuli = self.analyzed_data['stimulation_data/stimulus_times'].value
  203. startStim = np.array(stimuli / dt, dtype=int)
  204. endStim = int(stimuli[-1] / dt) + int(stimRinging / dt) # eDataReplaced = copy(eData)
  205. # high-pass filter recording #################################
  206. # eDataHP = self.analysisTools.highpass(eData,highpassfreq,rate,corners=4,zerophase=True)
  207. eDataHP = butter_highpass(eData, rate, highpassfreq, order=4)
  208. #self.h5pyTools.createOverwriteDS(grp_spikes, 'ephys_data_high-pass', eDataHP)
  209. # detect spikes ################################################
  210. app = QtGui.QApplication([])
  211. win = pg.GraphicsWindow(title="Data plotting")
  212. win.resize(1800, 600)
  213. win.setWindowTitle('high-pass filtered recording')
  214. label = pg.LabelItem(justify='right')
  215. win.addItem(label)
  216. pg.setConfigOptions(antialias=True)
  217. x2 = np.linspace(-100, 100, 1000)
  218. data2 = np.sin(x2) / x2
  219. p8 = win.addPlot(row=1, col=0, title="set threshold for spike detection with mouse click")
  220. p8.plot(eDataHP, pen=(255, 255, 255, 200))
  221. # lr = pg.LinearRegionItem([400,700])
  222. # lr.setZValue(-10)
  223. # p8.addItem(lr)
  224. vLine = pg.InfiniteLine(angle=90, movable=False)
  225. hLine = pg.InfiniteLine(angle=0, movable=False)
  226. hLineSpikes = pg.InfiniteLine(angle=0, pen=pg.mkPen(0, 255, 0), movable=False)
  227. hLineArtifacts = pg.InfiniteLine(angle=0, pen=pg.mkPen(255, 0, 0), movable=False)
  228. p8.addItem(vLine, ignoreBounds=True)
  229. p8.addItem(hLine, ignoreBounds=True)
  230. p8.addItem(hLineSpikes, ignoreBounds=True)
  231. p8.addItem(hLineArtifacts, ignoreBounds=True)
  232. vb = p8.vb
  233. # detectionTreshold = empty(0)
  234. def detectSpikeTimes(tresh):
  235. global detectionTreshold
  236. excursion = eDataHP < tresh # threshold ephys trace
  237. excursionInt = np.array(excursion, dtype=int) # convert boolean array into array of zeros and ones
  238. diff = excursionInt[1:] - excursionInt[:-1] # calculate difference
  239. spikeStart = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == 1))] # a difference of one is the start of a spike
  240. spikeEnd = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == -1))] # a difference of -1 is the spike end
  241. if (spikeEnd[0] - spikeStart[0]) < 0.: # if trace starts below threshold
  242. spikeEnd = spikeEnd[1:]
  243. if (spikeEnd[-1] - spikeStart[-1]) < 0.: # if trace ends below threshold
  244. spikeStart = spikeStart[:-1]
  245. if len(spikeStart) != len(spikeEnd): # unequal lenght of starts and ends is a problem of course
  246. print('problem in length of spikeStart and spikeEnd')
  247. sys.exit(1)
  248. spikeT = []
  249. for i in range(len(spikeStart)):
  250. if (spikeEnd[i] - spikeStart[i]) > 10: # ignore if difference between end and start is smaller than 15 points
  251. nMin = np.argmin(eDataHP[spikeStart[i]:spikeEnd[i]]) + spikeStart[i]
  252. spikeT.append(nMin)
  253. # detectionTreshold = tresh
  254. return spikeT
  255. def mouseMoved(evt):
  256. pos = evt[0] ## using signal proxy turns original arguments into a tuple
  257. if p8.sceneBoundingRect().contains(pos):
  258. mousePoint = vb.mapSceneToView(pos)
  259. index = int(mousePoint.x())
  260. if index > 0 and index < len(eDataHP):
  261. label.setText("<span style='font-size: 12pt'>x=%0.1f, <span style='font-size: 12pt'>y=%s</span>" % (mousePoint.x(), eDataHP[index]))
  262. vLine.setPos(mousePoint.x())
  263. hLine.setPos(mousePoint.y())
  264. pointSpikes = [0]
  265. pointArtifacts = [0]
  266. sSpikes = pg.ScatterPlotItem(size=10, pen=pg.mkPen(None), brush=pg.mkBrush(0, 255, 0))
  267. sSpikes.addPoints(x=pointSpikes, y=len(pointSpikes) * [0])
  268. p8.addItem(sSpikes)
  269. sArtifacts = pg.ScatterPlotItem(size=10, pen=pg.mkPen(None), brush=pg.mkBrush(255, 0, 0))
  270. sArtifacts.addPoints(x=pointArtifacts, y=len(pointArtifacts) * [0])
  271. p8.addItem(sArtifacts)
  272. def mouseClickedSpikes(evt):
  273. posClick = evt.pos()
  274. if p8.sceneBoundingRect().contains(posClick):
  275. mousePointC = vb.mapSceneToView(posClick)
  276. hLineSpikes.setPos(mousePointC.y())
  277. threshold = mousePointC.y()
  278. pointSpikes = detectSpikeTimes(threshold)
  279. # print 'spikes clicked:',pointSpikes
  280. sSpikes.setData(x=pointSpikes, y=eDataHP[pointSpikes])
  281. def mouseClickedArtifacts(evt):
  282. posClick = evt.pos()
  283. if p8.sceneBoundingRect().contains(posClick):
  284. mousePointC = vb.mapSceneToView(posClick)
  285. hLineArtifacts.setPos(mousePointC.y())
  286. # print type(evt)
  287. threshold = mousePointC.y()
  288. pointArtifacts = detectSpikeTimes(threshold)
  289. # sArtifacts.clear()
  290. sSpikes.setData(x=spikesRaw[0], y=spikesRaw[1])
  291. sArtifacts.setData(x=pointArtifacts, y=eDataHP[pointArtifacts])
  292. # first graphical dialog to set spike treshold
  293. proxy = pg.SignalProxy(p8.scene().sigMouseMoved, rateLimit=60, slot=mouseMoved)
  294. p8.scene().sigMouseClicked.connect(mouseClickedSpikes)
  295. pdb.set_trace() # input_ = input("Chose spike treshold in graphical window. Press any button to continue.")
  296. spikesRaw = sSpikes.getData()
  297. spikeTreshold = hLineSpikes.getPos()[1] # copy(detectionTreshold)
  298. # second graphical dialog to set treshold for artifacts
  299. p8.scene().sigMouseClicked.disconnect(mouseClickedSpikes)
  300. p8.scene().sigMouseClicked.connect(mouseClickedArtifacts)
  301. pdb.set_trace()
  302. # input_ = input("Chose artifact treshold in graphical window. Press any button to continue.")
  303. falseSpikes = sArtifacts.getData()
  304. artifactTreshold = hLineArtifacts.getPos()[1]
  305. while True:
  306. input_ = input("Enter pairs of indicies of regions to excluce from spike detection (e.g. [[0,700],[5450,5560]]). Press any number if None. : ")
  307. try:
  308. aaa = len(input_)
  309. except:
  310. print('No regions to exclude specified.')
  311. exclusionBorders = None
  312. break
  313. else:
  314. print('recorded')
  315. exclusionBorders = input_
  316. break
  317. while True:
  318. input2_ = input("Enter steps to exclude after stimulus onset - length of stimulus artifact (e.g. 200 corresponding to 2 ms). Press any key if None. : ")
  319. try:
  320. type(input2_)
  321. except:
  322. print('No regions to exclude specified.')
  323. artifactLength = None
  324. break
  325. else:
  326. artifactLength = input2_
  327. break
  328. #
  329. # pdb.set_trace()
  330. lspikes = spikesRaw[0].tolist()
  331. lartif = falseSpikes[0].tolist()
  332. spikes0 = [x for x in lspikes if x not in lartif]
  333. # add spike artifacts to regions to remove
  334. if artifactLength:
  335. if exclusionBorders == None:
  336. exclusionBorders = []
  337. if stim:
  338. for i in range(len(startStim)):
  339. exclusionBorders.append([startStim[i], startStim[i] + artifactLength])
  340. # remove spikes which fall in to regions to exclude
  341. if exclusionBorders:
  342. spikes1 = list(spikes0)
  343. for n in range(len(exclusionBorders)):
  344. spikes1 = [x for x in spikes1 if not (x > exclusionBorders[n][0] and x < exclusionBorders[n][1])]
  345. spikeTimes = eTime.value[np.array(spikes1, dtype=int)]
  346. #firingRate = brian.firing_rate(spikeTimes)
  347. #cv = brian.CV(spikeTimes)
  348. # pdb.set_trace()
  349. ######################################################
  350. # convolv original spike trains with Gaussian kernels
  351. binnedspikes, _ = np.histogram(spikeTimes, tbins)
  352. spikesconv = scipy.ndimage.filters.gaussian_filter1d(np.array(binnedspikes, float), sigma=nspikecountwindow)
  353. # convert the convolved spike trains to units of spikes/sec
  354. spikesconv *= 1. / binWidth
  355. # save data
  356. #self.h5pyTools.createOverwriteDS(grp_spikes, 'spikeTreshold', array([spikeTreshold]))
  357. #self.h5pyTools.createOverwriteDS(grp_spikes, 'artifactTreshold', array([artifactTreshold]))
  358. #self.h5pyTools.createOverwriteDS(grp_spikes, 'spikes', spikeTimes)
  359. #self.h5pyTools.createOverwriteDS(grp_spikes, 'firing_rate_evolution', spikesconv, ['dt', binWidth])
  360. #self.h5pyTools.createOverwriteDS(grp_spikes, 'firing_rate', array([firingRate]))
  361. #self.h5pyTools.createOverwriteDS(grp_spikes, 'CV', array([cv]))
  362. def NormalizeData(data):
  363. return (data - np.min(data)) / (np.max(data) - np.min(data))
  364. #################################################################################
  365. # detect spikes in ephys trace
  366. #################################################################################
  367. def detectSpikeTimes(tresh,eDataHP,ephysTimes,positive=True,plot=False):
  368. #global detectionTreshold
  369. while True:
  370. if positive :
  371. excursion = eDataHP > tresh # threshold ephys trace
  372. else:
  373. excursion = eDataHP < tresh
  374. excursionInt = np.array(excursion, dtype=int) # convert boolean array into array of zeros and ones
  375. diff = excursionInt[1:] - excursionInt[:-1] # calculate difference
  376. spikeStart = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == 1))] # a difference of one is the start of a spike
  377. spikeEnd = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == -1))] # a difference of -1 is the spike end
  378. if len(spikeEnd)>0 and len(spikeStart)>0:
  379. if (spikeEnd[0] - spikeStart[0]) < 0.: # if trace starts below threshold
  380. spikeEnd = spikeEnd[1:]
  381. if (spikeEnd[-1] - spikeStart[-1]) < 0.: # if trace ends below threshold
  382. spikeStart = spikeStart[:-1]
  383. if len(spikeStart) != len(spikeEnd): # unequal lenght of starts and ends is a problem of course
  384. print('problem in length of spikeStart and spikeEnd')
  385. sys.exit(1)
  386. spikeT = []
  387. spikeStart = spikeStart[spikeStart>100]
  388. #for i in range(len(spikeStart)):
  389. # #if (spikeEnd[i] - spikeStart[i]) > 10: # ignore if difference between end and start is smaller than 15 points
  390. # nMin = spikeStart[i] #np.argmin(eDataHP[spikeStart[i]:spikeEnd[i]]) + spikeStart[i]
  391. # spikeT.append(nMin)
  392. # detectionTreshold = tresh
  393. #pdb.set_trace()
  394. spikeTimes = ephysTimes[spikeStart]
  395. if plot:
  396. fig = plt.figure(figsize=(12,8))
  397. ax = fig.add_subplot(111)
  398. ax.plot(ephysTimes,eDataHP)
  399. ax.plot(ephysTimes[spikeStart], eDataHP[spikeStart],'.')
  400. ax.axhline(y=tresh, ls='--', c='0.5')
  401. plt.show()
  402. if plot:
  403. print('Is threshold ok? ->No : type new threshold value ->Yes : press Enter')
  404. recInput = input()
  405. if recInput == "":
  406. break
  407. else:
  408. tresh = float(recInput)
  409. print('new threshold : ', tresh)
  410. else:
  411. break
  412. #recInputIdx = [int(i) for i in recInput.split(',')]
  413. return (spikeTimes,spikeStart,tresh)
  414. #################################################################################
  415. # maps an abritray input array to the entire range of X-bit encoding
  416. #################################################################################
  417. def mapToXbit(inputArray,xBitEncoding):
  418. oldMin = np.min(inputArray)
  419. oldMax = np.max(inputArray)
  420. newMin = 0.
  421. newMax = 2**xBitEncoding-1.
  422. normXBit = newMin + (inputArray - oldMin) * newMax / (oldMax - oldMin)
  423. normXBitInt = np.array(normXBit, dtype=int)
  424. return normXBitInt
  425. #################################################################################
  426. # maps an abritray input array to the entire range of X-bit encoding
  427. #################################################################################
  428. # dataAnalysis.determineFrameTimes(exposureArray[0],arrayTimes,frames)
  429. def determineFrameTimes(exposureArray,arrayTimes,frames,rec=None):
  430. display = False
  431. #pdb.set_trace()
  432. numberOfFrames = len(frames)
  433. exposure = exposureArray > 20 # threshold trace
  434. exposureInt = np.array(exposure, dtype=int) # convert boolean array into array of zeros and ones
  435. difference = np.diff(exposureInt) # calculate difference
  436. expStart = np.arange(len(exposureArray))[np.concatenate((np.array([False]), difference == 1))] # a difference of one is the start of a spike
  437. expEnd = np.arange(len(exposureArray))[np.concatenate((np.array([False]), difference == -1))] # a difference of -1 is the spike end
  438. if (expEnd[0] - expStart[0]) < 0.: # if trace starts above threshold
  439. expEnd = expEnd[1:]
  440. if (expEnd[-1] - expStart[-1]) < 0.: # if trace ends above threshold
  441. expStart = expStart[:-1]
  442. frameDuration = expEnd - expStart
  443. midExposure = (expStart + expEnd)/2
  444. expStartTime = arrayTimes[expStart.astype(int)]
  445. expEndTime = arrayTimes[expEnd.astype(int)]
  446. #framesIdxDuringRec = np.array(len(softFrameTimes))[(arrayTimes[expEnd[0]]+0.002) < softFrameTimes]
  447. #framesIdxDuringRec = framesIdxDuringRec[:len(expStart)]
  448. if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
  449. recordedFrames = frames[3:(len(midExposure) + 3)]
  450. elif arrayTimes[int(midExposure[0])]<0.003:
  451. recordedFrames = frames[2:(len(midExposure) + 2)]
  452. else:
  453. recordedFrames = frames[:len(midExposure)]
  454. print('number of tot. frames, recorded frames, exposures start, end :',numberOfFrames,len(recordedFrames), len(expStart), len(expEnd))
  455. if display:
  456. ledON = np.zeros(len(exposureArray))
  457. for i in range(11):
  458. ledON[((i*1.)<=arrayTimes) & ((i*1.+0.2)>arrayTimes)] = 1.
  459. ledON[29.<=arrayTimes] = 1.
  460. data = np.loadtxt('/home/mgraupe/2019.04.01_000-%s.csv' % (rec[-3:]),delimiter=',',skiprows=1,usecols=(0,1))
  461. print(len(data))
  462. plt.plot(arrayTimes,exposureArray/32.)
  463. plt.plot(arrayTimes,ledON)
  464. print('fist frame at %s sec' % arrayTimes[int(midExposure[0])],end='')
  465. if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
  466. plt.plot(arrayTimes[midExposure.astype(int)], (data[3:(len(midExposure) + 3), 1] - 148.6) / 105.4, 'o-')
  467. print(3)
  468. elif arrayTimes[int(midExposure[0])]<0.003:
  469. plt.plot(arrayTimes[midExposure.astype(int)], (data[2:(len(midExposure) + 2), 1] - 148.6) / 105.4, 'o-')
  470. print(2)
  471. else:
  472. plt.plot(arrayTimes[midExposure.astype(int)], (data[:len(midExposure), 1] - 148.6) / 105.4, 'o-')
  473. print(0)
  474. #plt.plot(softFrameTimes[data[:-6,0].astype(np.int)+6],(data[:-6,1]-148.6)/105.4)
  475. #plt.plot(softFrameTimes,np.ones(len(softFrameTimes)),'|')
  476. plt.show()
  477. pdb.set_trace()
  478. return (expStartTime,expEndTime,recordedFrames)
  479. #################################################################################
  480. def generatePlotWithSTD(data,std=[2,3,4],names = None):
  481. matplotlib.use('TkAgg')
  482. nData = len(data)
  483. fig = plt.figure(figsize=(15,15))
  484. for i in range(nData):
  485. ax0 = fig.add_subplot(nData,1,i+1)
  486. if names is not None:
  487. ax0.set_title('%s' % names[i])
  488. STD = np.std(data[i])
  489. MM = np.mean(data[i])
  490. for n in range(len(std)):
  491. ax0.axhline(y=MM+std[n]*STD,ls='--',c=plt.cm.RdYlBu(n/len(std)),label='%s STD' % std[n])
  492. ax0.axhline(y=MM-std[n]*STD, ls='--', c=plt.cm.RdYlBu(n/len(std)))
  493. ax0.axhline(y=MM,c='C1',label='mean')
  494. ax0.plot(data[i],c='C0')
  495. if i==2:
  496. ax0.plot(np.abs(data[i]), c='C4')
  497. ax0.legend()
  498. plt.show()
  499. #################################################################################
  500. # maps an abritray input array to the entire range of X-bit encoding
  501. #################################################################################
  502. def determineFramesToExclude(frames,probIdx):
  503. listOfFramesToExclude = []
  504. canBeUsed = True
  505. # first let's decide on how many LED's (if any) are present in the FOV
  506. for i in range(len(probIdx)):
  507. currIdx = probIdx[i]
  508. continueDetectLoop = True
  509. while continueDetectLoop:
  510. print('checking idx, started at :', currIdx, probIdx[i])
  511. #frame8bit = np.array(np.transpose(frames[currIdx]), dtype=np.uint8)
  512. img = cv2.cvtColor(frames[currIdx], cv2.COLOR_GRAY2BGR)
  513. # rungs = []
  514. #imgPure = img.copy()
  515. cv2.imshow("PureImage", img)
  516. print('e if to exclude; r to remove from exclude; left right arrows to go back-forward one frame; o to specify another idx; f to move to next; x if recording contains too many errors and cannot be used :')
  517. PressedKey = cv2.waitKey(0)
  518. print(PressedKey)
  519. if PressedKey == 81: # left arrow key
  520. currIdx -=1
  521. elif PressedKey == 83: # right arrow key
  522. currIdx +=1
  523. elif PressedKey == 101: # y key
  524. print('%s added to exclude list' %currIdx)
  525. listOfFramesToExclude.append(currIdx)
  526. elif PressedKey == 114: # e key
  527. print('%s removed from exclude list' % currIdx)
  528. listOfFramesToExclude.remove(currIdx)
  529. elif PressedKey == 111 : # o key
  530. nIdx = input('specify a new idx to check :')
  531. currIdx = int(nIdx)
  532. elif PressedKey == 102: # f key
  533. continueDetectLoop = False
  534. elif PressedKey == 120: # x key
  535. canBeUsed = False
  536. break
  537. else:
  538. print('Key not recognized, try again.')
  539. print('current exclude list :',listOfFramesToExclude)
  540. if not canBeUsed:
  541. break
  542. cv2.destroyWindow("PureImage") # only destroy window at the end of the exploration
  543. lofEx = list(dict.fromkeys(listOfFramesToExclude)) # removes duplicates
  544. lofEx.sort()
  545. print('starting list of indexes :', probIdx)
  546. print('indexes to exclude :', lofEx)
  547. lofEx = np.asarray(lofEx,dtype=int)
  548. #pdb.set_trace()
  549. return (lofEx, canBeUsed)
  550. #################################################################################
  551. # maps an abritray input array to the entire range of X-bit encoding
  552. #################################################################################
  553. # ([ledTraces,ledCoordinates,frames,softFrameTimes,imageMetaInfo],[exposureDAQArray,exposureDAQArrayTimes],[ledDAQControlArray, ledDAQControlArrayTimes],verbose=True)
  554. def determineErroneousFrames(frames):
  555. # first threshold metrics of the movie to detect and exclude erroneous frames with horizontal lines, flash-back frames #########################
  556. frameDiff = []
  557. lineDiff = []
  558. print('calculating frame and line diffs ... ',end='')
  559. for i in range(len(frames)):
  560. if i>0:
  561. frameDiffAllPix = cv2.absdiff(frames[i],frames[i-1])
  562. fD = np.average(frameDiffAllPix)
  563. frameDiff.append(fD)
  564. lineDiffAllLines = cv2.absdiff(frames[i][:,1:],frames[i][:,:-1])
  565. lD = np.average(lineDiffAllLines,axis=0)
  566. lineDiff.append(lD)
  567. #pdb.set_trace()
  568. print('done!')
  569. frameDiff = np.asarray(frameDiff)
  570. lineDiff = np.asarray(lineDiff)
  571. lineDiffSum = np.sum(lineDiff,axis=1)
  572. frameDiffDiff = np.diff(frameDiff)
  573. generatePlotWithSTD([lineDiffSum,frameDiff,frameDiffDiff],std=[3,3.5,4],names=['lineDiffSum','frameDiff','diff of FrameDiff'])
  574. # trick to display the above image
  575. #frame8bit = np.array(np.transpose(frames[0]), dtype=np.uint8)
  576. img = cv2.cvtColor(frames[0], cv2.COLOR_GRAY2BGR)
  577. cv2.imshow('HoldImage',img)
  578. cv2.waitKey(0) #cv2.imshow()
  579. cv2.destroyWindow('HoldImage')
  580. thresholdingInput = input("Specify which trace to use (lineDiffSum - 1, frameDiff - 2, diff of frameDiff - 3; and which multiple of the STD (e.g. 1 3.5); type '4 0' if recording cannote be used (too many errors); '5 0' for recordings without errors : ")
  581. threshold = [float(i) for i in thresholdingInput.split()]
  582. print('choice :', threshold)
  583. #pdb.set_trace()
  584. if threshold[0] == 5.:
  585. idxToExclude = np.array([], dtype=np.int64)
  586. canBeUsed = True
  587. plt.close('all')
  588. return(idxToExclude,canBeUsed)
  589. elif threshold[0] == 1.:
  590. thresholded = lineDiffSum > np.mean(lineDiffSum) + np.std(lineDiffSum)*threshold[1]
  591. outlierIdx = np.arange(len(lineDiffSum))[thresholded] # use indices taking into account missed frames
  592. elif threshold[0] == 2.:
  593. thresholded = frameDiff > np.mean(frameDiff) + np.std(frameDiff) * threshold[1]
  594. outlierIdx = np.arange(len(frameDiff))[thresholded] # use indices taking into account missed frames
  595. outlierIdx += 1 # this is since the difference trace does not start at at the first frame but at the difference between first and second frame
  596. elif threshold[0] == 3.:
  597. thresholded = np.abs(frameDiffDiff) > np.mean(frameDiffDiff) + np.std(frameDiffDiff)*threshold[1]
  598. outlierIdx = np.arange(len(frameDiffDiff))[thresholded] # use indices taking into account missed frames
  599. outlierIdx += 2 # this is since the difference trace does not start at at the first frame but at the difference between first and second frame
  600. elif threshold[0] == 4.:
  601. canBeUsed = False
  602. idxToExclude = np.array([], dtype=np.int64)
  603. plt.close('all')
  604. return (idxToExclude, canBeUsed)
  605. print('length and identity of possible erronous frames :' , len(outlierIdx), outlierIdx)
  606. (idxExclude,canBeUsed) = determineFramesToExclude(frames,outlierIdx)
  607. #excludeMask = np.ones(len(ledVideoRoi[2]),dtype=bool)
  608. # add indicies for equivalent frames
  609. sameFrames = (frameDiff == 0)
  610. sameFrameIdx = np.arange(len(frameDiff))[sameFrames]
  611. sameFrameIdx += 1
  612. print('same frames were recorded here :', sameFrameIdx)
  613. idxToExclude = np.sort(np.concatenate((sameFrameIdx, idxExclude)))
  614. #pdb.set_trace()
  615. #excludeMask[idxToExclude] = False
  616. plt.close('all')
  617. return (idxToExclude,canBeUsed)
  618. #################################################################################
  619. # maps an abritray input array to the entire range of X-bit encoding
  620. #################################################################################
  621. # ([ledTraces,ledCoordinates,frames,softFrameTimes,imageMetaInfo,idxToExclude],[exposureDAQArray,exposureDAQArrayTimes],[ledDAQControlArray, ledDAQControlArrayTimes],verbose=True)
  622. def determineFrameTimesBasedOnLED(ledVideoRoi, cameraExposure, ledDAQc, pc, verbose=False, tail=False,manualThreshold=False):
  623. ##############################################################################################################
  624. # auxiliary function to convert bimodal trace into boolean array
  625. def traceToBinary(trace,threshold=None):
  626. rescaledTrace = (trace - np.min(trace)) / (np.max(trace) - np.min(trace))
  627. if threshold is None:
  628. rescaledTraceBin = rescaledTrace > 0.3
  629. else:
  630. rescaledTraceBin = rescaledTrace > threshold
  631. return (rescaledTrace,rescaledTraceBin)
  632. ##############################################################################################################
  633. def traceToBinaryForChangingMaxMin(trace,threshold=None):
  634. maxTrace = ndimage.maximum_filter(trace, size=5*2)
  635. minTrace = ndimage.minimum_filter(trace, size=5*2)
  636. rescaledTrace = (trace - minTrace) / (maxTrace - minTrace)
  637. if threshold is None:
  638. rescaledTraceBin = rescaledTrace > 0.3
  639. else:
  640. rescaledTraceBin = rescaledTrace > threshold
  641. return (rescaledTrace,rescaledTraceBin)
  642. ##############################################################################################################
  643. # maps LED daq control trace to boolean array ################################################################
  644. # TODO this number is zero on the behavior setup and 4 here
  645. if pc == 'behaviorPC':
  646. LEDcontrolIdx = 0 # which trace of the DAQ recording is linked to the !!!
  647. elif pc == '2photonPC':
  648. LEDcontrolIdx = 4
  649. else:
  650. print('Make sure the computer of the recording is specified.')
  651. ledDAQcontrolBin = traceToBinary(ledDAQc[0][LEDcontrolIdx])[1] # here the threshold is not important as the trace is binary to start out with
  652. # convert LED roi traces from video to boolean arrays
  653. ledVideoRoiBins = []
  654. ledVideoRoiRescaled = []
  655. allLEDVideoRoiValues = []
  656. # determine threshold [ledTraces,ledCoordinates,frames,softFrameTimes,imageMetaInfo,idxToExclude]
  657. # tail covering the LEDs for some
  658. if tail:
  659. matplotlib.use('TkAgg')
  660. print(' in tail ...')
  661. anticipateCorrectValues = True
  662. for i in range(ledVideoRoi[1][0]):
  663. plt.plot(ledVideoRoi[0][i],'o-',ms=2,label='%s' % i)
  664. plt.legend(loc=1)
  665. plt.show()
  666. if anticipateCorrectValues:
  667. inputA = input('Index until which the recording is not affected by the tail (integer; type 0 if recording is ok) :')
  668. #inputA=350
  669. untilOKidx = int(inputA)
  670. if untilOKidx != 0:
  671. period = [7, 7, 7, 5]
  672. for i in range(ledVideoRoi[1][0]):
  673. maxVal = np.max(ledVideoRoi[0][i][20:untilOKidx])
  674. minVal = np.min(ledVideoRoi[0][i][20:untilOKidx])
  675. for n in range(period[i]):
  676. # repeatValue(ledVideoRoi[0][i][(untilOKidx+n):],7)
  677. isHigh = [True if abs(ledVideoRoi[0][i][(untilOKidx + n)] - maxVal) < abs(ledVideoRoi[0][i][(untilOKidx + n)] - minVal) else False]
  678. if isHigh:
  679. ledVideoRoi[0][i][(untilOKidx + n):][::period[i]] = ledVideoRoi[0][i][(untilOKidx + n)]
  680. else:
  681. ledVideoRoi[0][i][(untilOKidx + n):][::period[i]] = ledVideoRoi[0][i][(untilOKidx + n)] # ledVideoRoi[0][i][]
  682. else:
  683. maxV = [254,251,250,213]
  684. minV = [200,174,217,147]
  685. idxMaxV = [[8580,8582,8583,8585],
  686. [],
  687. [8589],
  688. []]
  689. idxMinV = [[8581,8584,8586,8588],
  690. [8579,8580],
  691. [8590,8591],
  692. [8584,8585,8586,8589,8590,8591]]
  693. for i in range(4):
  694. for n in idxMaxV[i]:
  695. ledVideoRoi[0][i][n] = maxV[i]
  696. for m in idxMinV[i]:
  697. ledVideoRoi[0][i][m] = minV[i]
  698. # fig = plt.figure()
  699. # for i in range(ledVideoRoi[1][0]):
  700. # #ax = fig.add_subplot(3,1,i)
  701. # plt.plot(ledVideoRoi[0][i],'o-',ms=2,label='%s' % i)
  702. # #ax.set_xlim(8)
  703. # plt.legend(loc=1)
  704. # plt.show()
  705. # pdb.set_trace()
  706. ###########
  707. for i in range(ledVideoRoi[1][0]):
  708. allLEDVideoRoiValues.extend(traceToBinaryForChangingMaxMin(ledVideoRoi[0][i])[0]) # rescale all values to [0,1] and stack them
  709. allLEDVideoRoiValues = np.sort(np.asarray(allLEDVideoRoiValues)) # convert to array and sort
  710. luminocityDifferences = np.diff(allLEDVideoRoiValues)
  711. idxMaxDiff = np.argmax(luminocityDifferences)
  712. LEDVideoThreshold = allLEDVideoRoiValues[idxMaxDiff] + (allLEDVideoRoiValues[idxMaxDiff+1] - allLEDVideoRoiValues[idxMaxDiff])/2.
  713. if pc == 'behaviorPC':
  714. #illumLEDcontrolThreshold = LEDVideoThreshold**4.49185827 # mapping, i.e. exponent, from tools/fitOfIlluminationValues
  715. illumLEDcontrolThreshold = LEDVideoThreshold**18.37008924
  716. elif pc == '2photonPC':
  717. illumLEDcontrolThreshold = LEDVideoThreshold**2.61290794 # 2pinvivo
  718. # if (illumLEDcontrolThreshold) < 0.08 or manualThreshold:
  719. # print('thresholds before: ',LEDVideoThreshold, illumLEDcontrolThreshold)
  720. # #print('LED threshold extremly low! Fixed by setting both threshold to 0.8.')
  721. # fig = plt.figure(figsize=(10,10))
  722. # ax = fig.add_subplot(111)
  723. # #ax.plot(np.ones(len(allLEDVideoRoiValues)),allLEDVideoRoiValues,'.',ms=0.5)
  724. # ax.axhline(y=LEDVideoThreshold,ls='--',c='C0')
  725. # ax.plot(allLEDVideoRoiValues,'.',ms=0.5,c='C0')
  726. # #plt.plot(np.ones(len(ledDAQcontrolBin)),ledDAQcontrolBin,'.')
  727. # #ax.plot(ledDAQcontrolBin,'.',ms=0.5)
  728. # plt.show()
  729. # # thresholdInput = '0.8,0.8'
  730. # # thresholdInput = ''
  731. # thresholdInput = input('Provide alternative thresholds (e.g. 0.8,0.7), otherwise press enter to keep current thresholds : ')
  732. # if not thresholdInput == '':
  733. # newThresholds = [float(i) for i in thresholdInput.split(',')]
  734. # LEDVideoThreshold = newThresholds[0]
  735. # illumLEDcontrolThreshold = newThresholds[1]
  736. # find start and end of camera exposure period ################################################################
  737. exposureInt = np.array(cameraExposure[0][0], dtype=int) # convert boolean array into array of zeros and ones
  738. difference = np.diff(exposureInt) # calculate difference
  739. expStart = np.arange(len(exposureInt))[np.concatenate((np.array([False]), difference > 0))] # a difference of one is the start of the exposure
  740. expEnd = np.arange(len(exposureInt))[np.concatenate((np.array([False]), difference < 0))] # a difference of -1 is the end the exposure period
  741. #pdb.set_trace()
  742. if (expEnd[0] - expStart[0]) < 0.: # if trace starts above threshold
  743. print('exposure at start of recording')
  744. expEnd = expEnd[1:]
  745. exposureAtStart = True
  746. else:
  747. exposureAtStart = False
  748. if (expEnd[-1] - expStart[-1]) < 0.: # if trace ends above threshold
  749. print('exposure during end of recording')
  750. exposureAtEnd = True
  751. expStart = expStart[:-1]
  752. else:
  753. exposureAtEnd = False
  754. expStart = expStart.astype(int)
  755. expEnd = expEnd.astype(int)
  756. expStartTime = cameraExposure[1][expStart] # everything was based on indicies up to this point : here indicies -> time
  757. expEndTime = cameraExposure[1][expEnd] # everything was based on indicies up to this point : here indicies -> time
  758. frameDuration = expEndTime - expStartTime
  759. print('first frame started at ', expStartTime[0]*1000., 'ms' )
  760. ## based on exposure start-stop, how bright should the DAQ LED signal be ##########################################
  761. startEndExposureTime = np.column_stack((expStartTime, expEndTime))
  762. startEndExposurepIdx = np.column_stack((expStart,expEnd)) # create a 2-column array with 1st column containing start and 2nd column containing end index
  763. illumLEDcontrol = [np.mean(ledDAQcontrolBin[b[0]:b[1]]) for b in startEndExposurepIdx] # extract MEAN illumination value - from LED control trace - during exposure period
  764. illumLEDcontrol = np.asarray(illumLEDcontrol)
  765. adjusted = False
  766. if manualThreshold:
  767. while True:
  768. if (illumLEDcontrolThreshold) < 0.08 or manualThreshold:
  769. sortedIllumLEDcontrol = np.sort(illumLEDcontrol)
  770. fig = plt.figure(figsize=(10,10))
  771. ax = fig.add_subplot(111)
  772. #ax.plot(np.ones(len(allLEDVideoRoiValues)),allLEDVideoRoiValues,'.',ms=0.5)
  773. print('thresholds before: ',LEDVideoThreshold, illumLEDcontrolThreshold)
  774. if (illumLEDcontrolThreshold < 0.08) and not (adjusted):
  775. illumLEDcontrolThreshold = 0.2
  776. adjusted = True # do it only once
  777. print('adjusted thresholds: ', LEDVideoThreshold, illumLEDcontrolThreshold)
  778. print('video (up, down, down fraction) : ', np.sum(sortedIllumLEDcontrol > LEDVideoThreshold), np.sum(sortedIllumLEDcontrol < LEDVideoThreshold), np.sum(sortedIllumLEDcontrol < LEDVideoThreshold)/len(sortedIllumLEDcontrol))
  779. print('illum (up, down, down fraction): ', np.sum(allLEDVideoRoiValues>illumLEDcontrolThreshold), np.sum(allLEDVideoRoiValues<illumLEDcontrolThreshold), np.sum(allLEDVideoRoiValues<illumLEDcontrolThreshold)/len(allLEDVideoRoiValues))
  780. ax.axhline(y=illumLEDcontrolThreshold,ls='--',c='C0')
  781. ax.plot(np.linspace(0,1,len(sortedIllumLEDcontrol)),sortedIllumLEDcontrol,'.',ms=0.5,c='C0')
  782. ax.axhline(y=LEDVideoThreshold, ls='--', c='C1')
  783. ax.plot(np.linspace(0,1,len(allLEDVideoRoiValues)),allLEDVideoRoiValues, '.', ms=0.5,c='C1')
  784. #plt.plot(np.ones(len(ledDAQcontrolBin)),ledDAQcontrolBin,'.')
  785. #ax.plot(ledDAQcontrolBin,'.',ms=0.5)
  786. plt.show()
  787. # thresholdInput = '0.8,0.8'
  788. # thresholdInput = ''
  789. thresholdInput = input('Provide alternative thresholds (e.g. 0.8,0.2), otherwise press enter exit loop : ')
  790. if not thresholdInput == '':
  791. newThresholds = [float(i) for i in thresholdInput.split(',')]
  792. LEDVideoThreshold = newThresholds[0]
  793. illumLEDcontrolThreshold = newThresholds[1]
  794. else:
  795. break
  796. else:
  797. LEDVideoThreshold = LEDVideoThreshold
  798. illumLEDcontrolThreshold = illumLEDcontrolThreshold
  799. # print('press any key to check/redefine thresholds; exit loop with space or enter:')
  800. # PressedKey = cv2.waitKey(0)
  801. # if PressedKey == 13 or PressedKey == 32: # Enter or Space
  802. # break
  803. # else:
  804. # pass
  805. print('thresholds final: ',LEDVideoThreshold, illumLEDcontrolThreshold)
  806. # pdb.set_trace()
  807. # LEDVideoThreshold = 0.8
  808. # illumLEDcontrolThreshold = 0.8
  809. #print('adjusted thresholds : ', LEDVideoThreshold, illumLEDcontrolThreshold)
  810. (illumLEDcontrolrescaled, illumLEDcontrolBin) = traceToBinary(illumLEDcontrol, threshold=illumLEDcontrolThreshold) # 0.2 and 0.15 before
  811. # pdb.set_trace()
  812. # threshold and convert to binary
  813. for i in range(ledVideoRoi[1][0]):
  814. ledVideoRoiBins.append(traceToBinary(ledVideoRoi[0][i],threshold=LEDVideoThreshold)[1]) # 0.6 before 0.4
  815. ledVideoRoiRescaled.append(traceToBinary(ledVideoRoi[0][i])[0])
  816. #plt.plot(allLEDVideoRoiValues,illumLEDcontrol, '.', ms=0.5)
  817. #plt.show()
  818. #plt.plot(illumLEDcontrol)
  819. ## loop over frame numbers and extract binary number shown by leds ################################################
  820. nFrames = len(ledVideoRoiBins[0])
  821. binNumbers = np.array([[False,False,False],[True,False,False],[False,True,False],[True,True,False],[False,False,True],[True,False,True],[False,True,True],[True,True,True]])
  822. recordedFrames = 0
  823. frameCount = []
  824. binNumberInFrame = np.column_stack((ledVideoRoiBins[0],ledVideoRoiBins[1],ledVideoRoiBins[2]))
  825. frameNBefore = 0
  826. oldI = -1
  827. exceptionsInFrameCount = []
  828. idxToExclude = ledVideoRoi[5]
  829. for i in range(nFrames):
  830. if i not in idxToExclude:
  831. matchBool = np.all(np.equal(binNumberInFrame[i],binNumbers),axis=1) # which of the boolean number corresponds to the current frame pattern : return is a boolean list from 0 to 8 with one TRUE entry
  832. matchFrameN = np.arange(len(binNumbers))[matchBool][0] # converts the boolean list into the index corresponding to the match
  833. frameDiff = matchFrameN - frameNBefore # difference in count to previous frame
  834. if frameDiff < 0: # else : negative difference indicates that the counter restarted
  835. frameDiff+=7
  836. if (frameDiff != 1) and (frameDiff != -6):
  837. print(i,oldI,i-oldI,matchFrameN,frameNBefore,frameDiff,binNumberInFrame[i],binNumberInFrame[i-1])
  838. exceptionsInFrameCount.append([i,oldI,i-oldI,matchFrameN,frameNBefore,frameDiff,binNumberInFrame[i],binNumberInFrame[i-1]])
  839. if matchFrameN == 0: # counter will start at 0 and possibly go back to zero after end of recording
  840. if (i>70) and (i<(nFrames-10)):
  841. print(i,oldI,matchFrameN,frameNBefore,frameDiff,binNumberInFrame[i],binNumberInFrame[i-1])
  842. print('strange, zero frame in the middle of recording')
  843. pdb.set_trace()
  844. #frameDiff = 0
  845. frameCount.append([i,matchFrameN,frameDiff,int(ledVideoRoiBins[3][i]),oldI])
  846. frameNBefore = matchFrameN
  847. oldI = i
  848. frameCount = np.asarray(frameCount,dtype=int) # convert list to integer array
  849. idxRecordedFrames = np.cumsum(frameCount[:,2]) # use the frame differences to generate new index corresponding to video recording
  850. idxCounting = np.argwhere(idxRecordedFrames>0) # start and end index with first and last frame recording the counter
  851. idxFramesDuringRecording = idxRecordedFrames[idxCounting[0,0]:(idxCounting[-1,0]+1)] - 1 # remove leading and trailing zeros, and remove one to have the new index start with zero, cumsum makes first index to be 1
  852. #pdb.set_trace()
  853. if exposureAtStart: # remove first frame if exposure was active during start of recording, i.e., at t = 0 s
  854. idxFramesDuringRecording = idxFramesDuringRecording[1:] - 1
  855. if exposureAtEnd:
  856. idxFramesDuringRecording = idxFramesDuringRecording[:-1]
  857. #pdb.set_trace()
  858. idxMissingFrames = np.delete(np.arange(idxFramesDuringRecording[-1]+1),idxFramesDuringRecording)
  859. #idxTestMask = idxFramesDuringRecording < len(illumLEDcontrolBin) # index should not exceed length of array
  860. #illum = illumLEDcontrolBin[idxFramesDuringRecording[idxTestMask]]
  861. ## the excluded frames - based on distortions - need to be removed from the video sequence
  862. videoRoi = ledVideoRoiBins[3]
  863. mask = np.ones(len(videoRoi),dtype=bool)
  864. mask[idxToExclude] = False
  865. #pdb.set_trace()
  866. ## first loop to align the START of the frame recording - in illumLEDcontrolBin - with the video recording
  867. if any(idxToExclude < 20):
  868. print('Early frames to exclude. Problem!')
  869. pdb.set_trace()
  870. else:
  871. tmpIdx = np.where(videoRoi==True) # Index of the first ON frame for the 4th LED
  872. idxFirstFrameRec = tmpIdx[0][0] # extract index of first frame during recording
  873. if exposureAtStart:
  874. idxFirstFrameRec+=1 # increase that
  875. for j in range(70):
  876. videoRoiWOEX = videoRoi[mask][j:]
  877. if np.all(videoRoiWOEX[:20] == illumLEDcontrolBin[:20]): # note that illumLEDcontrolBin already accounts for a recording during stat of rec, this frame is removed
  878. missedFramesBegin = j
  879. break
  880. try:
  881. a=missedFramesBegin
  882. #a = lllll
  883. except:
  884. print("bad alignement !!!!!!!!!!!!!!!!!!!")
  885. print(videoRoi[mask][:20],illumLEDcontrolBin[:20])
  886. plt.plot(illumLEDcontrolrescaled,'.',ms=0.5)
  887. plt.plot(ledVideoRoiRescaled[3],'.',ms=0.5)
  888. plt.show()
  889. pdb.set_trace()
  890. if (idxFirstFrameRec == missedFramesBegin) or (missedFramesBegin == 0):
  891. print('Number of frames recorded before first full exposed frame during recording :', missedFramesBegin, idxFirstFrameRec, illumLEDcontrolBin[:20],videoRoi[:20] )
  892. videoRoiWOEX = videoRoi[mask][missedFramesBegin:]
  893. elif idxFirstFrameRec == (missedFramesBegin+1):
  894. missedFramesBegin+=1
  895. print('Number of frames recorded before first full exposed frame during recording (increased by one):', missedFramesBegin, idxFirstFrameRec, illumLEDcontrolBin[:20],videoRoi[:20])
  896. videoRoiWOEX = videoRoi[mask][missedFramesBegin:]
  897. else:
  898. print('Number of frames recorded before first full exposed frame during recording :', missedFramesBegin, idxFirstFrameRec, illumLEDcontrolBin[:20],videoRoi[:20] )
  899. print('Problem with determining index of first recorded frame.')
  900. pdb.set_trace()
  901. #pdb.set_trace()
  902. ## second loop in order to align the
  903. shiftDifference = []
  904. lengthOfIllumLEDcontrol = len(illumLEDcontrolBin)
  905. lengthOfROIinVideo = len(videoRoiWOEX)
  906. lengthOfIdxCount = idxFramesDuringRecording[-1] + 1
  907. print('length of illumLEDcontrolBin and videoRoiWOEX and idxFramesDuringRecording[-1] : ', lengthOfIllumLEDcontrol, lengthOfROIinVideo, lengthOfIdxCount)
  908. for i in range(-10,11,1): # loop to shift the mask over
  909. idxTemp = idxFramesDuringRecording + i # shift the array by increasing the indicies by a certain number
  910. idxIllum = idxTemp[(idxTemp>=0)&(idxTemp<lengthOfIllumLEDcontrol)] # indicies have to be larger than zero and should not be larger than the length of the illumLEDcontrolBin array
  911. #idxIllum = idx[idx<lengthOfIllumLEDcontrol] # indicies should not be larger than the length of the illumLEDcontrolBin array
  912. illum = illumLEDcontrolBin[idxIllum] # illumination at these indicies
  913. # idxMissing = np.delete(np.arange(idxIllum[-1]), idxIllum) #[i:]
  914. # idxMissing = np.delete(np.arange(idxFramesDuringRecording[-1]), idxIllum) # [i:]
  915. idxMissing = np.delete(np.arange(lengthOfIllumLEDcontrol), idxIllum)
  916. NidxRemovedAtExtremities = idxIllum[0] + ((lengthOfIllumLEDcontrol-1) - idxIllum[-1]) # counts number of frames missing in the beginning and end
  917. NidxRemovedAtExtremities -= np.sum((idxMissingFrames<idxIllum[0]) | (idxMissingFrames>idxIllum[-1])) # reduce if missing frames are in the extrimities
  918. #if i<0:
  919. # frameOverlap = [0 if ((lengthOfIllumLEDcontrol+np.abs(i)+1)<(lengthOfROIinVideo+len(idxMissingFrames))) else ((lengthOfROIinVideo+len(idxMissingFrames)) - (lengthOfIllumLEDcontrol+np.abs(i)+1))]
  920. #elif i>=0:
  921. # frameOverlap = [i if (lengthOfIllumLEDcontrol<(lengthOfROIinVideo+len(idxMissingFrames))) else (lengthOfIllumLEDcontrol-(lengthOfROIinVideo+len(idxMissingFrames)+i+1))]
  922. #print('overlap :',frameOverlap[0])
  923. #print('test',len(np.intersect1d(idxMissing,idxMissingFrames)), idxMissingFrames, idxMissing,(-len(idxMissing)),NidxRemovedAtExtremities)
  924. compareIdx = len(np.intersect1d(idxMissing,idxMissingFrames)) - len(idxMissing) + np.abs(NidxRemovedAtExtremities) # abs(i)
  925. len0 = len(illum)
  926. len1 = lengthOfROIinVideo
  927. if len0 < len1:
  928. compare = np.sum(np.equal(illum,videoRoiWOEX[:len0]))
  929. versch = compare - len0
  930. totLength = len0
  931. largeLength = len1
  932. else:
  933. compare = np.sum(np.equal(illum[:len1],videoRoiWOEX))
  934. versch = compare - len1
  935. totLength = len1
  936. largeLength = len0
  937. #pdb.set_trace()
  938. shiftDifference.append([i, versch, totLength, compareIdx,NidxRemovedAtExtremities])
  939. print(i, versch, totLength, compareIdx, NidxRemovedAtExtremities, idxMissing, idxMissingFrames)
  940. #if i >=0 :
  941. #pdb.set_trace()
  942. #compare = np.equal()
  943. shiftDifference = np.asarray(shiftDifference)
  944. if len(idxToExclude)==0: # without erronous frames ...
  945. shiftDifference = shiftDifference[shiftDifference[:,0]>=0] # ... use only the shifts which are larger than zero
  946. shiftToZero = shiftDifference[:,0][(shiftDifference[:,1]==0) & (shiftDifference[:,3]==0)]
  947. finalLength = shiftDifference[:,2][(shiftDifference[:,1]==0) & (shiftDifference[:,3]==0)]
  948. if len(shiftToZero)>1 or len(shiftToZero)==0:
  949. if len(shiftToZero)==0:
  950. idxTemp = idxFramesDuringRecording + 0
  951. idx = idxTemp[idxTemp >= 0]
  952. idxIllum = idx[idx < len(illumLEDcontrolBin)]
  953. # pdb.set_trace()
  954. shortest = [len(ledVideoRoiRescaled[3][mask][missedFramesBegin:]) if len(ledVideoRoiRescaled[3][mask][missedFramesBegin:]) < len(illumLEDcontrolrescaled[idxIllum]) else len(
  955. illumLEDcontrolrescaled[idxIllum])]
  956. plt.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]], illumLEDcontrolrescaled[idxIllum][:shortest[0]], 'o', ms=1)
  957. plt.show()
  958. pdb.set_trace()
  959. fig = plt.figure(figsize=(20,10))
  960. ax = fig.add_subplot(111)
  961. print('difference at :',)
  962. ax.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]], 'o-',ms=0.5,lw=0.3)
  963. ax.plot(illumLEDcontrolrescaled[idxIllum][:shortest[0]], 'o-',ms=0.5,lw=0.3)
  964. plt.show()
  965. #plt.clf()
  966. fig = plt.figure(figsize=(20, 10))
  967. ax = fig.add_subplot(111)
  968. ax.axhline(y=LEDVideoThreshold,c='C0',ls='--')
  969. ax.axhline(y=illumLEDcontrolThreshold,c='C1',ls='--')
  970. ledVideo = ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]]
  971. ledIllumDAQ = illumLEDcontrolrescaled[idxIllum][:shortest[0]]
  972. ax.plot(np.arange(len(ledVideo))[ledVideo > LEDVideoThreshold],ledVideo[ledVideo > LEDVideoThreshold], 'v', c='C0', ms=2)
  973. ax.plot(np.arange(len(ledVideo))[ledVideo < LEDVideoThreshold],ledVideo[ledVideo < LEDVideoThreshold], 'o', c='C0', ms=2)
  974. ax.plot(np.arange(len(ledIllumDAQ))[ledIllumDAQ > illumLEDcontrolThreshold],ledIllumDAQ[ledIllumDAQ>illumLEDcontrolThreshold], 'v',c='C1', ms=2)
  975. ax.plot(np.arange(len(ledIllumDAQ))[ledIllumDAQ < illumLEDcontrolThreshold],ledIllumDAQ[ledIllumDAQ<illumLEDcontrolThreshold], 'o', c='C1', ms=2)
  976. plt.show()
  977. pdb.set_trace()
  978. elif shiftToZero[1] == (shiftToZero[0]+5):
  979. print('Multiple shifts to zero, so multiple perfect overlays exist. First overlay with shift %s will be used.' % shiftToZero[0])
  980. pass
  981. else:
  982. print('Problem! More than one shift led to perfect overlay!')
  983. #np.arange(np.diff(idxRecordedFrames)>1)
  984. #
  985. idxTemp = idxFramesDuringRecording + 0
  986. idx = idxTemp[idxTemp>=0]
  987. idxIllum = idx[idx<len(illumLEDcontrolBin)]
  988. #pdb.set_trace()
  989. shortest = [len(ledVideoRoiRescaled[3][mask][missedFramesBegin:]) if len(ledVideoRoiRescaled[3][mask][missedFramesBegin:])<len(illumLEDcontrolrescaled[idxIllum]) else len(illumLEDcontrolrescaled[idxIllum])]
  990. plt.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]],illumLEDcontrolrescaled[idxIllum][:shortest[0]],'o',ms=1)
  991. plt.show()
  992. pdb.set_trace()
  993. plt.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]],'o-')
  994. plt.plot(illumLEDcontrolrescaled[idxIllum][:shortest[0]],'o-')
  995. plt.show()
  996. pdb.set_trace()
  997. finalShiftToZero = shiftToZero[0]
  998. print('final shift to zero and final length : ', finalShiftToZero,finalLength[0])
  999. idxTemp = idxFramesDuringRecording + finalShiftToZero
  1000. idx = idxTemp[idxTemp>=0]
  1001. idxIllumFinal = idx[idx<len(illumLEDcontrolBin)][:finalLength[0]]
  1002. compareIllumination = False
  1003. if compareIllumination:
  1004. illum = illumLEDcontrolrescaled[idxIllumFinal]
  1005. videoROI = ledVideoRoiRescaled[3][mask][missedFramesBegin:]
  1006. shortest = [len(videoROI) if len(videoROI) < len(illum) else len(illum)]
  1007. bothCombined = np.column_stack((videoROI[:shortest[0]],illum[:shortest[0]]))
  1008. plt.plot(videoROI[:shortest[0]], illum[:shortest[0]], 'o')
  1009. plt.show()
  1010. pdb.set_trace()
  1011. try :
  1012. illumValues = pickle.load( open('illuminatoinValues.p', 'rb' ) )
  1013. except :
  1014. illumValues = bothCombined
  1015. else:
  1016. illumValues = np.row_stack((illumValues,bothCombined))
  1017. pickle.dump(illumValues, open('illuminatoinValues.p', 'wb'))
  1018. frameTimes = startEndExposureTime[idxIllumFinal]
  1019. frameStartStopIdx = startEndExposurepIdx[idxIllumFinal]
  1020. videoIdx = np.arange(len(ledVideoRoiBins[3]))[mask][missedFramesBegin:][:finalLength[0]]
  1021. #recFrames = ledVideoRoi[2][videoIdx]
  1022. ddd = np.diff(idxIllumFinal)
  1023. print('Total number of dropped and excluded frames : ', np.sum(ddd-1), 'out of',len(ledVideoRoi[2]),'frame in total.')
  1024. print('Excluded frames :', len(idxToExclude))
  1025. print('Dropped frames :', np.sum(ddd-1)-len(idxToExclude))
  1026. frameSummary = np.array([len(ledVideoRoi[2]),np.sum(ddd-1),len(idxToExclude), np.sum(ddd-1)-len(idxToExclude)])
  1027. #pdb.set_trace()
  1028. return (idxIllumFinal,frameTimes,frameStartStopIdx,videoIdx,frameSummary)
  1029. ##############################################################################################################################
  1030. #pdb.set_trace()
  1031. for i in range(10):
  1032. #print(i)
  1033. #idxTest = idxRecordedFramesCleaned[1:-3] - 1
  1034. #compare = ledVideoRoiBins[3][2:-3] == illumLEDcontrolBin[idxTest]
  1035. #shortestLength = [len(illum) if (0<(len(ledVideoRoiBins)-(len(illum)+i))) else ]
  1036. videoRoi = ledVideoRoiBins[3][i]
  1037. if len(videoRoi) > len(illum):
  1038. #shortestLength = len(illum)
  1039. #else:
  1040. #shortestLength = len(videoRoi)
  1041. print('problem in length relations')
  1042. pdb.set_trace()
  1043. compare = illum[:len(videoRoi)] == videoRoi
  1044. differences = np.sum(np.invert(compare))
  1045. print('number of differences :', i, differences,i,len(illum)+i,len(videoRoi))
  1046. shifting.append([i,differences,i,len(illum)+i,len(videoRoi)])
  1047. shifting = np.asarray(shifting)
  1048. correctShift = np.argwhere(shifting[:,1]==0)
  1049. if len(correctShift) == 0:
  1050. print('No perfect overlay has been found')
  1051. print(shifting)
  1052. pdb.set_trace()
  1053. elif len(correctShift)>1:
  1054. print('Multiple corret overlays have been found. Suspicious!')
  1055. print(shifting)
  1056. pdb.set_trace()
  1057. elif len(correctShift) == 1:
  1058. rightShift = shifting[correctShift[0][0]]
  1059. print('The correct shift is ', rightShift)
  1060. print('Number of recorded videos :', len(ledVideoRoi[2][rightShift[2]:rightShift[3]]))
  1061. print('Number of associated time points :', len(startEndExposurepIdx[idxRecordedFramesCleaned[idxTestMask]][:rightShift[4]]))
  1062. ddd = np.diff(idxRecordedFramesCleaned[idxTestMask][:rightShift[4]])
  1063. print('Number of gaps, number of lost frames, and size of gaps :', len(ddd[ddd>1]),np.sum(ddd[ddd>1]) - len(ddd[ddd>1]), ddd[ddd>1])
  1064. idxVideo = np.arange(rightShift[2],rightShift[3])
  1065. idxTimePoints = idxRecordedFramesCleaned[idxTestMask][:rightShift[4]]
  1066. #pdb.set_trace()
  1067. return (idxVideo,idxTimePoints,startEndExposureTime,startEndExposurepIdx,rightShift)
  1068. if compare == 0:
  1069. #plt.plot(ledVideoRoi[0][3][2:-3], 'o-', label='ledVideoRoi')
  1070. ii = 3
  1071. plt.plot(ledVideoRoiRescaled[3][ii:len(illum)+ii], 'o-', label='ledVideoRoiRescaled')
  1072. plt.plot(ledVideoRoiBins[3][ii:len(illum)+ii],'o-',label='ledVideoRoiBins')
  1073. plt.plot(illumLEDcontrol[idxRecordedFramesCleaned[idxTestMask]],'o-',label='illumLEDcontrol')
  1074. plt.plot(illumLEDcontrolBin[idxRecordedFramesCleaned[idxTestMask]], 'o-', label='illumLEDcontrolBin')
  1075. plt.legend()
  1076. plt.show()
  1077. pdb.set_trace()
  1078. #compare = illumLEDcontrol[idxTest] ==
  1079. #idxTest = idxRecordedFramesCleaned[1:-1]-1
  1080. totLength = len(illumLEDcontrol[idxTest])
  1081. ret = np.array_equal(illumLEDcontrol[idxRecordedFramesCleaned],ledVideoRoiBins[3][i:(totLength+i)])
  1082. print(i,ret)
  1083. pdb.set_trace()
  1084. pdb.set_trace()
  1085. if len(illumLEDcontrol) <= len(ledVIDEOroi):
  1086. illuminationLonger = True
  1087. ledVIDEOroiMask = np.arange(len(ledVIDEOroi)) < len(illumination)
  1088. illuminationMask = np.arange(len(illumination)) < len(illumination)
  1089. cc = crosscorr(1,illumination,ledVIDEOroi[ledVIDEOroiMask],20) # calculate cross-correlation between LED in video and LED from DAQ array
  1090. else:
  1091. illumniationLonger = False
  1092. ledVIDEOroiMask = np.arange(len(ledVIDEOroi)) < len(ledVIDEOroi)
  1093. illuminationMask = np.arange(len(illumination)) < len(ledVIDEOroi)
  1094. cc = crosscorr(1, illumination[illuminationMask], ledVIDEOroi, 3) # calculate cross-correlation between LED in video and LED from DAQ array
  1095. peaks = find_peaks(cc[:,1],height=0)
  1096. if len(peaks[0]) > 1:
  1097. print('MULTIPLE peaks found in cross-correlogram between LED brigthness and DAQ array')
  1098. pdb.set_trace()
  1099. elif len(peaks[0]) == 0:
  1100. print('NO peaks were found in cross-correlogram between LED brigthness and DAQ array')
  1101. pdb.set_trace()
  1102. else:
  1103. pdb.set_trace()
  1104. shift = cc[:,0][peaks[0][0]]
  1105. shiftInt = int(shift)
  1106. print('video trace has to be shifted by (float and int number) ', shift, shiftInt)
  1107. #print(len(ledVIDEOroi),len(illumination))
  1108. #pdb.set_trace()
  1109. if verbose:
  1110. if shiftInt >= 0:
  1111. plt.plot(ledVIDEOroi[ledVIDEOroiMask][shiftInt:],'o-',ms=1,label='Video roi (shifted)')
  1112. else:
  1113. plt.plot(ledVIDEOroi[ledVIDEOroiMask][:shiftInt], 'o-', ms=1, label='Video roi (shifted)')
  1114. plt.plot(illumination[illuminationMask],'o-',ms=1,label='from LED daq control')
  1115. plt.legend()
  1116. plt.show()
  1117. frameIdx = np.arange(len(ledVIDEOroi))
  1118. recordedFramesIdx = frameIdx[shiftInt:(len(illumination)+shiftInt)]
  1119. #pdb.set_trace()
  1120. return (startEndExpTime,startEndExpIdx,recordedFramesIdx)
  1121. ##### end of current implementation ##############################################################################################
  1122. #framesIdxDuringRec = np.array(len(softFrameTimes))[(arrayTimes[expEnd[0]]+0.002) < softFrameTimes]
  1123. #framesIdxDuringRec = framesIdxDuringRec[:len(expStart)]
  1124. if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
  1125. recordedFrames = frames[3:(len(midExposure) + 3)]
  1126. elif arrayTimes[int(midExposure[0])]<0.003:
  1127. recordedFrames = frames[2:(len(midExposure) + 2)]
  1128. else:
  1129. recordedFrames = frames[:len(midExposure)]
  1130. print('number of tot. frames, recorded frames, exposures start, end :',numberOfFrames,len(recordedFrames), len(expStart), len(expEnd))
  1131. if display:
  1132. ledON = np.zeros(len(exposureArray))
  1133. for i in range(11):
  1134. ledON[((i*1.)<=arrayTimes) & ((i*1.+0.2)>arrayTimes)] = 1.
  1135. ledON[29.<=arrayTimes] = 1.
  1136. data = np.loadtxt('/home/mgraupe/2019.04.01_000-%s.csv' % (rec[-3:]),delimiter=',',skiprows=1,usecols=(0,1))
  1137. print(len(data))
  1138. plt.plot(arrayTimes,exposureArray/32.)
  1139. plt.plot(arrayTimes,ledON)
  1140. print('fist frame at %s sec' % arrayTimes[int(midExposure[0])],end='')
  1141. if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
  1142. plt.plot(arrayTimes[midExposure.astype(int)], (data[3:(len(midExposure) + 3), 1] - 148.6) / 105.4, 'o-')
  1143. print(3)
  1144. elif arrayTimes[int(midExposure[0])]<0.003:
  1145. plt.plot(arrayTimes[midExposure.astype(int)], (data[2:(len(midExposure) + 2), 1] - 148.6) / 105.4, 'o-')
  1146. print(2)
  1147. else:
  1148. plt.plot(arrayTimes[midExposure.astype(int)], (data[:len(midExposure), 1] - 148.6) / 105.4, 'o-')
  1149. print(0)
  1150. #plt.plot(softFrameTimes[data[:-6,0].astype(np.int)+6],(data[:-6,1]-148.6)/105.4)
  1151. #plt.plot(softFrameTimes,np.ones(len(softFrameTimes)),'|')
  1152. plt.show()
  1153. pdb.set_trace()
  1154. return (expStartTime,expEndTime,recordedFrames)
  1155. #################################################################################
  1156. # detect spikes in ephys trace
  1157. #################################################################################
  1158. def applyImageNormalizationMask(frames,imageMetaInfo,normFrame,normImageMetaInfo,mouse, date, rec):
  1159. print(imageMetaInfo, normImageMetaInfo)
  1160. pixelRange = 10
  1161. print('small, large frame : ', np.shape(frames), np.shape(normFrame))
  1162. print('pixel-ratio, x-ratio, y-ratio', imageMetaInfo[4]/normImageMetaInfo[4],end='')
  1163. fig = plt.figure()
  1164. rect1 = patches.Rectangle(normImageMetaInfo[:2], normImageMetaInfo[2], normImageMetaInfo[3],linewidth=1,edgecolor='C0',facecolor='none')
  1165. rect2 = patches.Rectangle(imageMetaInfo[:2],imageMetaInfo[2],imageMetaInfo[3],linewidth=1,edgecolor='C1',facecolor='none')
  1166. framesF = np.array(frames,dtype=float)
  1167. avgFrame = np.average(frames[:,:,:,0],axis=0)
  1168. # rescale image stack to the resolution of the normalization image
  1169. framesRescaled = scipy.ndimage.zoom(framesF, [1,imageMetaInfo[4]/normImageMetaInfo[4],imageMetaInfo[4]/normImageMetaInfo[4],1], order=3)
  1170. # average across all time points of image stack
  1171. #avgFrameZ = np.average(framesRescaled[:,:,:,0],axis=0)
  1172. # rescale the average image to match pixel-size of normalization image, the re-scaling factor of the ratio of the pixel-sizes : stack/norm
  1173. avgFrameZ = scipy.ndimage.zoom(avgFrame, imageMetaInfo[4]/normImageMetaInfo[4], order=3)
  1174. # x,y location in pixel indices of the stack in the normalization image
  1175. xLoc = int(np.round((imageMetaInfo[0] - normImageMetaInfo[0]) / normImageMetaInfo[4]))
  1176. yLoc = int(np.round((imageMetaInfo[1] - normImageMetaInfo[1]) / normImageMetaInfo[4]))
  1177. # dimensions of the rescaled image
  1178. xDim = np.shape(avgFrameZ)[0]
  1179. yDim = np.shape(avgFrameZ)[1]
  1180. #####################
  1181. ax0 = fig.add_subplot(2,3,4)
  1182. ax0.set_title('avg. of image stack',size=7)
  1183. ax0.imshow(np.transpose(avgFrame))
  1184. ax1 = fig.add_subplot(2,3,5)
  1185. ax1.set_title('avg. of image stack : rescaled to norm. image pixel size',size=7)
  1186. ax1.imshow(np.transpose(avgFrameZ))
  1187. print('image stack : ', np.shape(framesRescaled))
  1188. scipy.io.savemat('%s_%s_%s_imageStackBeforeRescaling.mat' % (mouse, date, rec), mdict={'dataArray': framesF})
  1189. scipy.io.savemat('%s_%s_%s_imageStack.mat' % (mouse, date, rec), mdict={'dataArray': framesRescaled})
  1190. #img_stack_uint8 = mapToXbit(avgFrameZ,8)
  1191. #pdb.set_trace()
  1192. #tiff.imsave('avg_imageStack_scaled.tif', np.array(img_stack_uint8, dtype=np.uint8))
  1193. ax2 = fig.add_subplot(2,3,1)
  1194. ax2.set_title('normalization image with image stack rectangle',size=7)
  1195. ret = patches.Rectangle([(imageMetaInfo[0]-normImageMetaInfo[0])/normImageMetaInfo[4],(imageMetaInfo[1]-normImageMetaInfo[1])/normImageMetaInfo[4]],imageMetaInfo[2]/normImageMetaInfo[4],imageMetaInfo[3]/normImageMetaInfo[4],linewidth=1,edgecolor='r',facecolor='none')
  1196. ax2.imshow(np.transpose(normFrame[0,:,:,0]))
  1197. ax2.add_patch(ret)
  1198. scipy.io.savemat('%s_%s_%s_registrationImage.mat' % (mouse, date, rec), mdict={'dataArray': normFrame[0,:,:,0]})
  1199. ax2 = fig.add_subplot(2,3,6)
  1200. ax2.set_title('area of image stack from normalization image',size=7)
  1201. #ret = patches.Rectangle([(imageMetaInfo[1]-normImageMetaInfo[1])/normImageMetaInfo[4],(imageMetaInfo[0]-normImageMetaInfo[0])/normImageMetaInfo[4]],imageMetaInfo[3]/normImageMetaInfo[4],imageMetaInfo[2]/normImageMetaInfo[4],linewidth=1,edgecolor='r',facecolor='none')
  1202. ax2.imshow(np.transpose(normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0]))
  1203. #ax2.add_patch(ret)
  1204. print('norm. image : ', np.shape(normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0]))
  1205. scipy.io.savemat('%s_%s_%s_normalizationImage.mat' % (mouse, date, rec), mdict={'dataArray': normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0]})
  1206. normFrameF = np.array(normFrame, dtype=float)
  1207. test1 = scipy.ndimage.gaussian_filter1d(normFrameF[:,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),:], 2, axis=1)
  1208. test2 = scipy.ndimage.gaussian_filter1d(test1, 2, axis=2)
  1209. #test1 = scipy.ndimage.gaussian_filter1d(framesF, 2, axis=1)
  1210. #test2 = scipy.ndimage.gaussian_filter1d(test1, 2, axis=2)
  1211. #norm8bit = mapToXbit(test2,8)
  1212. filter1D = scipy.ndimage.gaussian_filter1d(framesRescaled, 2, axis=1)
  1213. filter2D = scipy.ndimage.gaussian_filter1d(filter1D, 2, axis=2)
  1214. norm = filter2D / test2
  1215. norm8bit = mapToXbit(norm,8)
  1216. ax2 = fig.add_subplot(2,3,2)
  1217. ax2.set_title('normalized average image',size=7)
  1218. #ret = patches.Rectangle([(imageMetaInfo[1]-normImageMetaInfo[1])/normImageMetaInfo[4],(imageMetaInfo[0]-normImageMetaInfo[0])/normImageMetaInfo[4]],imageMetaInfo[3]/normImageMetaInfo[4],imageMetaInfo[2]/normImageMetaInfo[4],linewidth=1,edgecolor='r',facecolor='none')
  1219. ax2.imshow(np.transpose(np.average(norm[:,:,:,0],axis=0)))
  1220. plt.show()
  1221. #pdb.set_trace()
  1222. errMatrix = np.zeros((pixelRange*2+1,pixelRange*2+1))
  1223. # #row, col = np.indices(err)
  1224. xRange = np.arange(pixelRange*2+1)
  1225. yRange = np.copy(xRange)
  1226. # #avgFrameZ = np.copy(normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0])
  1227. for xy in itertools.product(xRange, yRange):
  1228. #err = np.abs((normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0] - avgFrameZ) ** 2).sum() / (xDim*yDim)
  1229. xStart = xLoc + xy[0] - pixelRange
  1230. yStart = yLoc + xy[1] - pixelRange
  1231. normImg = normFrameF[0,xStart:(xStart+xDim),yStart:(yStart+yDim),0]
  1232. normImgNorm = mapToXbit(normImg,8) #normImg - np.average(normImg)
  1233. avgFrameZNorm = mapToXbit(np.average(framesRescaled[:,:,:,0],axis=0),8) #avgFrameZ - np.average(avgFrameZ)
  1234. errMatrix[xy[0],xy[1]] = ((normImgNorm - avgFrameZNorm) ** 2).sum() / (xDim*yDim)
  1235. minimumIndices = np.argwhere(errMatrix == np.min(errMatrix))
  1236. print('MI :', minimumIndices)
  1237. #pdb.set_trace()
  1238. #xNorm = np.linspace(normImageMetaInfo[0],normImageMetaInfo[0]+
  1239. #ax.add_patch(rect1)
  1240. #ax.add_patch(rect2)
  1241. #ax.set_ylim(normImageMetaInfo[1]-10,normImageMetaInfo[1]+normImageMetaInfo[3]+10)
  1242. #ax.set_xlim(normImageMetaInfo[0]-10,normImageMetaInfo[0]+normImageMetaInfo[2]+10)
  1243. #plt.patches.Rectangle(normImageMetaInfo[:2],normImageMetaInfo[2],normImageMetaInfo[3])
  1244. #plt.patches.Rectangle(imageMetaInfo[:2],imageMetaInfo[2],imageMetaInfo[3])
  1245. #plt.show()
  1246. #pdb.set_trace()
  1247. return norm8bit
  1248. #################################################################################
  1249. # detect spikes in ephys trace
  1250. #################################################################################
  1251. def detectPawTrackingOutlies(pawTraces,pawMetaData):
  1252. jointNames = pawMetaData['data']['DLC-model-config file']['all_joints_names']
  1253. threshold = 60
  1254. def findOutliersBasedOnMaxSpeed(onePawData,jointName,i): # should be an 3 column array frame#, x, y
  1255. frDisplOrig = np.sqrt((np.diff(onePawData[:, 1])) ** 2 + (np.diff(onePawData[:, 2])) ** 2) / np.diff(onePawData[:, 0])
  1256. onePawDataTmp = np.copy(onePawData)
  1257. onePawIndicies = np.arange(len(onePawData))
  1258. # excursionsBoolOld = np.zeros(len(pawDataTmp)-1,dtype=bool)
  1259. nIt = 0
  1260. while True: # cycle as long as there are large displacements
  1261. frDispl = np.sqrt((np.diff(onePawDataTmp[:,1])) ** 2 + (np.diff(onePawDataTmp[:,2])) ** 2) / np.diff(onePawDataTmp[:, 0]) # calculate displacement
  1262. excursionsBoolTmp = frDispl > threshold # threshold displacement
  1263. print(nIt, sum(excursionsBoolTmp))
  1264. nIt += 1
  1265. if sum(excursionsBoolTmp) == 0: # no excursions above threshold are found anymore -> exit loop
  1266. break
  1267. else:
  1268. onePawDataTmp = onePawDataTmp[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
  1269. onePawIndicies = onePawIndicies[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
  1270. print('%s # of positions, # of detected mis-trackings, fraction : ' % (jointName), len(onePawData), len(onePawData) - len(onePawDataTmp), (len(onePawData) - len(onePawDataTmp)) / len(onePawData))
  1271. if jointName=='tail_base_bottom':
  1272. pdb.set_trace()
  1273. return (len(onePawData),len(onePawDataTmp),onePawIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig)
  1274. pawTrackingOutliers = []
  1275. for i in range(len(jointNames)):
  1276. (tot,correct,correctIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig) = findOutliersBasedOnMaxSpeed(np.column_stack((pawTraces[:,0],pawTraces[:,(i*3+1)],pawTraces[:,(i*3+2)])),jointNames[i],i)
  1277. pawTrackingOutliers.append([i,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig])
  1278. return pawTrackingOutliers
  1279. #pdb.set_trace()
  1280. #################################################################################
  1281. def detectPawTrackingOutliersObstacle(pawTraces,pawMetaData):
  1282. jointNames = pawMetaData['data']['DLC-model-config file']['all_joints_names']
  1283. threshold = 70
  1284. print(jointNames)
  1285. def findOutliersBasedOnMaxSpeedObstacle(onePawData,jointName,i): # should be an 3 column array frame#, x, y
  1286. # if jointName=='obstacle':
  1287. # threshold=100
  1288. # else:
  1289. # threshold = 80
  1290. frDisplOrig = np.sqrt((np.diff(onePawData[:, 1])) ** 2 + (np.diff(onePawData[:, 2])) ** 2) / np.diff(onePawData[:, 0])
  1291. onePawDataTmp = np.copy(onePawData)
  1292. onePawIndicies = np.arange(len(onePawData))
  1293. # excursionsBoolOld = np.zeros(len(pawDataTmp)-1,dtype=bool)
  1294. nIt = 0
  1295. while True: # cycle as long as there are large displacements
  1296. frDispl = np.sqrt((np.diff(onePawDataTmp[:,1])) ** 2 + (np.diff(onePawDataTmp[:,2])) ** 2) / np.diff(onePawDataTmp[:, 0]) # calculate displacement
  1297. excursionsBoolTmp = frDispl > threshold # threshold displacement
  1298. print(nIt, sum(excursionsBoolTmp))
  1299. nIt += 1
  1300. if sum(excursionsBoolTmp) == 0: # no excursions above threshold are found anymore -> exit loop
  1301. break
  1302. else:
  1303. onePawDataTmp = onePawDataTmp[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
  1304. onePawIndicies = onePawIndicies[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
  1305. print('%s # of positions, # of detected mis-trackings, fraction : ' % (jointName), len(onePawData), len(onePawData) - len(onePawDataTmp), (len(onePawData) - len(onePawDataTmp)) / len(onePawData))
  1306. # if jointName=='tail_base_bottom':
  1307. # # pdb.set_trace()
  1308. return (len(onePawData),len(onePawDataTmp),onePawIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig)
  1309. pawTrackingOutliersDic = {}
  1310. pawTrackingOutliersList=[]
  1311. pawTrackingOutliersBot_paw=[]
  1312. b=0
  1313. bot_paw = ['front_left_bottom', 'front_right_bottom', 'hind_left_bottom', 'hind_right_bottom']
  1314. for i in range(len(jointNames)):
  1315. pawTrackingOutliersDic[jointNames[i]] = {}
  1316. (tot,correct,correctIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig) = findOutliersBasedOnMaxSpeedObstacle(np.column_stack((pawTraces[:,0],pawTraces[:,(i*3+1)],pawTraces[:,(i*3+2)])),jointNames[i],i)
  1317. pawTrackingOutliersList.append([i,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig]) #all pf these are parameters that we stock in each jointName
  1318. parameters = [i,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig]
  1319. parameters_string = ['i','tot','correct','correctIndicies','jointName','onePawData','onePawDataTmp','frDispl','frDisplOrig']
  1320. if any([x in jointNames[i] for x in bot_paw]):
  1321. pawTrackingOutliersBot_paw.append([b,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig])
  1322. b+=1
  1323. for j in range(len(parameters)):
  1324. pawTrackingOutliersDic[jointNames[i]][parameters_string[j]] = parameters[j]
  1325. return pawTrackingOutliersDic,pawTrackingOutliersList,pawTrackingOutliersBot_paw
  1326. #################################################################################
  1327. def detectPawTrackingOutliersObstacleVids(pawTraces, pawMetaData):
  1328. jointNames = pawMetaData['data']['DLC-model-config file']['all_joints_names']
  1329. threshold = 70
  1330. print(jointNames)
  1331. def findOutliersBasedOnMaxSpeedObstacle(onePawData, jointName, i): # should be an 3 column array frame#, x, y
  1332. # if jointName=='obstacle':
  1333. # threshold=100
  1334. # else:
  1335. # threshold = 80
  1336. frDisplOrig = np.sqrt((np.diff(onePawData[:, 1])) ** 2 + (np.diff(onePawData[:, 2])) ** 2) / np.diff(
  1337. onePawData[:, 0])
  1338. onePawDataTmp = np.copy(onePawData)
  1339. onePawIndicies = np.arange(len(onePawData))
  1340. # excursionsBoolOld = np.zeros(len(pawDataTmp)-1,dtype=bool)
  1341. nIt = 0
  1342. while True: # cycle as long as there are large displacements
  1343. frDispl = np.sqrt((np.diff(onePawDataTmp[:, 1])) ** 2 + (np.diff(onePawDataTmp[:, 2])) ** 2) / np.diff(
  1344. onePawDataTmp[:, 0]) # calculate displacement
  1345. excursionsBoolTmp = frDispl > threshold # threshold displacement
  1346. print(nIt, sum(excursionsBoolTmp))
  1347. nIt += 1
  1348. if sum(excursionsBoolTmp) == 0: # no excursions above threshold are found anymore -> exit loop
  1349. break
  1350. else:
  1351. onePawDataTmp = onePawDataTmp[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
  1352. onePawIndicies = onePawIndicies[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
  1353. print('%s # of positions, # of detected mis-trackings, fraction : ' % (jointName), len(onePawData),
  1354. len(onePawData) - len(onePawDataTmp), (len(onePawData) - len(onePawDataTmp)) / len(onePawData))
  1355. # if jointName=='tail_base_bottom':
  1356. # # pdb.set_trace()
  1357. return (len(onePawData), len(onePawDataTmp), onePawIndicies, onePawData, onePawDataTmp, frDispl, frDisplOrig)
  1358. pawTrackingOutliersDic = {}
  1359. pawTrackingOutliersList = []
  1360. pawTrackingOutliersBot_paw = []
  1361. b = 0
  1362. bot_paw = ['front_left_bottom', 'front_right_bottom', 'hind_left_bottom', 'hind_right_bottom']
  1363. for i in range(len(jointNames)):
  1364. pawTrackingOutliersDic[jointNames[i]] = {}
  1365. tot=[]
  1366. correct=[]
  1367. correctIndicies=np.array([])
  1368. onePawData=np.empty((3))
  1369. onePawDataTmp=np.empty((3))
  1370. frDispl=np.array([])
  1371. frDisplOrig=np.array([])
  1372. for v in np.unique(pawMetaData['obs_number']):
  1373. # print('obstacle ids', np.unique(pawMetaData['obs_number']))
  1374. vmask=pawMetaData['obs_number']==v
  1375. try:
  1376. (tot_v, correct_v, correctIndicies_v, onePawData_v, onePawDataTmp_v, frDispl_v,frDisplOrig_v) = findOutliersBasedOnMaxSpeedObstacle(np.column_stack((pawTraces[:, 0][vmask], pawTraces[:, (i * 3 + 1)][vmask], pawTraces[:, (i * 3 + 2)][vmask])), jointNames[i],i)
  1377. except:
  1378. print('missmatch between frame labeled and obstacle frame numbers for label', jointNames[i], len(pawTraces[:, 0]), len(pawMetaData['obs_number']), 'please regenerate video with proper angle range and analyze with DLC')
  1379. pdb.set_trace()
  1380. # print(len(onePawDataTmp_v), len(frDispl_v), len(onePawData_v), len(frDisplOrig_v))
  1381. correctIndicies=np.concatenate((correctIndicies,correctIndicies_v))
  1382. onePawData=np.vstack((onePawData,onePawData_v))
  1383. onePawDataTmp=np.vstack((onePawDataTmp,onePawDataTmp_v))
  1384. frDispl = np.concatenate((frDispl, frDispl_v))
  1385. frDisplOrig=np.concatenate((frDisplOrig, frDisplOrig_v))
  1386. onePawData=onePawData[1:]
  1387. onePawDataTmp=onePawDataTmp[1:]
  1388. # pdb.set_trace()
  1389. pawTrackingOutliersList.append([i, tot, correct, correctIndicies, jointNames[i], onePawData, onePawDataTmp, frDispl,frDisplOrig]) # all pf these are parameters that we stock in each jointName
  1390. parameters = [i, tot, correct, correctIndicies, jointNames[i], onePawData, onePawDataTmp, frDispl,frDisplOrig]
  1391. parameters_string = ['i', 'tot', 'correct', 'correctIndicies', 'jointName', 'onePawData', 'onePawDataTmp','frDispl', 'frDisplOrig']
  1392. if any([x in jointNames[i] for x in bot_paw]):
  1393. pawTrackingOutliersBot_paw.append(
  1394. [b, tot, correct, correctIndicies, jointNames[i], onePawData, onePawDataTmp, frDispl, frDisplOrig])
  1395. b += 1
  1396. for j in range(len(parameters)):
  1397. pawTrackingOutliersDic[jointNames[i]][parameters_string[j]] = parameters[j]
  1398. return pawTrackingOutliersDic, pawTrackingOutliersList, pawTrackingOutliersBot_paw
  1399. #################################################################################
  1400. #################################################################################
  1401. # convert ca traces in easily usable numpy array
  1402. #################################################################################
  1403. def getCaWheelPawInterpolatedDictsPerDay(nSess,allCorrDataPerSession,allStepData,showFig = False):
  1404. baselineTime = 5.
  1405. # calcium traces ##############################################################
  1406. trialStartUnixTimes = []
  1407. fTraces = allCorrDataPerSession[nSess]['caImg']['Fluo'] #[3][0][0]
  1408. timeStamps = allCorrDataPerSession[nSess]['caImg']['timeStamps'] # [3][0][3] # the array containing the time-stamp array
  1409. recordings = np.unique(timeStamps[:, 1]) # determine how many recordings where performed
  1410. caTracesDict = {}
  1411. for n in range(len(recordings)):
  1412. mask = (timeStamps[:, 1] == recordings[n])
  1413. triggerStart = timeStamps[:, 5][mask]
  1414. trialStartUnixTimes.append(timeStamps[:, 3][mask][0])
  1415. if n > 0:
  1416. if oldTriggerStart > triggerStart[0]:
  1417. print('problem in trial order')
  1418. sys.exit(1)
  1419. # for i in range(len(fTraces)):
  1420. # triggerstart - time of the acq start trigger for the current acquisition
  1421. # timeStamps[:, 4][mask] - time of the first pixel in the frame passed since acqModeEpoch
  1422. caTracesTime = (timeStamps[:, 4][mask] - triggerStart) # triggerStart is negative
  1423. #pdb.set_trace()
  1424. caTracesFluo = fTraces[:, mask]
  1425. # pdb.set_trace()
  1426. # caTraces.append(np.column_stack((caTracesTime,caTracesFluo)))
  1427. caTracesDict[n] = np.row_stack((caTracesTime, caTracesFluo))
  1428. #print(np.shape(np.row_stack((caTracesTime, caTracesFluo))))
  1429. oldTriggerStart = triggerStart[0]
  1430. # wheel speed ######################################################
  1431. # also find calmest pre-motorization period
  1432. minPreMotorMeanV = 1000.
  1433. wheelTracks = allCorrDataPerSession[nSess]['wheel'] #[1]
  1434. nRec = 0
  1435. # print(len(wheelTracks))
  1436. wheelSpeedDict = {}
  1437. for n in range(len(wheelTracks)):
  1438. wheelRecStartTime = wheelTracks[n]['timeStamp']#[3]
  1439. if (trialStartUnixTimes[nRec] - wheelRecStartTime) < 1.:
  1440. # if not wheelTracks[n][4]:
  1441. # recStartTime = wheelTracks[0][3]
  1442. if nRec > 0:
  1443. if oldRecStartTime > wheelRecStartTime:
  1444. print('problem in trial order')
  1445. sys.exit(1)
  1446. wheelTime = wheelTracks[n]['sTimes']#[2]
  1447. wheelSpeed = wheelTracks[n]['linearSpeed']#[1] # linear wheel speed in cm/s
  1448. angleSpeed = wheelTracks[n]['angluarSpeed']#[0]
  1449. angleTime = wheelTracks[n]['angleTimes']#[5]
  1450. wheelSpeedDict[nRec] = np.row_stack((wheelTime, wheelSpeed))
  1451. #pdb.set_trace()
  1452. preMMask = (wheelTime < baselineTime)
  1453. preMmeanV = np.mean(np.abs(wheelSpeed[preMMask]))
  1454. #print(nSess, nRec, preMmeanV)
  1455. if preMmeanV < minPreMotorMeanV:
  1456. slowestRec = nRec
  1457. minPreMotorMeanV = np.copy(preMmeanV)
  1458. nRec += 1
  1459. oldRecStartTime = wheelRecStartTime
  1460. print('trials with slowest baseline period :', slowestRec)
  1461. # normalize ca-traces by baseline fluorescence : fluorescence during the first baselineTime seconds in the least active recording #############################################
  1462. mask = (caTracesDict[slowestRec][0] < baselineTime)
  1463. F0 = np.mean(caTracesDict[slowestRec][1:][:, mask], axis=1)
  1464. #pdb.set_trace()
  1465. for n in range(len(recordings)):
  1466. normalizedCaTraces = (caTracesDict[n][1:] - F0[:, np.newaxis]) / F0[:, np.newaxis]
  1467. caTracesDict[n][1:] = np.copy(normalizedCaTraces)
  1468. # pdb.set_trace()
  1469. # paw speed ######################################################
  1470. pawTracks = allCorrDataPerSession[nSess]['paws']#[2]
  1471. nRec = 0
  1472. pawTracksDict = {}
  1473. pawID = []
  1474. for n in range(len(pawTracks)):
  1475. # if not wheelTracks[n][4]:
  1476. pawRecStartTime = pawTracks[n]['recStartTime']#[4]
  1477. if (trialStartUnixTimes[nRec] - pawRecStartTime) < 1.:
  1478. if nRec > 0:
  1479. if oldRecStartTime > pawRecStartTime:
  1480. print('problem in trial order')
  1481. sys.exit(1)
  1482. pawTracksDict[nRec] = {}
  1483. for i in range(4):
  1484. # pdb.set_trace()
  1485. if nRec == 0:
  1486. pawID.append(pawTracks[n]['jointNamesFramesInfo'][i][0])
  1487. # pawTracksDict[nFig][i]['pawID'] = pawTracks[n][2][i][0]
  1488. pawSpeedTime = pawTracks[n]['pawSpeed'][i][:,0] # times of cleared paw speed
  1489. pawSpeed = pawTracks[n]['pawSpeed'][i][:,1] # 1 is combined speed in 2-d plane of the camera view from below,
  1490. pawTracksDict[nRec][i] = np.row_stack((pawSpeedTime,pawSpeed)) # interp = interp1d(pawSpeedTime, pawSpeed) # newPawSpeedAtCaTimes = interp(caTracesTime[nFig]) # pawTracksDict[i]['pawSpeed'].extend(newPawSpeedAtCaTimes)
  1491. oldRecStartTime = pawRecStartTime
  1492. nRec += 1
  1493. # interpolation #############################################################################
  1494. # interp = interp1d(wheelTime, wheelSpeed)
  1495. # interpMask = (caTracesTime[nFig] >= wheelTime[0]) & (caTracesTime[nFig] <= wheelTime[-1])
  1496. # newWheelSpeedAtCaTimes = interp(caTracesTime[nFig][interpMask])
  1497. # wheelSpeedAll.extend(newWheelSpeedAtCaTimes)
  1498. wheelSpeedDictInterp = wheelSpeedDict.copy()
  1499. pawTracksDictInterp = pawTracksDict.copy()
  1500. caTracesDictInterp = caTracesDict.copy()
  1501. for nrec in range(len(caTracesDict)):
  1502. # determine interpolation range
  1503. startInterpTime = np.max((caTracesDict[nrec][0, 0], wheelSpeedDict[nrec][0, 0], pawTracksDict[nrec][0][0, 0], pawTracksDict[nrec][1][0, 0], pawTracksDict[nrec][2][0, 0], pawTracksDict[nrec][3][0, 0]))
  1504. endInterpTime = np.min((caTracesDict[nrec][0, -1], wheelSpeedDict[nrec][0, -1], pawTracksDict[nrec][0][0, -1], pawTracksDict[nrec][1][0, -1], pawTracksDict[nrec][2][0, -1], pawTracksDict[nrec][3][0, -1]))
  1505. interpMask = (caTracesDict[nrec][0] >= startInterpTime) & (caTracesDict[nrec][0] <= endInterpTime)
  1506. # restrict ca-traces to interpolation range
  1507. #pdb.set_trace()
  1508. #matrix = np.copy(caTracesDict[nrec])
  1509. caTracesDictInterp[nrec] = caTracesDict[nrec][:,interpMask]
  1510. # interpolate wheel speed
  1511. interpWheel = interp1d(wheelSpeedDict[nrec][0], wheelSpeedDict[nrec][1])#,kind='cubic')
  1512. newWheelSpeedAtCaTimes = interpWheel(caTracesDict[nrec][0][interpMask])
  1513. wheelSpeedDictInterp[nrec] = np.row_stack((caTracesDict[nrec][0][interpMask], newWheelSpeedAtCaTimes))
  1514. # interpolate paw speed
  1515. for i in range(4):
  1516. interpPaw = interp1d(pawTracksDict[nrec][i][0], pawTracksDict[nrec][i][1])#,kind='cubic')
  1517. newPawSpeedAtCaTimes = interpPaw(caTracesDict[nrec][0][interpMask])
  1518. pawTracksDictInterp[nrec][i] = np.row_stack((caTracesDict[nrec][0][interpMask], newPawSpeedAtCaTimes))
  1519. #pawAll[i].extend(newPawSpeedAtCaTimes)
  1520. if showFig:
  1521. cc = ['C0','C1']
  1522. fig, ax = plt.subplots(figsize=(10, 6))
  1523. for i in range(2):
  1524. ax.plot(pawTracksDictInterp[nrec][i][0],pawTracksDictInterp[nrec][i][1]/np.max(pawTracksDictInterp[nrec][i][1]),c=cc[i])
  1525. idxSwings = allStepData[nSess][4][nrec][3][i][1]
  1526. # print('Wow nSess',nDay,allStepData[nDay-1][0])
  1527. recTimes = allStepData[nSess][4][nrec][4][i][2]
  1528. # pdb.set_trace()
  1529. idxSwings = np.asarray(idxSwings)
  1530. for k in range(len(idxSwings)): # loop over all swings
  1531. startSwingTime = recTimes[idxSwings[k, 0]]
  1532. endSwingTime = recTimes[idxSwings[k, 1]]
  1533. ax.fill_between((startSwingTime,endSwingTime), 0, 1, color='0.6', alpha=0.5, transform=ax.get_xaxis_transform())
  1534. ax.fill_between((startSwingTime,endSwingTime), 0, 1, color='0.4', alpha=0.5, transform=ax.get_xaxis_transform())
  1535. #ax.plot(caTracesDictInterp[nrec][0],caTracesDictInterp[nrec][1]/np.max(caTracesDictInterp[nrec][1]),c='black')
  1536. plt.show()
  1537. pdb.set_trace()
  1538. return (wheelSpeedDictInterp,pawTracksDictInterp,caTracesDictInterp,wheelSpeedDict,pawTracksDict,caTracesDict,slowestRec)
  1539. #################################################################################
  1540. # calculate correlations between ca-imaging, wheel speed and paw speed
  1541. #################################################################################
  1542. def doRegressionAnalysis(mouse,allCorrDataPerSession,allStepData,borders=None,figShow=False):
  1543. matplotlib.use('TkAgg') # WxAgg
  1544. from sklearn.linear_model import LinearRegression
  1545. from sklearn.svm import SVR
  1546. #SVR(kernel='rbf', C=1e3, gamma=0.1)
  1547. #from sklearn.ensemble import RandomForestRegressor
  1548. regressionN = 6
  1549. Rvalues = []
  1550. #for nSess in range(len(allCorrDataPerSession)):
  1551. for nDay in range(len(allCorrDataPerSession)):
  1552. print(nDay,allCorrDataPerSession[nDay]['folder'])
  1553. (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP, wheelSpeedDict, pawTracksDict, caTracesDict, slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay,allCorrDataPerSession,allStepData)
  1554. # ATTENTION : all of the arrays also contain a time array
  1555. # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTimeSame]
  1556. # dims of pawTracksDictInterP : [nSessions][nPaw][2][valuesOverTimeSame]
  1557. # dims of caTracesDictInterP : [nSessions][nRois+1][valuesOverTimeSame]
  1558. # dims of wheelSpeedDict : [nSessions][2][valuesOverTime]
  1559. # dims of pawTracksDict : [nSessions][nPaw][2][valuesOverTime]
  1560. # dims of caTracesDict : [nSessions][nRois+1][valuesOverTime]
  1561. #(wheelSpeedDict, pawTracksDict, caTracesDict,aa,bb,cc) = getCaWheelPawInterpolatedDictsPerDay(nSess, allCorrDataPerSession)
  1562. nRecWheel = len(wheelSpeedDictInterP)
  1563. nRecPaw = len(pawTracksDictInterP)
  1564. nRecCa = len(caTracesDictInterP)
  1565. print('Recording length :', nRecWheel, nRecPaw, nRecCa)
  1566. if (nRecWheel != nRecPaw ) or (nRecWheel != nRecCa):
  1567. print('problem in number of recordings listed in dictionaries')
  1568. # loop over 5 different regressions, each using a different combination of test and train samples
  1569. recs = range(nRecCa)
  1570. RTempValues = []
  1571. for reg in range(nRecCa): # loop over all recordings
  1572. recsForTraining = recs.copy()
  1573. recsForTraining.remove(reg)
  1574. recsForTest = [reg]
  1575. Rval1 = []
  1576. for d in range(regressionN): # loop over wheel speed, the four paw speeds and the combined speed
  1577. #print('recording iteration %s, variable %s' %(reg,d))
  1578. #pdb.set_trace()
  1579. # concatenate data
  1580. if borders is not None:
  1581. timeMaskTrain = (caTracesDictInterP[recsForTraining[0]][0]>=borders[0])&(caTracesDictInterP[recsForTraining[0]][0]<=borders[1])
  1582. timeMaskTest = (caTracesDictInterP[recsForTest[0]][0] >= borders[0]) & (caTracesDictInterP[recsForTest[0]][0] <= borders[1])
  1583. else:
  1584. timeMaskTrain = (caTracesDictInterP[recsForTraining[0]][0]>=0)&(caTracesDictInterP[recsForTraining[0]][0]<=1000.)
  1585. timeMaskTest = (caTracesDictInterP[recsForTest[0]][0]>=0)&(caTracesDictInterP[recsForTest[0]][0]<=1000.)
  1586. X = np.copy(caTracesDictInterP[recsForTraining[0]][1:][:,timeMaskTrain])
  1587. Xtest = np.copy(caTracesDictInterP[recsForTest[0]][1:][:,timeMaskTest])
  1588. #pdb.set_trace()
  1589. if d == 0:
  1590. Y = np.copy(wheelSpeedDictInterP[recsForTraining[0]][1:][:,timeMaskTrain])
  1591. Ytest = np.copy(wheelSpeedDictInterP[recsForTest[0]][1:][:,timeMaskTest])
  1592. YtestTime = np.copy(wheelSpeedDictInterP[recsForTest[0]][0][timeMaskTest])
  1593. elif (d>0) and (d<5):
  1594. pawId = d-1
  1595. Y = np.copy(pawTracksDictInterP[recsForTraining[0]][pawId][1:][:,timeMaskTrain])
  1596. Ytest = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][1:][:,timeMaskTest])
  1597. YtestTime = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][0][timeMaskTest])
  1598. elif d==5: # case where all four paw speeds are added together
  1599. pawSpeedTrain = []
  1600. pawSpeedTest = []
  1601. for i in range(4):
  1602. pawSpeedTrain.append(np.copy(pawTracksDictInterP[recsForTraining[0]][i][1:][:, timeMaskTrain]))
  1603. pawSpeedTest.append(np.copy(pawTracksDictInterP[recsForTest[0]][i][1:][:, timeMaskTest]))
  1604. Y = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
  1605. Ytest = pawSpeedTest[0] + pawSpeedTest[1] + pawSpeedTest[2] + pawSpeedTest[3]
  1606. #pdb.set_trace()
  1607. for t in recsForTraining[1:]:
  1608. if borders is not None:
  1609. timeMaskTrain = (caTracesDictInterP[t][0] >= borders[0]) & (caTracesDictInterP[t][0] <= borders[1])
  1610. else:
  1611. timeMaskTrain = (caTracesDictInterP[t][0] >= 0) & (caTracesDictInterP[t][0] <= 1000.)
  1612. X = np.column_stack((X,caTracesDictInterP[t][1:][:,timeMaskTrain]))
  1613. if d == 0:
  1614. Y = np.column_stack((Y,wheelSpeedDictInterP[t][1:][:,timeMaskTrain]))
  1615. elif (d>0) and (d<5):
  1616. Y = np.column_stack((Y, pawTracksDictInterP[t][pawId][1:][:,timeMaskTrain]))
  1617. elif d==5:
  1618. pawSpeedTrain = []
  1619. for i in range(4):
  1620. pawSpeedTrain.append(np.copy(pawTracksDictInterP[t][i][1:][:, timeMaskTrain]))
  1621. speedTemp = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
  1622. Y = np.column_stack((Y, speedTemp))
  1623. #pdb.set_trace()
  1624. Y = Y[0]
  1625. X = np.transpose(X)
  1626. Ytest = Ytest[0]
  1627. Xtest = np.transpose(Xtest)
  1628. # linear regression ########################################
  1629. linReg = LinearRegression()
  1630. linReg.fit(X,Y)
  1631. #svm_rbf = SVR(kernel='rbf', C=1e3, gamma=0.1)
  1632. #svm_rbf.fit(X,Y)
  1633. YTrainPred = linReg.predict(X)
  1634. YTestPred = linReg.predict(Xtest)
  1635. R2trainLR = linReg.score(X, Y)
  1636. R2testLR = linReg.score(Xtest, Ytest) # 1. - np.sum((Ytest-YTestPred)**2)/np.sum((Ytest - np.mean(Ytest))**2)#linReg.score(Xtest, Ytest)
  1637. #print(linReg.coef_)
  1638. #print(linReg.intercept_)
  1639. #yPred = linReg.predict(np.transpose(X))
  1640. # random forest ##############################################
  1641. #randForestReg = RandomForestRegressor(n_estimators=20)
  1642. #randForestReg.fit(X, Y)
  1643. #R2trainRF = randForestReg.score(X, Y)
  1644. #R2testRF= randForestReg.score(Xtest, Ytest)
  1645. #
  1646. if figShow :
  1647. print('R2 test :',R2testLR)
  1648. fig = plt.figure()
  1649. ax = fig.add_subplot(111)
  1650. ax.plot(YtestTime,Ytest,lw=2)
  1651. ax.plot(YtestTime,YTestPred,lw=2)
  1652. ax.spines['top'].set_visible(False)
  1653. ax.spines['right'].set_visible(False)
  1654. ax.spines['bottom'].set_position(('outward', 10))
  1655. ax.spines['left'].set_position(('outward', 10))
  1656. ax.yaxis.set_ticks_position('left')
  1657. ax.xaxis.set_ticks_position('bottom')
  1658. plt.show()
  1659. Rval1.extend([R2trainLR,R2testLR])
  1660. RTempValues.append(Rval1)
  1661. #pdb.set_trace()
  1662. Rs = np.zeros(regressionN*2)
  1663. for reg in range(5):
  1664. Rs += RTempValues[reg]
  1665. Rs /=5.
  1666. Rvalues.append(Rs)
  1667. return Rvalues
  1668. #################################################################################
  1669. # perform linear regression between spiking activity and behavioral measures
  1670. #################################################################################
  1671. def crossValidatedRegression(regModels,X,y,t,fold,visualize=False):
  1672. cols = ['C0','C1','C2','C3','C4','C5','C6','C7','C8','C9','C10']
  1673. regOutput = {}
  1674. ###########
  1675. if visualize:
  1676. plt.plot(t,y,color='black')
  1677. # get coeffss : apply the regression models consecutively
  1678. for j in range(len(regModels)):
  1679. regOutput[j] = {}
  1680. print('applying', regModels[j][0])
  1681. regModel = regModels[j][1]
  1682. regModel.fit(X, y)
  1683. # print(regModel.coef_)
  1684. regOutput[j]['name'] = regModels[j][0]
  1685. regOutput[j]['coefficients'] = regModel.coef_
  1686. regOutput[j]['fitScore'] = regModel.score(X, y)
  1687. regOutput[j]['scores'] = np.zeros(fold*5)
  1688. if visualize:
  1689. plt.plot(t,regModel.predict(X),c=cols[j],label=regModels[j][0]+': %s ' % (np.round(regOutput[j]['fitScore'],3)))
  1690. del regModel
  1691. if visualize:
  1692. plt.xlabel('time (s)')
  1693. plt.ylabel('firing rate')
  1694. plt.legend(frameon=False)
  1695. plt.show()
  1696. ########
  1697. if fold>0:
  1698. print('performing cross-validation')
  1699. # to k-fold cross-validated regression to access the score
  1700. for j in range(len(regModels)):
  1701. #kf = KFold(n_splits=fold)
  1702. i = 0
  1703. #scores = cross_val_score(regModels[j][1], X, y, cv=fold)
  1704. #print(regModels[j][0],scores)
  1705. rkf = RepeatedKFold(n_splits=fold, n_repeats=5, random_state=42)
  1706. #for train_index, test_index in kf.split(X):
  1707. for train_index, test_index in rkf.split(X):
  1708. #print(i,len(train_index),len(test_index))
  1709. X_train, X_test = X[train_index], X[test_index]
  1710. y_train, y_test = y[train_index], y[test_index]
  1711. regModel = regModels[j][1]
  1712. regModel.fit(X_train, y_train)
  1713. regOutput[j]['scores'][i] = regModel.score(X_test, y_test)
  1714. #if visualize:
  1715. # plt.plot(t[test_index],regModel.predict(X_test),label=(regModels[j][0] if i==0 else None))
  1716. del regModel
  1717. i+=1
  1718. for j in range(len(regModels)):
  1719. print(regOutput[j]['name'],regOutput[j]['fitScore'],np.mean(regOutput[j]['scores']),np.std(regOutput[j]['scores']))#,regOutput[j]['scores'])
  1720. return regOutput
  1721. # # print(regModel.intercept_)
  1722. # sc = regModel.score(Xregressors_scaled[tmask], YspikeCount[tmask])
  1723. # print('score:', sc)
  1724. # Ypred = regModel.predict(Xregressors_scaled[tmask])
  1725. # del regModel
  1726. # plt.plot(tbinCenters[tmask], Ypred, label=regs[j][0] + ' %s' % np.round(sc, 3))
  1727. # plt.xlabel('time (s)')
  1728. # plt.ylabel('firing rate') # pass
  1729. #################################################################################
  1730. # shuffles variable within chunks
  1731. #################################################################################
  1732. def shuffleVariable(y,ttime,dt,chunkSize):
  1733. nChunk = int(chunkSize/dt)
  1734. shuffleIterations = int(np.ceil(len(ttime)/nChunk))
  1735. for i in range(shuffleIterations):
  1736. np.random.shuffle(y[(i*nChunk):((i+1)*nChunk)])
  1737. return y
  1738. #################################################################################
  1739. # perform linear regression between spiking activity and behavioral measures
  1740. #################################################################################
  1741. # cPawPos,pawSpeed,ephys,swingStanceDict,sTimes,linearSpeed)
  1742. def performGLManalysis(date, rec, pawPos,pawSpeed,ephys,swingStanceD,sTimes,linearSpeed):
  1743. matplotlib.use('TkAgg')
  1744. # create spike-count vector
  1745. spikeTimes = ephys[0]
  1746. print('firing rate :', 1./np.mean(np.diff(spikeTimes)))
  1747. dt = 0.01
  1748. shiftRange = 0.2 # binary columns are shifted back- and forth in time by this delay in s
  1749. nShift = int(shiftRange/dt)
  1750. tbins = np.linspace(0., 60., int(60 / dt) + 1, endpoint=True)
  1751. tbinCenters = (tbins[1:]+tbins[:-1])/2
  1752. binnedspikes, _ = np.histogram(spikeTimes, tbins)
  1753. spikecountwindow = 0.02
  1754. nspikecountwindow = int(spikecountwindow / dt) # + 0.5)
  1755. #YspikeCount = np.convolve(binnedspikes, np.ones(nspikecountwindow), 'same')
  1756. binnedspikes=np.array(binnedspikes,dtype=float)
  1757. YspikeCount = scipy.ndimage.gaussian_filter1d(binnedspikes, nspikecountwindow,axis=0) # convolve with Gaussian kernel
  1758. print(len(YspikeCount),len(binnedspikes),nspikecountwindow)
  1759. # create regressor matrix
  1760. #Xregressors = np.zeros((len(tbinCenters),1+4*4))
  1761. Xregressors = np.zeros((len(tbinCenters), 1))
  1762. # interpolate wheel speed
  1763. interpWheel = interp1d(sTimes, linearSpeed,fill_value='extrapolate')#,kind='cubic')
  1764. Xregressors[:,0] = interpWheel(tbinCenters)
  1765. # interpolate paw position and paw speed
  1766. for i in range(4):
  1767. interpPawPos = interp1d(pawPos[i][:,0],pawPos[i][:,1],fill_value='extrapolate')
  1768. #Xregressors[:,1+i] = interpPawPos(tbinCenters)
  1769. Xregressors = np.column_stack((Xregressors,interpPawPos(tbinCenters)))
  1770. for i in range(4):
  1771. interpPawSpeed = interp1d(pawSpeed[i][:,0],pawSpeed[i][:,2],fill_value='extrapolate') # pawSpeed[i][1] is total speed, 2 is x, 3 is y speed
  1772. #Xregressors[:,5+i] = interpPawSpeed(tbinCenters)
  1773. Xregressors = np.column_stack((Xregressors, interpPawSpeed(tbinCenters)))
  1774. for i in range(4):
  1775. idxSwings = swingStanceD['swingP'][i][1]
  1776. recTimes = swingStanceD['forFit'][i][2]
  1777. idxSwings = np.asarray(idxSwings)
  1778. binnedSwingStartTimes, _ = np.histogram(recTimes[idxSwings[:,0]], tbins)
  1779. binnedSwingEndTimes, _ = np.histogram(recTimes[idxSwings[:, 1]], tbins)
  1780. startTshift = np.zeros((len(binnedSwingStartTimes),2*nShift+1))
  1781. endTshift = np.zeros((len(binnedSwingEndTimes),2*nShift+1))
  1782. n = 0
  1783. for j in range(-nShift,nShift+1):
  1784. startTshift[:,n] = np.roll(binnedSwingStartTimes,j)
  1785. endTshift[:,n] = np.roll(binnedSwingEndTimes, j)
  1786. #plt.plot(startTshift[:,n])
  1787. n+=1
  1788. #plt.show()
  1789. Xregressors = np.column_stack((Xregressors, startTshift))
  1790. Xregressors = np.column_stack((Xregressors, endTshift))
  1791. #Xregressors[:,9+i] = binnedSwingStartTimes
  1792. #Xregressors[:,13+i] = binnedSwingEndTimes
  1793. #pdb.set_trace()
  1794. print('shape or regressor matrix :', np.shape(Xregressors))
  1795. ## preprocessing of the data
  1796. #scaler = preprocessing.MinMaxScaler().fit(Xregressors) # This estimator scales and translates each feature individually such that the maximal absolute value of each feature in the training set will be 1.0. It does not shift/center the data, and thus does not destroy any sparsity.
  1797. #Xregressors_scaled = scaler.transform(Xregressors)
  1798. #pdb.set_trace()
  1799. Xregressors_scaled = np.copy(Xregressors)
  1800. Xregressors_scaled[:,:9] = scipy.stats.zscore(Xregressors[:,:9],axis=0)
  1801. #Xregressors_scaled = np.copy(Xregressors)
  1802. #pdb.set_trace()
  1803. #Xregressors_scaled[:8] = Xregressors[:8] # preserve the sparse data
  1804. timeLimits = [10,50]
  1805. tmask = (tbinCenters>timeLimits[0]) & (tbinCenters<timeLimits[1])
  1806. # generate list of regression models
  1807. # alpha multiplies the penalty terms : for alpha=0 is equivalent to an ordinary least square
  1808. # ridge regression : l2 regularization
  1809. # elastic net : The ElasticNet mixing parameter, with 0 <= l1_ratio <= 1. For l1_ratio = 0 the penalty is an L2 penalty. For l1_ratio = 1 it is an L1 penalty. For 0 < l1_ratio < 1, the penalty is a combination of L1 and L2.
  1810. regs = [('Linear Regression',linear_model.LinearRegression()),('Ridge regression',linear_model.Ridge(alpha=100.))]#,('Elastic Net Regression',linear_model.ElasticNet(alpha=0.01,l1_ratio=0.5,random_state=0))]#,('GLM with log link function',linear_model.PoissonRegressor(alpha=1e-6/len(YspikeCount[tmask])))]
  1811. #regs = [('Elastic Net Regression',linear_model.ElasticNet(alpha=0.01,l1_ratio=0.5,random_state=0))]\
  1812. # ('Ridge regression', linear_model.Ridge(alpha=1.)),
  1813. # ('Elastic Net Regression', linear_model.ElasticNet(alpha=0.01, random_state=0))]
  1814. regResults = crossValidatedRegression(regs,Xregressors_scaled[tmask],YspikeCount[tmask],tbinCenters[tmask],fold=0,visualize=False)
  1815. print(len(Xregressors_scaled[tmask]))
  1816. #plt.plot(tbinCenters[tmask],YspikeCount[tmask])
  1817. #coeffss = []
  1818. #pdb.set_trace()
  1819. #plt.legend(frameon=False)
  1820. #plt.show()
  1821. return regResults
  1822. plt.clf()
  1823. fig = plt.figure(figsize=(12,4))
  1824. plt.subplots_adjust(left=0.05, right=0.96, top=0.94, bottom=0.1)
  1825. cols = ['C0','C1','C2','C3','C4']
  1826. pawID = ['FL','FR','HL','HR']
  1827. ax = fig.add_subplot(1,5,1)
  1828. for i in range(len(regs)):
  1829. ax.plot(regResults[i]['coefficients'][:9],'o-',label=regs[i][0])
  1830. ax.set_ylabel('beta-weight')
  1831. plt.xticks(np.arange(9),['wheel speed','x-pos FL','x-pos FR','x-pos HL','x-pos HR','v FL','v FR','v HL','v HR'],rotation=45, ha='right',fontsize=8)
  1832. #plt.setp(ax.get_xticklabels(), rotation=45, ha="right",rotation_mode="anchor")
  1833. plt.legend(frameon=False)
  1834. tVector = np.linspace(-nShift,nShift,2*nShift+1,endpoint=True)*dt
  1835. shifts = 2*nShift + 1
  1836. for j in range(4):
  1837. ax = fig.add_subplot(1,5,j+2)
  1838. for i in range(len(regs)):
  1839. ax.set_title(pawID[j])
  1840. ax.plot(tVector,regResults[i]['coefficients'][(9+(2*j)*shifts):(9+(2*j+1)*shifts)],c=cols[i],ls=':',label=(None if j<3 else 'swingStart '+regs[i][0]))
  1841. ax.plot(tVector,regResults[i]['coefficients'][(9+(2*j+1)*shifts):(9+(2*j+2)*shifts)],c=cols[i],ls='-',label=(None if j<3 else 'swingEnd '+regs[i][0]))
  1842. #ax.set_ylim(-0.7,1.6)
  1843. ax.axvline(x=0,ls=':',c='0.4')
  1844. ax.set_xlabel('time (s)')
  1845. ax.set_ylabel('beta-weight')
  1846. plt.legend(frameon=False)
  1847. plt.show()
  1848. pdb.set_trace()
  1849. # perform linear regression : no regularization
  1850. print('Linear regression')
  1851. Lreg = linear_model.LinearRegression()
  1852. Lreg.fit(Xregressors_scaled, YspikeCount)
  1853. print(Lreg.coef_)
  1854. print(Lreg.intercept_)
  1855. print('score:',Lreg.score(Xregressors_scaled,YspikeCount))
  1856. Ypred = Lreg.predict(Xregressors_scaled)
  1857. plt.title('Linear reg.')
  1858. plt.plot(YspikeCount)
  1859. plt.plot(Ypred)
  1860. plt.show()
  1861. def constructDesignMatrix(pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
  1862. Xregressors = np.zeros((len(tbinCenters), 1))
  1863. # interpolate wheel speed
  1864. interpWheel = interp1d(sTimes, linearSpeed, fill_value='extrapolate') # ,kind='cubic')
  1865. newWheelSpeed = interpWheel(tbinCenters)
  1866. if (shuffle) and ('wheelSpeed' in variable):
  1867. newWheelSpeed = shuffleVariable(newWheelSpeed,tbinCenters,dt,chunkSize)
  1868. Xregressors[:, 0] = newWheelSpeed
  1869. # interpolate paw position and paw speed
  1870. for i in range(4):
  1871. interpPawPos = interp1d(pawPos[i][:, 0], pawPos[i][:, 1], fill_value='extrapolate')
  1872. # Xregressors[:,1+i] = interpPawPos(tbinCenters)
  1873. newPawPos = interpPawPos(tbinCenters)
  1874. if (shuffle) and (i in pawID) and ('pawPosition' in variable):
  1875. newPawPos = shuffleVariable(newPawPos,tbinCenters,dt,chunkSize)
  1876. Xregressors = np.column_stack((Xregressors, newPawPos))
  1877. for i in range(4):
  1878. interpPawSpeed = interp1d(pawSpeed[i][:, 0], pawSpeed[i][:, 2], fill_value='extrapolate') # pawSpeed[i][1] is total speed, 2 is x, 3 is y speed
  1879. # Xregressors[:,5+i] = interpPawSpeed(tbinCenters)
  1880. newPawSpeed = interpPawSpeed(tbinCenters)
  1881. if (shuffle) and (i in pawID) and ('pawSpeed' in variable):
  1882. newPawSpeed = shuffleVariable(newPawSpeed, tbinCenters, dt, chunkSize)
  1883. Xregressors = np.column_stack((Xregressors, newPawSpeed))
  1884. for i in range(4):
  1885. idxSwings = swingStanceD['swingP'][i][1]
  1886. recTimes = swingStanceD['forFit'][i][2]
  1887. idxSwings = np.asarray(idxSwings)
  1888. binnedSwingStartTimes, _ = np.histogram(recTimes[idxSwings[:, 0]], tbins)
  1889. binnedSwingEndTimes, _ = np.histogram(recTimes[idxSwings[:, 1]], tbins)
  1890. if (shuffle) and (i in pawID) and ('swingStart' in variable):
  1891. binnedSwingStartTimes = shuffleVariable(binnedSwingStartTimes, tbinCenters, dt, chunkSize)
  1892. if (shuffle) and (i in pawID) and ('stanceStart' in variable):
  1893. binnedSwingEndTimes = shuffleVariable(binnedSwingEndTimes, tbinCenters, dt, chunkSize)
  1894. startTshift = np.zeros((len(binnedSwingStartTimes), 2 * nShift + 1))
  1895. endTshift = np.zeros((len(binnedSwingEndTimes), 2 * nShift + 1))
  1896. n = 0
  1897. for j in range(-nShift, nShift + 1):
  1898. startTshift[:, n] = np.roll(binnedSwingStartTimes, j)
  1899. endTshift[:, n] = np.roll(binnedSwingEndTimes, j)
  1900. n += 1
  1901. Xregressors = np.column_stack((Xregressors, startTshift))
  1902. Xregressors = np.column_stack((Xregressors, endTshift)) # Xregressors[:,9+i] = binnedSwingStartTimes # Xregressors[:,13+i] = binnedSwingEndTimes
  1903. return Xregressors
  1904. matplotlib.use('TkAgg')
  1905. # create spike-count vector
  1906. spikeTimes = ephys[0]
  1907. print('firing rate :', 1./np.mean(np.diff(spikeTimes)))
  1908. dt = 0.01
  1909. shiftRange = 0.2 # binary columns are shifted back- and forth in time by this delay in s
  1910. nShift = int(shiftRange/dt)
  1911. tbins = np.linspace(0., 60., int(60 / dt) + 1, endpoint=True)
  1912. tbinCenters = (tbins[1:]+tbins[:-1])/2
  1913. binnedspikes, _ = np.histogram(spikeTimes, tbins)
  1914. spikecountwindow = 0.03 # sigma of the Gaussian kernel
  1915. nspikecountwindow = int(spikecountwindow / dt + 0.5)
  1916. binnedspikes = np.array(binnedspikes,dtype=float)
  1917. #YspikeCount = scipy.ndimage.gaussian_filter1d(binnedspikes, nspikecountwindow,axis=0) # convolve with Gaussian kernel
  1918. YspikeCount = np.convolve(binnedspikes, np.ones(nspikecountwindow), 'same') # convolve spike-count with square kernel
  1919. #print(len(YspikeCount),len(binnedspikes),nspikecountwindow)
  1920. #plt.hist(YspikeCount,bins=30)
  1921. #plt.show()
  1922. #pdb.set_trace()
  1923. # create regressor matrix
  1924. #Xregressors = np.zeros((len(tbinCenters),1+4*4))
  1925. timeLimits = [10,50]
  1926. tmask = (tbinCenters>timeLimits[0]) & (tbinCenters<timeLimits[1])
  1927. regs = [('Ridge Regression', linear_model.Ridge(alpha=1))]
  1928. Xregressors = constructDesignMatrix()
  1929. #pdb.set_trace()
  1930. print('shape or regressor matrix :', np.shape(Xregressors))
  1931. ## preprocessing of the data
  1932. scaler = preprocessing.MinMaxScaler().fit(Xregressors) # This estimator scales and translates each feature individually such that the maximal absolute value of each feature in the training set will be 1.0. It does not shift/center the data, and thus does not destroy any sparsity.
  1933. Xregressors_scaled = scaler.transform(Xregressors)
  1934. regResultsFullModel = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=10, visualize=False)
  1935. coffs = regResultsFullModel[0]['coefficients']
  1936. shuffleForWeights = False
  1937. if shuffleForWeights:
  1938. shuffCoeffsWheelSpeed = []
  1939. for i in range(100):
  1940. # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
  1941. Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['wheelSpeed'],chunkSize=2.,shuffle=True)
  1942. # pdb.set_trace()
  1943. #print('shape or regressor matrix :', np.shape(Xregressors))
  1944. ## preprocessing of the data
  1945. scaler = preprocessing.MinMaxScaler().fit(Xregressors) # This estimator scales and translates each feature individually such that the maximal absolute value of each feature in the training set will be 1.0. It does not shift/center the data, and thus does not destroy any sparsity.
  1946. Xregressors_scaled = scaler.transform(Xregressors)
  1947. regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
  1948. shuffCoeffsWheelSpeed.append(regResultsShuffle[0]['coefficients'])
  1949. shuffCoeffsPawPos = []
  1950. for i in range(100):
  1951. # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
  1952. Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['pawPosition'],chunkSize=2.,shuffle=True)
  1953. # pdb.set_trace()
  1954. #print('shape or regressor matrix :', np.shape(Xregressors))
  1955. ## preprocessing of the data
  1956. scaler = preprocessing.MinMaxScaler().fit(Xregressors) # This estimator scales and translates each feature individually such that the maximal absolute value of each feature in the training set will be 1.0. It does not shift/center the data, and thus does not destroy any sparsity.
  1957. Xregressors_scaled = scaler.transform(Xregressors)
  1958. regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
  1959. shuffCoeffsPawPos.append(regResultsShuffle[0]['coefficients'])
  1960. #
  1961. shuffCoeffsPawSpeed = []
  1962. for i in range(100):
  1963. # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
  1964. Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['pawSpeed'],chunkSize=2.,shuffle=True)
  1965. # pdb.set_trace()
  1966. #print('shape or regressor matrix :', np.shape(Xregressors))
  1967. ## preprocessing of the data
  1968. scaler = preprocessing.MinMaxScaler().fit(Xregressors) # This estimator scales and translates each feature individually such that the maximal absolute value of each feature in the training set will be 1.0. It does not shift/center the data, and thus does not destroy any sparsity.
  1969. Xregressors_scaled = scaler.transform(Xregressors)
  1970. regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
  1971. shuffCoeffsPawSpeed.append(regResultsShuffle[0]['coefficients'])
  1972. shuffCoeffsSwingStance = []
  1973. for i in range(100):
  1974. # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
  1975. Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['swingStart','stanceStart'],chunkSize=2.,shuffle=True)
  1976. # pdb.set_trace()
  1977. #print('shape or regressor matrix :', np.shape(Xregressors))
  1978. ## preprocessing of the data
  1979. scaler = preprocessing.MinMaxScaler().fit(Xregressors) # This estimator scales and translates each feature individually such that the maximal absolute value of each feature in the training set will be 1.0. It does not shift/center the data, and thus does not destroy any sparsity.
  1980. Xregressors_scaled = scaler.transform(Xregressors)
  1981. regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
  1982. shuffCoeffsSwingStance.append(regResultsShuffle[0]['coefficients'])
  1983. coffs = np.asarray(coffs)
  1984. shuffCoeffsWheelSpeed = np.asarray(shuffCoeffsWheelSpeed)
  1985. shuffCoeffsPawPos = np.asarray(shuffCoeffsPawPos)
  1986. shuffCoeffsPawSpeed = np.asarray(shuffCoeffsPawSpeed)
  1987. shuffCoeffsSwingStance = np.asarray(shuffCoeffsSwingStance)
  1988. shuffCoeffs = np.copy(shuffCoeffsWheelSpeed)
  1989. shuffCoeffs[:,1:5] = shuffCoeffsPawPos[:,1:5]
  1990. shuffCoeffs[:,5:10] = shuffCoeffsPawSpeed[:,5:10]
  1991. shuffCoeffs[:, 10:] = shuffCoeffsSwingStance[:,10:]
  1992. #pdb.set_trace()
  1993. fig = plt.figure(figsize=(12,4))
  1994. plt.subplots_adjust(left=0.05, right=0.96, top=0.94, bottom=0.1)
  1995. cols = ['C0','C1','C2','C3','C4']
  1996. #pawID = ['FL','FR','HL','HR']
  1997. ax0 = fig.add_subplot(1,5,1)
  1998. ax0.axhline(y=0,ls='--',c='0.5')
  1999. ax0.plot(coffs[:9],'o-')
  2000. ax0.plot(np.mean(shuffCoeffs,axis=0)[:9])
  2001. ax0.fill_between(np.arange(9),np.percentile(shuffCoeffs,5,axis=0)[:9],np.percentile(shuffCoeffs,95,axis=0)[:9],alpha=0.5)
  2002. #ax1 = fig.add_subplot(1,2,2)
  2003. tVector = np.linspace(-nShift,nShift,2*nShift+1,endpoint=True)*dt
  2004. shifts = 2*nShift + 1
  2005. for j in range(4):
  2006. ax1 = fig.add_subplot(1,5,j+2)
  2007. for i in range(len(regs)):
  2008. #ax.set_title(pawID[j])
  2009. ax1.plot(tVector,coffs[(9+(2*j)*shifts):(9+(2*j+1)*shifts)])
  2010. ax1.plot(tVector,np.mean(shuffCoeffs[:,(9+(2*j)*shifts):(9+(2*j+1)*shifts)],axis=0)) #label=(None if j<3 else 'swingStart '+regs[i][0]))
  2011. ax1.fill_between(tVector, np.percentile(shuffCoeffs[:,(9+(2*j)*shifts):(9+(2*j+1)*shifts)], 5, axis=0), np.percentile(shuffCoeffs[:,(9+(2*j)*shifts):(9+(2*j+1)*shifts)], 95, axis=0), alpha=0.5)
  2012. plt.show()
  2013. pdb.set_trace()
  2014. #Xregressors_scaled = np.copy(Xregressors)
  2015. #pdb.set_trace()
  2016. #Xregressors_scaled[:8] = Xregressors[:8] # preserve the sparse data
  2017. # generate list of regression models
  2018. # alpha multiplies the penalty terms : for alpha=0 is equivalent to an ordinary least square
  2019. # ridge regression : l2 regularization
  2020. # elastic net : The ElasticNet mixing parameter, with 0 <= l1_ratio <= 1. For l1_ratio = 0 the penalty is an L2 penalty. For l1_ratio = 1 it is an L1 penalty. For 0 < l1_ratio < 1, the penalty is a combination of L1 and L2.
  2021. regs = [('Linear Regression',linear_model.LinearRegression()),('Ridge regression',linear_model.Ridge(alpha=1.)),('Elastic Net Regression',linear_model.ElasticNet(alpha=0.01,l1_ratio=0.5,random_state=0)),('GLM with log link function',linear_model.PoissonRegressor(alpha=1e-6/len(YspikeCount[tmask])))]
  2022. #regs = [('Ridge Regression',linear_model.Ridge(alpha=1.))]#,('GLM with log link function',linear_model.PoissonRegressor(alpha=1e-6/len(YspikeCount[tmask])))]
  2023. # ('Ridge regression', linear_model.Ridge(alpha=1.)),
  2024. # ('Elastic Net Regression', linear_model.ElasticNet(alpha=0.01, random_state=0))]
  2025. RRidge = []
  2026. #import statsmodels.api as sm
  2027. #gamma_model = sm.GLM(YspikeCount[tmask],Xregressors_scaled[tmask], family=sm.families.Poisson())
  2028. #gamma_results = gamma_model.fit()
  2029. #print(gamma_results.summary())
  2030. #pdb.set_trace()
  2031. checkAlphaVariable = False
  2032. if checkAlphaVariable :
  2033. for i in range(100):
  2034. al = float(i) #1./(1.1220184543019633**i)
  2035. #regs = [('GLM with log link function',linear_model.PoissonRegressor(alpha=al))]
  2036. regs = [('Ridge Regression', linear_model.Ridge(alpha=al))]
  2037. regResults = crossValidatedRegression(regs,Xregressors_scaled[tmask],YspikeCount[tmask],tbinCenters[tmask],fold=10,visualize=False)
  2038. RRidge.append([i,al,regResults[0]['fitScore'],np.mean(regResults[0]['scores'])])
  2039. RRidge = np.asarray(RRidge)
  2040. plt.title('Ridge Regression with log-link function')
  2041. plt.plot(RRidge[:,1], RRidge[:,2], 'o-',label='fit score')
  2042. plt.plot(RRidge[:,1],RRidge[:,3],'o-',label='cross validated score')
  2043. plt.xlabel('alpha (penalty weight)')
  2044. plt.ylabel('R^2')
  2045. #plt.xscale('log')
  2046. #plt.legend(frameon=False)
  2047. plt.show()
  2048. pdb.set_trace()
  2049. print(len(Xregressors_scaled[tmask]))
  2050. #plt.plot(tbinCenters[tmask],YspikeCount[tmask])
  2051. #coeffss = []
  2052. #pdb.set_trace()
  2053. #plt.legend(frameon=False)
  2054. #plt.show()
  2055. regs = [('Linear Regression',linear_model.LinearRegression()),('Ridge regression',linear_model.Ridge(alpha=1.)),('Elastic Net Regression',linear_model.ElasticNet(alpha=0.01,l1_ratio=0.5,random_state=0)),('GLM with log link function',linear_model.PoissonRegressor(alpha=1e-6/len(YspikeCount[tmask])))]
  2056. #regResultsFullModel = {}
  2057. #for i in range(len(regs)):
  2058. regResultsFullModel = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
  2059. plt.clf()
  2060. fig = plt.figure(figsize=(12,4))
  2061. plt.subplots_adjust(left=0.05, right=0.96, top=0.94, bottom=0.1)
  2062. cols = ['C0','C1','C2','C3','C4']
  2063. pawID = ['FL','FR','HL','HR']
  2064. ax = fig.add_subplot(1,5,1)
  2065. for i in range(len(regs)):
  2066. ax.plot(regResultsFullModel[i]['coefficients'][:9],'o-',label=regs[i][0])
  2067. ax.set_ylabel('beta-weight')
  2068. plt.xticks(np.arange(9),['wheel speed','x-pos FL','x-pos FR','x-pos HL','x-pos HR','v FL','v FR','v HL','v HR'],rotation=45, ha='right',fontsize=8)
  2069. #plt.setp(ax.get_xticklabels(), rotation=45, ha="right",rotation_mode="anchor")
  2070. plt.legend(frameon=False)
  2071. tVector = np.linspace(-nShift,nShift,2*nShift+1,endpoint=True)*dt
  2072. shifts = 2*nShift + 1
  2073. for j in range(4):
  2074. ax = fig.add_subplot(1,5,j+2)
  2075. for i in range(len(regs)):
  2076. ax.set_title(pawID[j])
  2077. ax.plot(tVector,regResultsFullModel[i]['coefficients'][(9+(2*j)*shifts):(9+(2*j+1)*shifts)],c=cols[i],ls='-',label=(None if j<3 else 'swingStart '+regs[i][0]))
  2078. ax.plot(tVector,regResultsFullModel[i]['coefficients'][(9+(2*j+1)*shifts):(9+(2*j+2)*shifts)],c=cols[i],ls=':',label=(None if j<3 else 'swingEnd '+regs[i][0]))
  2079. ax.set_ylim(-0.7,1.6)
  2080. ax.axvline(x=0,ls=':',c='0.4')
  2081. ax.set_xlabel('time (s)')
  2082. ax.set_ylabel('beta-weight')
  2083. plt.legend(frameon=False)
  2084. plt.show()
  2085. pdb.set_trace()
  2086. # perform linear regression : no regularization
  2087. print('Linear regression')
  2088. Lreg = linear_model.LinearRegression()
  2089. Lreg.fit(Xregressors_scaled, YspikeCount)
  2090. print(Lreg.coef_)
  2091. print(Lreg.intercept_)
  2092. print('score:',Lreg.score(Xregressors_scaled,YspikeCount))
  2093. Ypred = Lreg.predict(Xregressors_scaled)
  2094. plt.title('Linear reg.')
  2095. plt.plot(YspikeCount)
  2096. plt.plot(Ypred)
  2097. plt.show()
  2098. # perform elastic net regression : Linear regression with combined L1 and L2 priors as regularizer
  2099. # Elastic-net is useful when there are multiple features which are correlated with one another. Lasso is likely to pick one of these at random, while elastic-net is likely to pick both.
  2100. print('Elastic net regression')
  2101. ENreg = linear_model.ElasticNet(alpha=0.1,random_state=0)
  2102. ENreg.fit(Xregressors_scaled, YspikeCount)
  2103. print(ENreg.coef_)
  2104. print(ENreg.intercept_)
  2105. print('score:',ENreg.score(Xregressors_scaled,YspikeCount))
  2106. Ypred = ENreg.predict(Xregressors_scaled)
  2107. plt.title('Elastic net reg.')
  2108. plt.plot(YspikeCount)
  2109. plt.plot(Ypred)
  2110. plt.show()
  2111. #pdb.set_trace()
  2112. # perform Generalized Linear Model with a Poisson distribution. This regressor uses the ‘log’ link function.
  2113. # from sklearn.ensemble import HistGradientBoostingRegressor
  2114. print('GLM with log link function')
  2115. #Preg = HistGradientBoostingRegressor(loss="poisson",l2_regularization=1, max_leaf_nodes=128) #
  2116. Preg = linear_model.PoissonRegressor(alpha=0)
  2117. Preg.fit(Xregressors_scaled, YspikeCount)
  2118. print(Preg.coef_)
  2119. print(Preg.intercept_)
  2120. print('score:',Preg.score(Xregressors_scaled, YspikeCount))
  2121. Ypred = Preg.predict(Xregressors_scaled)
  2122. plt.title('GLM with log link function')
  2123. plt.plot(YspikeCount)
  2124. plt.plot(Ypred)
  2125. plt.show()
  2126. pdb.set_trace()
  2127. #################################################################################
  2128. # calculate correlations between ca-imaging, wheel speed and paw speed
  2129. #################################################################################
  2130. def doContinuousRegressionAnalysis(mouse,allCorrDataPerSession,allStepData,borders=None,figShow=False):
  2131. matplotlib.use('TkAgg') # WxAgg
  2132. from sklearn.linear_model import LinearRegression
  2133. from sklearn.model_selection import train_test_split
  2134. from sklearn.svm import SVR
  2135. #SVR(kernel='rbf', C=1e3, gamma=0.1)
  2136. #from sklearn.ensemble import RandomForestRegressor
  2137. regressionN = 6
  2138. Rvalues = []
  2139. #for nSess in range(len(allCorrDataPerSession)):
  2140. for nDay in range(len(allCorrDataPerSession)):
  2141. print(nDay,allCorrDataPerSession[nDay]['folder'])
  2142. (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP, wheelSpeedDict, pawTracksDict, caTracesDict, slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay,allCorrDataPerSession,allStepData)
  2143. # ATTENTION : all of the arrays also contain a time array
  2144. # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTimeSame]
  2145. # dims of pawTracksDictInterP : [nSessions][nPaw][2][valuesOverTimeSame]
  2146. # dims of caTracesDictInterP : [nSessions][nRois+1][valuesOverTimeSame]
  2147. # dims of wheelSpeedDict : [nSessions][2][valuesOverTime]
  2148. # dims of pawTracksDict : [nSessions][nPaw][2][valuesOverTime]
  2149. # dims of caTracesDict : [nSessions][nRois+1][valuesOverTime]
  2150. #(wheelSpeedDict, pawTracksDict, caTracesDict,aa,bb,cc) = getCaWheelPawInterpolatedDictsPerDay(nSess, allCorrDataPerSession)
  2151. nRecWheel = len(wheelSpeedDictInterP)
  2152. nRecPaw = len(pawTracksDictInterP)
  2153. nRecCa = len(caTracesDictInterP)
  2154. print('Recording length :', nRecWheel, nRecPaw, nRecCa)
  2155. if (nRecWheel != nRecPaw ) or (nRecWheel != nRecCa):
  2156. print('problem in number of recordings listed in dictionaries')
  2157. # loop over 5 different regressions, each using a different combination of test and train samples
  2158. recs = range(nRecCa)
  2159. RTempValues = []
  2160. #for reg in range(nRecCa): # loop over all recordings
  2161. #recsForTraining = recs.copy()
  2162. #recsForTraining.remove(reg)
  2163. #recsForTest = [reg]
  2164. varr = ['wheel speed','paw speed 0', 'paw speed 1','paw speed 2','paw speed 3','paw speed 0+1+2+3']
  2165. for d in range(regressionN): # loop over wheel speed, the four paw speeds and the combined speed
  2166. print('regressing %s' %varr[d])
  2167. #pdb.set_trace()
  2168. # concatenate data
  2169. if borders is not None:
  2170. timeMaskTrain = (caTracesDictInterP[0][0]>=borders[0])&(caTracesDictInterP[0][0]<=borders[1])
  2171. timeMaskTest = (caTracesDictInterP[0][0] >= borders[0]) & (caTracesDictInterP[0][0] <= borders[1])
  2172. else:
  2173. timeMaskTrain = (caTracesDictInterP[0][0]>=0)&(caTracesDictInterP[0][0]<=1000.)
  2174. timeMaskTest = (caTracesDictInterP[0][0]>=0)&(caTracesDictInterP[0][0]<=1000.)
  2175. X = np.copy(caTracesDictInterP[0][1:][:,timeMaskTrain])
  2176. #Xtest = np.copy(caTracesDictInterP[0][1:][:,timeMaskTest])
  2177. #pdb.set_trace()
  2178. if d == 0:
  2179. Y = np.copy(wheelSpeedDictInterP[0][1:][:,timeMaskTrain])
  2180. #Ytest = np.copy(wheelSpeedDictInterP[0][1:][:,timeMaskTest])
  2181. #YtestTime = np.copy(wheelSpeedDictInterP[0][0][timeMaskTest])
  2182. elif (d>0) and (d<5):
  2183. pawId = d-1
  2184. Y = np.copy(pawTracksDictInterP[0][pawId][1:][:,timeMaskTrain])
  2185. #Ytest = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][1:][:,timeMaskTest])
  2186. #YtestTime = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][0][timeMaskTest])
  2187. elif d==5: # case where all four paw speeds are added together
  2188. pawSpeedTrain = []
  2189. #pawSpeedTest = []
  2190. for i in range(4):
  2191. pawSpeedTrain.append(np.copy(pawTracksDictInterP[0][i][1:][:, timeMaskTrain]))
  2192. #pawSpeedTest.append(np.copy(pawTracksDictInterP[recsForTest[0]][i][1:][:, timeMaskTest]))
  2193. Y = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
  2194. #Ytest = pawSpeedTest[0] + pawSpeedTest[1] + pawSpeedTest[2] + pawSpeedTest[3]
  2195. #pdb.set_trace()
  2196. for t in recs[1:]:
  2197. if borders is not None:
  2198. timeMaskTrain = (caTracesDictInterP[t][0] >= borders[0]) & (caTracesDictInterP[t][0] <= borders[1])
  2199. else:
  2200. timeMaskTrain = (caTracesDictInterP[t][0] >= 0) & (caTracesDictInterP[t][0] <= 1000.)
  2201. X = np.column_stack((X,caTracesDictInterP[t][1:][:,timeMaskTrain]))
  2202. if d == 0:
  2203. Y = np.column_stack((Y,wheelSpeedDictInterP[t][1:][:,timeMaskTrain]))
  2204. elif (d>0) and (d<5):
  2205. Y = np.column_stack((Y, pawTracksDictInterP[t][pawId][1:][:,timeMaskTrain]))
  2206. elif d==5:
  2207. pawSpeedTrain = []
  2208. for i in range(4):
  2209. pawSpeedTrain.append(np.copy(pawTracksDictInterP[t][i][1:][:, timeMaskTrain]))
  2210. speedTemp = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
  2211. Y = np.column_stack((Y, speedTemp))
  2212. #pdb.set_trace()
  2213. Y = Y[0]
  2214. X = np.transpose(X)
  2215. #Ytest = Ytest[0]
  2216. nRegressionIterations = 10
  2217. Rval1 = np.zeros(2)
  2218. for n in range(nRegressionIterations):
  2219. #print(n,end='')
  2220. X_train, X_test, y_train, y_test = train_test_split(X, Y, test_size = 0.2)
  2221. #Xtest = np.transpose(Xtest)
  2222. # linear regression ########################################
  2223. linReg = LinearRegression()
  2224. linReg.fit(X_train,y_train)
  2225. #svm_rbf = SVR(kernel='rbf', C=1e3, gamma=0.1)
  2226. #svm_rbf.fit(X,Y)
  2227. #YTrainPred = linReg.predict(X)
  2228. y_test_pred = linReg.predict(X_test)
  2229. R2trainLR = linReg.score(X_train, y_train)
  2230. R2testLR = linReg.score(X_test, y_test) # 1. - np.sum((Ytest-YTestPred)**2)/np.sum((Ytest - np.mean(Ytest))**2)#linReg.score(Xtest, Ytest)
  2231. #print(linReg.coef_)
  2232. #print(linReg.intercept_)
  2233. #yPred = linReg.predict(np.transpose(X))
  2234. # random forest ##############################################
  2235. #randForestReg = RandomForestRegressor(n_estimators=20)
  2236. #randForestReg.fit(X, Y)
  2237. #R2trainRF = randForestReg.score(X, Y)
  2238. #R2testRF= randForestReg.score(Xtest, Ytest)
  2239. #
  2240. if figShow :
  2241. print('R2 test :',R2testLR)
  2242. fig = plt.figure()
  2243. ax = fig.add_subplot(111)
  2244. ax.plot(y_test,lw=2)
  2245. ax.plot(y_test_pred,lw=2)
  2246. ax.spines['top'].set_visible(False)
  2247. ax.spines['right'].set_visible(False)
  2248. ax.spines['bottom'].set_position(('outward', 10))
  2249. ax.spines['left'].set_position(('outward', 10))
  2250. ax.yaxis.set_ticks_position('left')
  2251. ax.xaxis.set_ticks_position('bottom')
  2252. plt.show()
  2253. Rval1+= np.array([R2trainLR,R2testLR])
  2254. RTempValues.extend(Rval1/nRegressionIterations)
  2255. #pdb.set_trace()
  2256. #Rs = np.zeros(regressionN*2)
  2257. #for reg in range(nRegressionIterations):
  2258. # Rs += RTempValues[reg]
  2259. #Rs /=nRegressionIterations
  2260. Rvalues.append([nDay,allCorrDataPerSession[nDay]['folder'],RTempValues])
  2261. return Rvalues
  2262. #################################################################################
  2263. # calculate correlations between ca-imaging, wheel speed and paw speed
  2264. #################################################################################
  2265. def generateStepTriggeredCaTraces(mouse,allCorrDataPerSession,allStepData,trigger='swingOnset',calculateFast=False): # swingOnset or swingOffset
  2266. matplotlib.use('TkAgg')
  2267. # check for sanity
  2268. if len(allCorrDataPerSession) != len(allStepData):
  2269. print('both dictionaries are not of the same length')
  2270. print('CaWheelPawDict:',len(allCorrDataPerSession),' StepStanceDict:,',len(allStepData))
  2271. timeAxis = np.linspace(-0.4,0.6,int((0.4+0.6)/0.02)+1)
  2272. timeAxisRescaled = np.linspace(-1.,2.,int((1.+2.)/0.02)+1)
  2273. preStanceMask = timeAxis<-0.1
  2274. preStanceRescaledMask = timeAxisRescaled<-0.2
  2275. K = len(timeAxis)
  2276. KRescaled = len(timeAxisRescaled)
  2277. caTraces = []
  2278. maxTimeDelay = 1.
  2279. for nDay in range(len(allCorrDataPerSession)):
  2280. print(allCorrDataPerSession[nDay]['folder'], allStepData[nDay][0], nDay)
  2281. # consistency check
  2282. if not (allCorrDataPerSession[nDay]['folder'] == allStepData[nDay][0]):
  2283. print('All animal data and swing data not from the same day!')
  2284. #
  2285. caRecTime = allCorrDataPerSession[nDay]['caImg']['timeStamps'][0, 3]
  2286. wheelRecTime = allStepData[nDay][1][0][3] # check recording start of first recording allStepData[nDay][1][3]
  2287. pawRecTime = allStepData[nDay][2][0][4] # again, third indices picks first recording
  2288. #pdb.set_trace()
  2289. timeDiffCaWheel = np.abs(caRecTime-wheelRecTime)
  2290. timeDiffCaPaw = np.abs(caRecTime-pawRecTime)
  2291. timeDiffWheelPaw = np.abs(wheelRecTime-pawRecTime)
  2292. if any(np.array([timeDiffCaWheel,timeDiffCaPaw,timeDiffWheelPaw]) > maxTimeDelay):
  2293. print('PROBLEM in data consistency!')
  2294. print('recordings are separated by %s - %s - %s s min' % (timeDiffCaWheel,timeDiffCaPaw,timeDiffWheelPaw))
  2295. pdb.set_trace()
  2296. else:
  2297. print('Delay between recordings is :', timeDiffCaWheel,timeDiffCaPaw,timeDiffWheelPaw, 's')
  2298. #print(allCorrDataPerSession[nDay][0],nDay) getCaWheelPawInterpolatedDictsPerDay
  2299. (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP,wheelSpeedDict,pawTracksDict,caTracesDict,slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession,allStepData)
  2300. #if len(allStepData[nDay-1][4])==6:
  2301. # print('more recordings :',len(allStepData[nDay][4]))
  2302. # addIdx = 1
  2303. #else:
  2304. # addIdx = 0
  2305. #pdb.set_trace()
  2306. N = len(caTracesDict[0][1:]) # number of ROIs
  2307. caSnippets = [[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)]] #np.zeros((4,N,K))
  2308. caSnippetsRescaled = [[[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)],[[] for i in range(N)]]
  2309. recordingID = [[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)]]
  2310. recordingIDRescaled = [[[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)]]
  2311. NRecs=len(allStepData[nDay][4])
  2312. for nrec in range(NRecs): # loop over the five recordings of a day
  2313. for i in range(4): # loop over the four paws
  2314. #pdb.set_trace()
  2315. idxSwings = allStepData[nDay][4][nrec][3][i][1]
  2316. #print('Wow nSess',nDay,allStepData[nDay-1][0])
  2317. recTimes = allStepData[nDay][4][nrec][4][i][2]
  2318. #pdb.set_trace()
  2319. idxSwings = np.asarray(idxSwings)
  2320. if trigger == 'swingOnset':
  2321. NstepCycles = len(idxSwings)
  2322. elif trigger == 'swingOffset':
  2323. NstepCyles = len(idxSwings)-1
  2324. for k in range(NstepCyles): # loop over all swings
  2325. startSwingTime = recTimes[idxSwings[k, 0]]
  2326. endSwingTime = recTimes[idxSwings[k, 1]]
  2327. if trigger == 'swingOnset':
  2328. triggerTime = startSwingTime
  2329. duration = (endSwingTime-startSwingTime)
  2330. elif trigger == 'swingOffset':
  2331. triggerTime = endSwingTime
  2332. duration = recTimes[idxSwings[k+1, 0]] - endSwingTime
  2333. if len(caTracesDict[nrec][1:])!=N:
  2334. print('problem in number of ROIs')
  2335. pdb.set_trace(0)
  2336. for l in range(len(caTracesDict[nrec][1:])): # loop over all ROIs
  2337. interpCa = interp1d(caTracesDict[nrec][0]-triggerTime, caTracesDict[nrec][l+1])#,kind='cubic')
  2338. interpCaRescaled = interp1d((caTracesDict[nrec][0]-triggerTime)/(duration), caTracesDict[nrec][l+1])#,kind='cubic')
  2339. ############
  2340. try:
  2341. newCaTraceAtSwing = interpCa(timeAxis)
  2342. except ValueError:
  2343. pass
  2344. else:
  2345. #caSnippets[i,l,:] += newCaTraceAtSwing
  2346. caSnippets[i][l].append(newCaTraceAtSwing)
  2347. recordingID[i][l].append(nrec)
  2348. ############
  2349. try:
  2350. newCaTraceAtSwingRescaled = interpCaRescaled(timeAxisRescaled)
  2351. except ValueError:
  2352. #print('error')
  2353. pass
  2354. else:
  2355. #caSnippets[i,l,:] += newCaTraceAtSwing
  2356. caSnippetsRescaled[i][l].append(newCaTraceAtSwingRescaled)
  2357. recordingIDRescaled[i][l].append(nrec)
  2358. #pdb.set_trace()
  2359. # 4 paws
  2360. # N number of ROIS
  2361. # 2 mean and std
  2362. # K number of time points during average
  2363. alpha = 0.05
  2364. def get_CI(dat, alpha=.05):
  2365. return np.array(list(map(lambda x: boot.ci(x, alpha=alpha), dat)))
  2366. caSnippetsArray = np.zeros((4, N, 3 + 3*NRecs, K))
  2367. caSnippetsRescaledArray = np.zeros((4, N, 3 + 3*NRecs, KRescaled))
  2368. for i in range(4): # loop over four paws
  2369. for l in range(N): # loop over all ROIs
  2370. caTempArray = np.asarray(caSnippets[i][l])
  2371. #pdb.set_trace()
  2372. caSnippetsZscores = (caTempArray - np.mean(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]) #/np.std(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]
  2373. if calculateFast:
  2374. CI = np.zeros((len(timeAxis),2))
  2375. else:
  2376. print('calculating CIs of global average ...')
  2377. CI = np.array(list(map(lambda x: boot.ci(x, alpha=.05), caSnippetsZscores.T)))
  2378. print('done')
  2379. caTemp = np.mean(caSnippetsZscores,axis=0)
  2380. #caTempSTD = np.std(caSnippetsZscores,axis=0)
  2381. caSnippetsArray[i,l,0,:] = caTemp
  2382. caSnippetsArray[i,l,1,:] = CI[:,0]
  2383. caSnippetsArray[i,l,2,:] = CI[:,1]
  2384. caTempRec = []
  2385. for nrec in range(NRecs):
  2386. recMask = (np.asarray(recordingID[i][l]) == nrec)
  2387. caTempRec.append(caSnippetsZscores[recMask])
  2388. if calculateFast:
  2389. ci_cell_trial = np.zeros((NRecs, len(timeAxis), 2))
  2390. else:
  2391. print('calculating CIs of recording average for paw %s and roi %s ...' % (i, l))
  2392. ci_cell_trial = np.array(Parallel(n_jobs=numcores)(delayed(get_CI)(c.T, alpha) for c in caTempRec))
  2393. print('done')
  2394. for nrec in range(NRecs):
  2395. caSnippetsArray[i, l, 3 + nrec*3, :] = np.mean(caTempRec[nrec], axis=0)
  2396. caSnippetsArray[i, l, 4 + nrec*3, :] = ci_cell_trial[nrec][:,0]
  2397. caSnippetsArray[i, l, 5 + nrec*3, :] = ci_cell_trial[nrec][:,1]
  2398. #pdb.set_trace()
  2399. caTempRescaledArray = np.asarray(caSnippetsRescaled[i][l])
  2400. #pdb.set_trace()
  2401. caSnippetsRescaledZscores = (caTempRescaledArray - np.mean(caTempRescaledArray[:,preStanceRescaledMask],axis=1)[:,np.newaxis])# /np.std(caTempRescaledArray[:,preStanceRescaledMask],axis=1)[:,np.newaxis]
  2402. caTempRe = np.mean(caSnippetsRescaledZscores,axis=0)
  2403. if calculateFast:
  2404. CI = np.zeros((len(timeAxisRescaled),2))
  2405. else:
  2406. print('calculating CIs of global rescaled recording average')
  2407. CI = np.array(list(map(lambda x: boot.ci(x, alpha=.05), caSnippetsRescaledZscores.T)))
  2408. print('done')
  2409. #caTempReSTD = np.std(caSnippetsRescaledZscores,axis=0)
  2410. caSnippetsRescaledArray[i,l,0,:] = caTempRe
  2411. caSnippetsRescaledArray[i,l,1,:] = CI[:,0]
  2412. caSnippetsRescaledArray[i,l,2,:] = CI[:,1]
  2413. caTempRecRescaled = []
  2414. for nrec in range(NRecs):
  2415. recMask = (np.asarray(recordingIDRescaled[i][l]) == nrec)
  2416. caTempRecRescaled.append(caSnippetsRescaledZscores[recMask])
  2417. #caTemp = np.mean(caSnippetsRescaledZscores[recMask], axis=0)
  2418. #caTempSTD = np.std(caSnippetsRescaledZscores[recMask], axis=0)
  2419. #CI = np.array(list(map(lambda x: boot.ci(x, alpha=.05), caSnippetsRescaledZscores[recMask].T)))
  2420. if calculateFast :
  2421. ci_cell_trial_rescaled = np.zeros((NRecs,len(timeAxisRescaled),2))
  2422. else:
  2423. print('calculating CIs of rescaled recording average for paw %s and roi %s ...' % (i, l))
  2424. ci_cell_trial_rescaled = np.array(Parallel(n_jobs=numcores)(delayed(get_CI)(c.T, alpha) for c in caTempRecRescaled))
  2425. print('done')
  2426. for nrec in range(NRecs):
  2427. caSnippetsRescaledArray[i, l, 3 + nrec*3, :] = np.mean(caTempRecRescaled[nrec], axis=0)
  2428. caSnippetsRescaledArray[i, l, 4 + nrec*3, :] = ci_cell_trial_rescaled[nrec][:,0]
  2429. caSnippetsRescaledArray[i, l, 5 + nrec*3, :] = ci_cell_trial_rescaled[nrec][:,1]
  2430. caTraces.append([allCorrDataPerSession[nDay]['folder'],allStepData[nDay][0],nDay,timeAxis,caSnippetsArray,timeAxisRescaled,caSnippetsRescaledArray])
  2431. return caTraces
  2432. #################################################################################
  2433. # calculate correlations between ca-imaging, wheel speed and paw speed
  2434. #################################################################################
  2435. def generateStepTriggeredCaTracesAllPaws(mouse,allCorrDataPerSession,allStepData):
  2436. maxSeparation = 0.# min separation between swings in sec
  2437. # check for sanity
  2438. if len(allCorrDataPerSession) != len(allStepData):
  2439. print('both dictionaries are not of the same length')
  2440. print('CaWheelPawDict:',len(allCorrDataPerSession),' StepStanceDict:,',len(allStepData))
  2441. timeAxis = np.linspace(-0.4,0.6,(0.6+0.4)/0.02+1)
  2442. timeAxisRescaled = np.linspace(-1.,2.,(2+1)/0.02+1)
  2443. preStanceMask = timeAxis<-0.1
  2444. preStanceRescaledMask = timeAxisRescaled<-0.2
  2445. K = len(timeAxis)
  2446. KRescaled = len(timeAxisRescaled)
  2447. caTraces = []
  2448. swingT = []
  2449. for nDay in range(len(allCorrDataPerSession)):
  2450. (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP,wheelSpeedDict,pawTracksDict,caTracesDict) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession)
  2451. if len(allStepData[nDay-1][4])==6:
  2452. print('more recordings :',len(allStepData[nDay-1][4]))
  2453. addIdx = 1
  2454. else:
  2455. addIdx = 0
  2456. #pdb.set_trace()
  2457. N = len(caTracesDict[0][1:])
  2458. print(allCorrDataPerSession[nDay][0], allStepData[nDay - 1][0], nDay, N)
  2459. caSnippets = [[] for i in range(N)] #np.zeros((4,N,K))
  2460. caSnippetsRescaled = [[] for i in range(N)]
  2461. caSnippetsArray = np.zeros((N,2,K))
  2462. caSnippetsRescaledArray = np.zeros((N, 2, KRescaled))
  2463. swingSnippets = [[] for i in range(5)]
  2464. for nrec in range(5): # loop over the five recordings of a day
  2465. swingTimes = np.zeros(3)
  2466. for i in range(4): # loop over the four paws and lump all swing times together
  2467. idxSwings = allStepData[nDay-1][4][nrec+addIdx][3][i][1]
  2468. #print('Wow nSess',nDay,allStepData[nDay-1][0])
  2469. recTimes = allStepData[nDay-1][4][nrec+addIdx][4][i][2]
  2470. #pdb.set_trace()
  2471. idxSwings = np.asarray(idxSwings)
  2472. startSwingTimes = recTimes[idxSwings[:,0]]
  2473. endSwingTimes = recTimes[idxSwings[:,1]]
  2474. swingTimes = np.vstack((swingTimes,np.column_stack((startSwingTimes, endSwingTimes,np.repeat(i,len(startSwingTimes))))))
  2475. swingTimes = swingTimes[1:] # remove first element which was zeros only
  2476. # sort swing times according to swing start
  2477. swingTimes = swingTimes[swingTimes[:,0].argsort()]
  2478. # remove swings with fall within the minimum separation betweeen swings
  2479. diffSwings = np.diff(swingTimes[:,0]) # calculate inter-swing intervals
  2480. swingTimesSparse = swingTimes[np.concatenate((diffSwings>maxSeparation,np.array([True])))] # only use swings which fall above separation time
  2481. swingSnippets[nrec] = swingTimes
  2482. for k in range(len(swingTimesSparse)): # loop over all swings
  2483. startSwingTime = swingTimesSparse[k, 0]
  2484. endSwingTime = swingTimesSparse[k, 1]
  2485. if len(caTracesDict[nrec][1:])!=N: print('problem in number of ROIs')
  2486. for l in range(len(caTracesDict[nrec][1:])): # loop over all ROIs
  2487. interpCa = interp1d(caTracesDict[nrec][0]-startSwingTime, caTracesDict[nrec][l+1])#,kind='cubic')
  2488. interpCaRescaled = interp1d((caTracesDict[nrec][0]-startSwingTime)/(endSwingTime-startSwingTime), caTracesDict[nrec][l+1])#,kind='cubic')
  2489. ############
  2490. try:
  2491. newCaTraceAtSwing = interpCa(timeAxis)
  2492. except ValueError:
  2493. pass
  2494. else:
  2495. #caSnippets[i,l,:] += newCaTraceAtSwing
  2496. caSnippets[l].append(newCaTraceAtSwing)
  2497. ############
  2498. try:
  2499. newCaTraceAtSwingRescaled = interpCaRescaled(timeAxisRescaled)
  2500. except ValueError:
  2501. #print('error')
  2502. pass
  2503. else:
  2504. #caSnippets[i,l,:] += newCaTraceAtSwing
  2505. caSnippetsRescaled[l].append(newCaTraceAtSwingRescaled)
  2506. #pdb.set_trace()
  2507. #for i in range(4):
  2508. for l in range(N):
  2509. caTempArray = np.asarray(caSnippets[l])
  2510. caSnippetsZscores = (caTempArray - np.mean(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]) #/np.std(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]
  2511. caTemp = np.mean(caSnippetsZscores,axis=0)
  2512. caTempSTD = np.std(caSnippetsZscores,axis=0)
  2513. caSnippetsArray[l,0,:] = caTemp
  2514. caSnippetsArray[l,1,:] = caTempSTD
  2515. #
  2516. caTempRescaledArray = np.asarray(caSnippetsRescaled[l])
  2517. #pdb.set_trace()
  2518. caSnippetsRescaledZscores = (caTempRescaledArray - np.mean(caTempRescaledArray[:,preStanceRescaledMask],axis=1)[:,np.newaxis]) #/np.std(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]
  2519. caTempRe = np.mean(caSnippetsRescaledZscores,axis=0)
  2520. caTempReSTD = np.std(caSnippetsRescaledZscores,axis=0)
  2521. caSnippetsRescaledArray[l,0,:] = caTempRe
  2522. caSnippetsRescaledArray[l,1,:] = caTempReSTD
  2523. swingT.append([allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay,swingSnippets])
  2524. caTraces.append([allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay,caSnippetsArray,caSnippetsRescaledArray])
  2525. return caTraces
  2526. #################################################################################
  2527. def calcualteAllDeltaTValues(t1,t2):
  2528. l1 = len(t1)
  2529. l2 = len(t2)
  2530. fannedOut1 = np.tile(t1,(l2,1))
  2531. fannedOut2 = np.tile(t2,(l1,1))
  2532. transposedFannedOut2 = np.transpose(fannedOut2)
  2533. #print l1, l2
  2534. #print shape(fannedOut1), shape(transposedFannedOut2)
  2535. #pdb.set_trace()
  2536. differences = fannedOut1 - transposedFannedOut2 # subtract(fannedOut1,transposedFannedOut2)
  2537. return differences.flatten()
  2538. #################################################################################
  2539. # calculate correlations between ca-imaging, wheel speed and paw speed
  2540. #################################################################################
  2541. def generateInterstepTimeHistogram(mouse,allCorrDataPerSession,allStepData):
  2542. # check for sanity
  2543. if len(allCorrDataPerSession) != len(allStepData):
  2544. print('both dictionaries are not of the same length')
  2545. print('CaWheelPawDict:',len(allCorrDataPerSession),' StepStanceDict:,',len(allStepData))
  2546. #timeAxis = np.linspace(-0.4,0.6,(0.6+0.4)/0.02+1)
  2547. #timeAxisRescaled = np.linspace(-1.,2.,(2+1)/0.02+1)
  2548. #preStanceMask = timeAxis<-0.1
  2549. #preStanceRescaledMask = timeAxisRescaled<-0.2
  2550. #K = len(timeAxis)
  2551. #KRescaled = len(timeAxisRescaled)
  2552. pawSwingTimes = []
  2553. for nDay in range(1,len(allCorrDataPerSession)):
  2554. print(allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay)
  2555. (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP,wheelSpeedDict,pawTracksDict,caTracesDict) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession)
  2556. if len(allStepData[nDay-1][4])==6:
  2557. print('more recordings :',len(allStepData[nDay-1][4]))
  2558. addIdx = 1
  2559. else:
  2560. addIdx = 0
  2561. #pdb.set_trace()
  2562. #N = len(caTracesDict[0][1:])
  2563. #caSnippets = [[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)]] #np.zeros((4,N,K))
  2564. #caSnippetsRescaled = [[[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)]]
  2565. #caSnippetsArray = np.zeros((4,N,2,K))
  2566. #caSnippetsRescaledArray = np.zeros((4, N, 2, KRescaled))
  2567. allData = []
  2568. for nrec in range(5): # loop over the five recordings of a day
  2569. startSwingTimes = [[] for i in range(4)]
  2570. endSwingTimes = [[] for i in range(4)]
  2571. for i in range(4): # loop over the four paws
  2572. idxSwings = allStepData[nDay-1][4][nrec+addIdx][3][i][1]
  2573. #print('Wow nSess',nDay,allStepData[nDay-1][0])
  2574. recTimes = allStepData[nDay-1][4][nrec+addIdx][4][i][2]
  2575. #pdb.set_trace()
  2576. idxSwings = np.asarray(idxSwings)
  2577. startSwingT = recTimes[idxSwings[:,0]]
  2578. endSwingT = recTimes[idxSwings[:,1]]
  2579. startSwingTimes[i].append(startSwingT)
  2580. endSwingTimes[i].append(endSwingT)
  2581. allData.append([startSwingTimes,endSwingTimes])
  2582. interPawSwingTimes = [[] for i in range(4)]
  2583. interStepTimes = []
  2584. stepLengths = [[] for i in range(4)]
  2585. #pdb.set_trace()
  2586. for nrec in range(5):
  2587. for i in range(4):
  2588. interPawSwingTimes[i].extend(calcualteAllDeltaTValues(allData[nrec][0][i][0],allData[nrec][0][i][0]))
  2589. stepLengths[i].extend(allData[nrec][1][i][0]-allData[nrec][0][i][0])
  2590. interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][0][0],allData[nrec][0][1][0]))
  2591. interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][0][0], allData[nrec][0][2][0]))
  2592. interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][0][0], allData[nrec][0][3][0]))
  2593. interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][1][0], allData[nrec][0][2][0]))
  2594. interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][1][0], allData[nrec][0][3][0]))
  2595. interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][2][0], allData[nrec][0][3][0]))
  2596. #pdb.set_trace()
  2597. pawSwingTimes.append([allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay,interPawSwingTimes,stepLengths,interStepTimes])
  2598. return pawSwingTimes
  2599. #################################################################################
  2600. # remove empty columns and row - from the image registration routine
  2601. #################################################################################
  2602. def removeEmptyColumnAndRows(img):
  2603. # hmask = np.invert(np.sum(img, axis=0) == 0) # looking for zeros does not work as the boundary values are not zeros all the time
  2604. # vmask = np.invert(np.sum(img, axis=1) == 0)
  2605. htemp = (img == img[0,:]) # look instead for same values in row
  2606. hmask = np.invert(htemp.all(axis=0))
  2607. vtemp = (img == img[:,0])
  2608. vmask = np.invert(vtemp.all(axis=1))
  2609. idxH = np.arange(len(hmask))[np.hstack((False,np.diff(hmask)>0))]
  2610. idxV = np.arange(len(vmask))[np.hstack((False, np.diff(vmask)>0))]
  2611. #pdb.set_trace()
  2612. if len(idxV) == 0:
  2613. idxV = np.array([0,np.shape(img)[0]])
  2614. if len(idxH) == 0:
  2615. idxH = np.array([0,np.shape(img)[1]])
  2616. croppedImg = img[idxV[0]:idxV[1],idxH[0]:idxH[1]]
  2617. cutLengths = np.vstack((idxV,idxH))
  2618. #pdb.set_trace()
  2619. return cutLengths
  2620. #################################################################################
  2621. # remove empty columns and row - from the image registration routine
  2622. #################################################################################
  2623. def alignTwoImages(imgA,cutLengthsA,imgB,cutLengthsB,refDate,otherDate,movementValues,figSave=False,figDir=''):
  2624. #matplotlib.use('TkAgg')
  2625. column1 = np.maximum(cutLengthsA[:,0],cutLengthsB[:,0])
  2626. column2 = np.minimum(cutLengthsA[:,1],cutLengthsB[:,1])
  2627. cutLenghts = np.column_stack((column1,column2))
  2628. imgA = imgA[cutLenghts[0,0]:cutLenghts[0,1],cutLenghts[1,0]:cutLenghts[1,1]]
  2629. imgB = imgB[cutLenghts[0,0]:cutLenghts[0,1],cutLenghts[1,0]:cutLenghts[1,1]]
  2630. # Find size of ref image
  2631. sz = imgA.shape
  2632. corr = signal.correlate(imgA - imgA.mean(), imgB - imgB.mean(), mode='same', method='fft')
  2633. maxIdx = np.unravel_index(np.argmax(corr, axis=None), corr.shape)
  2634. shifty = np.shape(imgA)[0]/2. - maxIdx[0]
  2635. shiftx = np.shape(imgA)[1]/2. - maxIdx[1]
  2636. print('max of cross-correlation : ', shiftx, shifty )
  2637. #pdb.set_trace()
  2638. # Define the motion model
  2639. #warp_mode = cv2.MOTION_TRANSLATION #cv2.MOTION_EUCLIDEAN # cv2.MOTION_TRANSLATION # MOTION_EUCLIDEAN
  2640. warp_mode = cv2.MOTION_AFFINE #EUCLIDEAN #HOMOGRAPHY
  2641. warp_modes = [cv2.MOTION_AFFINE,cv2.MOTION_EUCLIDEAN,cv2.MOTION_TRANSLATION]
  2642. # Define 2x3 or 3x3 matrices and initialize the matrix to identity
  2643. if warp_mode == cv2.MOTION_HOMOGRAPHY:
  2644. warp_matrix = np.eye(3, 3, dtype=np.float32)
  2645. else:
  2646. warp_matrix = np.eye(2, 3, dtype=np.float32)
  2647. # Specify the number of iterations.
  2648. number_of_iterations = 1000
  2649. # Specify the threshold of the increment
  2650. # in the correlation coefficient between two iterations
  2651. termination_eps = 1e-10
  2652. # Define termination criteria
  2653. criteria = (cv2.TERM_CRITERIA_EPS | cv2.TERM_CRITERIA_COUNT, number_of_iterations, termination_eps)
  2654. # Run the ECC algorithm. The results are stored in warp_matrix.
  2655. #try:
  2656. imA_u8 = (((imgA-np.min(imgA))/(np.max(imgA)-np.min(imgA)))*255).astype(np.uint8) #cv2.cvtColor(imgA,cv2.COLOR_BGR2GRAY)
  2657. imB_u8 = (((imgB-np.min(imgB))/(np.max(imgB)-np.min(imgB)))*255).astype(np.uint8)
  2658. #pdb.set_trace()
  2659. warpResults = []
  2660. corrMax = []
  2661. for w in range(len(warp_modes)):
  2662. print('testing ', warp_modes[w])
  2663. # warp_matrix1[0, 2] = aS.xOffset
  2664. # warp_matrix1[1, 2] = aS.yOffset
  2665. if (movementValues[0] != 0) and (movementValues[1] != 0):
  2666. warp_matrix[0, 2] = movementValues[0] # -20.
  2667. warp_matrix[1, 2] = movementValues[1] # -40.
  2668. else:
  2669. warp_matrix[0, 2] = shiftx # -20.
  2670. warp_matrix[1, 2] = shifty # -40.
  2671. #try:
  2672. # (cc, warp_matrixRet) = cv2.findTransformECC(imA_u8, imB_u8, warp_matrix, warp_modes[w], criteria, inputMask = None, gaussFiltSize=5)
  2673. #except TypeError:
  2674. try :
  2675. (cc, warp_matrixRet) = cv2.findTransformECC(imA_u8, imB_u8, warp_matrix, warp_modes[w], criteria, inputMask = None)
  2676. except:
  2677. print('findTransformECC did not converge')
  2678. cc = -1
  2679. warp_matrixRet = warp_matrix
  2680. #print(warp_matrixRet,warp_matrix)
  2681. warpResults.append([w,warp_modes[w],np.copy(warp_matrixRet),np.copy(cc)])
  2682. corrMax.append(cc)
  2683. #if cc>0.8:
  2684. # break
  2685. print(warpResults)
  2686. corrMax = np.asarray(corrMax)
  2687. maxCorr = np.argmax(corrMax)
  2688. cc_max = warpResults[maxCorr][3]
  2689. warp_matrix_max = warpResults[maxCorr][2]
  2690. #except:
  2691. #print('findTransformECC output : ',cc,warp_matrixRet)
  2692. #print('find image transformation did not converge')
  2693. #warp_matrixRet = np.copy(warp_matrix_max)
  2694. #cc = None
  2695. #else:
  2696. #pass
  2697. #(cc2, warp_matrix2Ret) = cv2.findTransformECC(imBD, im820, warp_matrix2, warp_mode, criteria)
  2698. if warp_mode == cv2.MOTION_HOMOGRAPHY:
  2699. # Use warpPerspective for Homography
  2700. imgB_aligned = cv2.warpPerspective(imgB, warp_matrix_max, (sz[1], sz[0]), flags=cv2.INTER_LINEAR + cv2.WARP_INVERSE_MAP)
  2701. else:
  2702. # Use warpAffine for Translation, Euclidean and Affine
  2703. imgB_aligned = cv2.warpAffine(imgB, warp_matrix_max, (sz[1], sz[0]), flags=cv2.INTER_LINEAR + cv2.WARP_INVERSE_MAP);
  2704. print('result of image alignment-> warp-matrix and correlation coefficient : ', warp_matrix_max, cc_max)
  2705. if figSave :
  2706. ##################################################################
  2707. # Show final results
  2708. # figure #################################
  2709. fig_width = 10 # width in inches
  2710. fig_height = 10 # height in inches
  2711. fig_size = [fig_width, fig_height]
  2712. params = {'axes.labelsize': 11, 'axes.titlesize': 11, 'font.size': 11, 'xtick.labelsize': 11, 'ytick.labelsize': 11, 'figure.figsize': fig_size, 'savefig.dpi': 600,
  2713. 'axes.linewidth': 1.3, 'ytick.major.size': 4, # major tick size in points
  2714. 'xtick.major.size': 4 # major tick size in points
  2715. # 'edgecolor' : None
  2716. # 'xtick.major.size' : 2,
  2717. # 'ytick.major.size' : 2,
  2718. }
  2719. rcParams.update(params)
  2720. # set sans-serif font to Arial
  2721. rcParams['font.sans-serif'] = 'Arial'
  2722. # create figure instance
  2723. fig = plt.figure()
  2724. # define sub-panel grid and possibly width and height ratios
  2725. gs = gridspec.GridSpec(2, 2 # ,
  2726. # width_ratios=[1.2,1]
  2727. # height_ratios=[1,1]
  2728. )
  2729. # define vertical and horizontal spacing between panels
  2730. gs.update(wspace=0.3, hspace=0.3)
  2731. # possibly change outer margins of the figure
  2732. plt.subplots_adjust(left=0.05, right=0.95, top=0.92, bottom=0.06)
  2733. # sub-panel enumerations
  2734. # plt.figtext(0.06, 0.92, 'A',clip_on=False,color='black', weight='bold',size=22)
  2735. # first sub-plot #######################################################
  2736. # gssub = gridspec.GridSpecFromSubplotSpec(1, 2, subplot_spec=gs[0],hspace=0.2)
  2737. # ax0 = plt.subplot(gssub[0])
  2738. # fig = plt.figure(figsize=(10,10))
  2739. #plt.figtext(0.1, 0.95, '%s ' % (aS.animalID), clip_on=False, color='black', size=14)
  2740. ax0 = plt.subplot(gs[0])
  2741. ax0.set_title('reference image %s' % refDate)
  2742. ax0.imshow(imgA)
  2743. ax0 = plt.subplot(gs[1])
  2744. ax0.set_title('to-be-aligned image %s' % otherDate )
  2745. ax0.imshow(imgB)
  2746. ax0 = plt.subplot(gs[2])
  2747. ax0.set_title('overlay of both images')
  2748. overlayBefore = cv2.addWeighted(imgA/np.max(imgA), 1, imgB/np.max(imgB), 1, 0)
  2749. ax0.imshow(overlayBefore)
  2750. ax0 = plt.subplot(gs[3])
  2751. ax0.set_title('overlay after alignement c = %s \nof BD-AD images' % np.round(cc_max,4), fontsize=10)
  2752. overlayAfter = cv2.addWeighted(imgA/np.max(imgA), 1, imgB_aligned/np.max(imgB_aligned), 1, 0)
  2753. ax0.imshow(overlayAfter)
  2754. #plt.show()
  2755. plt.savefig(figDir + 'ImageAlignment_%s-%s.pdf' % (refDate,otherDate)) # plt.savefig(figOutDir+'ImageAlignment_%s.png' % aS.animalID) # plt.show()
  2756. plt.close()
  2757. return (warp_matrix_max,cc_max)
  2758. #################################################################################
  2759. # calculate correlations between ca-imaging, wheel speed and paw speed
  2760. #################################################################################
  2761. def alignROIsCheckOverlap(statRef,opsRef,statAlign,opsAlign,warp_matrix,refDate,otherDate,figSave=False,figDir=''):
  2762. ncellsRef= len(statRef)
  2763. ncellsAlign = len(statAlign)
  2764. imMaskRef = np.zeros((opsRef['Ly'], opsRef['Lx']))
  2765. imMaskAlign = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
  2766. intersectionROIs = []
  2767. intersectionROIsA = []
  2768. for n in range(0,ncellsRef):
  2769. imMaskRef[:] = 0
  2770. #if iscellBD[n][0]==1:
  2771. #pdb.set_trace()
  2772. ypixRef = statRef[n]['ypix']
  2773. xpixRef = statRef[n]['xpix']
  2774. imMaskRef[ypixRef,xpixRef] = 1
  2775. for m in range(0,ncellsAlign):
  2776. imMaskAlign[:] = 0
  2777. #if iscellAD[m][0]==1:
  2778. ypixAl = statAlign[m]['ypix']
  2779. xpixAl = statAlign[m]['xpix']
  2780. # perform homographic transform : rotation + translation
  2781. #pdb.set_trace()
  2782. points = np.column_stack((xpixAl,ypixAl))
  2783. newPoints = np.copy(points)
  2784. #pdb.set_trace()
  2785. warp_matrix_inverse = np.copy(warp_matrix)
  2786. cv2.invertAffineTransform(warp_matrix,warp_matrix_inverse)
  2787. #newPoints = cv2.transform(points,warp_matrix_inverse)
  2788. xpixAlPrime = np.rint(xpixAl*warp_matrix_inverse[0,0] + ypixAl*warp_matrix_inverse[0,1] + warp_matrix_inverse[0,2])
  2789. ypixAlPrime = np.rint(xpixAl*warp_matrix_inverse[1,0] + ypixAl*warp_matrix_inverse[1,1] + warp_matrix_inverse[1,2]) # - np.rint(warp_matrix[1,2])
  2790. xpixAlPrime = np.array(xpixAlPrime,dtype=int)
  2791. ypixAlPrime = np.array(ypixAlPrime,dtype=int)
  2792. #pdb.set_trace()
  2793. # make sure pixels remain within
  2794. xpixAlPrime2 = xpixAlPrime[(xpixAlPrime<opsAlign['Lx'])&(ypixAlPrime<opsAlign['Ly'])]
  2795. ypixAlPrime2 = ypixAlPrime[(xpixAlPrime<opsAlign['Lx'])&(ypixAlPrime<opsAlign['Ly'])]
  2796. imMaskAlign[ypixAlPrime2,xpixAlPrime2] = 1
  2797. #imMaskAlign[xpixAlPrime2,ypixAlPrime2] = 1
  2798. intersection = np.sum(np.logical_and(imMaskRef,imMaskAlign))
  2799. eitherOr = np.sum(np.logical_or(imMaskRef,imMaskAlign))
  2800. if intersection>0.2:
  2801. #print(n,m,intersection,eitherOr,intersection/eitherOr)
  2802. intersectionROIs.append([n,m,xpixRef,ypixRef,xpixAlPrime2,ypixAlPrime2,intersection,eitherOr,intersection/eitherOr])
  2803. intersectionROIsA.append([n,m,intersection,eitherOr,intersection/eitherOr])
  2804. # clean up intersection ROIs; each ROI should only overlap once
  2805. def removeDoubleCellOccurrences(interROIs,column):
  2806. uniquePerColumn = np.unique(interROIs[:,column],return_counts=True) # find unique occurrences
  2807. multipleCells = uniquePerColumn[0][uniquePerColumn[1]>1] # which cells occur more than once in the first column
  2808. indiciesToRemove = []
  2809. for i in multipleCells:
  2810. indicies = np.argwhere(interROIs[:,column]==i)
  2811. maxIdx = np.argmax(interROIs[indicies[:,0]][:,4])
  2812. delIndicies = np.delete(indicies,maxIdx)
  2813. indiciesToRemove.extend(delIndicies)
  2814. return indiciesToRemove
  2815. intersectionROIsA = np.asarray(intersectionROIsA)
  2816. removeIdicies0 = removeDoubleCellOccurrences(intersectionROIsA,0)
  2817. removeIdicies1 = removeDoubleCellOccurrences(intersectionROIsA,1)
  2818. removeIndicies = np.asarray(removeIdicies0 + removeIdicies1)
  2819. removeIndicies = np.unique(removeIndicies)
  2820. cleanedIntersectionROIs = []
  2821. for i in range(len(intersectionROIs)):
  2822. if i not in removeIndicies:
  2823. cleanedIntersectionROIs.append(intersectionROIs[i])
  2824. #pdb.set_trace()
  2825. if len(removeIndicies)>0:
  2826. intersectionROIsA = np.delete(intersectionROIsA,removeIndicies,axis=0)
  2827. #pdb.set_trace()
  2828. if figSave:
  2829. imRef = opsRef['meanImg']
  2830. imAlign = opsAlign['meanImg']
  2831. ##################################################################
  2832. # Show final results
  2833. fig = plt.figure(figsize=(15, 15)) ########################
  2834. plt.figtext(0.1, 0.95, '%s and %s' % (refDate,otherDate), clip_on=False, color='black', size=14)
  2835. ax0 = fig.add_subplot(3, 2, 1) #############################
  2836. ax0.set_title('reference img')
  2837. ax0.imshow(imRef)
  2838. ax0 = fig.add_subplot(3, 2, 2) #############################
  2839. ax0.set_title('image to be aligned')
  2840. ax0.imshow(imAlign)
  2841. ax0 = fig.add_subplot(3, 2, 3) #############################
  2842. ax0.set_title('ROIs in reference image')
  2843. imRef = np.zeros((opsRef['Ly'], opsRef['Lx']))
  2844. imRefB = np.zeros((opsRef['Ly'], opsRef['Lx']))
  2845. for n in range(0, ncellsRef):
  2846. ypixR = statRef[n]['ypix']
  2847. xpixR = statRef[n]['xpix']
  2848. imRef[ypixR, xpixR] = n + 1
  2849. imRefB[ypixR, xpixR] = 1
  2850. ax0.imshow(imRef, cmap='gist_ncar')
  2851. ax0 = fig.add_subplot(3, 2, 4) #############################
  2852. ax0.set_title('ROIs in aligned image')
  2853. imAlign = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
  2854. imAlignB = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
  2855. for n in range(0, ncellsAlign):
  2856. ypixA = statAlign[n]['ypix']
  2857. xpixA = statAlign[n]['xpix']
  2858. imAlign[ypixA, xpixA] = n + 1
  2859. imAlignB[ypixA, xpixA] = 2
  2860. ax0.imshow(imAlign, cmap='gist_ncar')
  2861. ax0 = fig.add_subplot(3, 2, 5) #############################
  2862. ax0.set_title('overlapping ROIs Ref-Aligned')
  2863. imRef = np.zeros((opsRef['Ly'], opsRef['Lx']))
  2864. imAlign = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
  2865. for n in range(0, len(cleanedIntersectionROIs)):
  2866. ypixR = cleanedIntersectionROIs[n][3]
  2867. xpixR = cleanedIntersectionROIs[n][2]
  2868. ypixA = cleanedIntersectionROIs[n][5]
  2869. xpixA = cleanedIntersectionROIs[n][4]
  2870. imRef[ypixR, xpixR] = 1
  2871. imAlign[ypixA, xpixA] = 2
  2872. overlayBothROIs1 = cv2.addWeighted(imRef, 1, imAlign, 1, 0)
  2873. #overlayBothROIs1B = cv2.addWeighted(imRefB, 1, imAlignB, 1, 0)
  2874. ax0.imshow(overlayBothROIs1)
  2875. ax0 = fig.add_subplot(3, 2, 6) #############################
  2876. ax0.set_title('fraction of ROI overlap Ref-Aligned')
  2877. interFractions1 = []
  2878. for n in range(0, len(cleanedIntersectionROIs)):
  2879. interFractions1.append(cleanedIntersectionROIs[n][8])
  2880. ax0.hist(interFractions1, bins=15)
  2881. plt.savefig(figDir + 'ROIalignment_%s-%s.pdf' % (refDate, otherDate))
  2882. #plt.show()
  2883. plt.close()
  2884. return (cleanedIntersectionROIs,intersectionROIsA)
  2885. #pickle.dump(intersectionROIs, open( dataOutDir + 'ROIintersections_%s.p' % aS.animalID, 'wb' ) )
  2886. #################################################################################
  2887. # find ROI recorded on ref day and on any other given day
  2888. #################################################################################
  2889. def findMatchingRois(mouse,allCorrDataPerSession,analysisLocation,refDate=0):
  2890. # check for sanity
  2891. nDays = len(allCorrDataPerSession)
  2892. refDay = allCorrDataPerSession[refDate][0]
  2893. print('fluo images will be aligned to recordings of :', refDay)
  2894. refDayCaData = allCorrDataPerSession[refDate][3][0]
  2895. refImg = refDayCaData[2]['meanImgE']
  2896. refImgCutLengths = removeEmptyColumnAndRows(refImg)
  2897. opsRef = refDayCaData[2]
  2898. statRef = refDayCaData[4]
  2899. # create list of recoridng day indicies
  2900. recDaysList = [i for i in range(nDays)]
  2901. movementValuesPreset = np.zeros((len(recDaysList), 2))
  2902. # movementValuesPreset[0] = np.array([-20,-47])
  2903. movementValuesPreset[1] = np.array([143, 153])
  2904. # movementValuesPreset[3] = np.array([-1,15])
  2905. # remove day used for referencing
  2906. recDaysList.remove(refDate)
  2907. if os.path.exists(analysisLocation+'/alignmentData.p'):
  2908. allDataRead = pickle.load(open(analysisLocation+'/alignmentData.p'))
  2909. else:
  2910. allDataRead = None
  2911. allData = []
  2912. for nDay in recDaysList:
  2913. print(allCorrDataPerSession[nDay][0],nDay)
  2914. #imgE = allCorrDataPerSession[nDay][3][0][2]['meanImgE']
  2915. img = allCorrDataPerSession[nDay][3][0][2]['meanImgE']
  2916. cutLengths = removeEmptyColumnAndRows(img)
  2917. if allDataRead is not None:
  2918. warp_matrix = allDataRead[nDay][3]
  2919. else:
  2920. (warp_matrix,cc) = alignTwoImages(refImg,refImgCutLengths,img,cutLengths,allCorrDataPerSession[refDate][0],allCorrDataPerSession[nDay][0],movementValuesPreset[nDay],figShow=True,)
  2921. opsAlign = allCorrDataPerSession[nDay][3][0][2]
  2922. statAlign = allCorrDataPerSession[nDay][3][0][4]
  2923. (cleanedIntersectionROIs,intersectionROIsA) = alignROIsCheckOverlap(statRef,opsRef,statAlign,opsAlign,warp_matrix,allCorrDataPerSession[refDate][0],allCorrDataPerSession[nDay][0],showFig=True)
  2924. print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(statRef), len(statAlign), len(cleanedIntersectionROIs))
  2925. allData.append([allCorrDataPerSession[nDay][0],nDay,cutLengths,warp_matrix,cc,cleanedIntersectionROIs,intersectionROIsA])
  2926. intersectingCellsInRefRecording = np.arange(len(statRef))
  2927. for nDay in recDaysList:
  2928. intersectingCellsInRefRecording = np.intersect1d(intersectingCellsInRefRecording,allData[nDay][5][:,0])
  2929. print(nDay,allCorrDataPerSession[nDay][0],intersectingCellsInRefRecording)
  2930. pdb.set_trace()
  2931. return 0
  2932. #################################################################################
  2933. # correlates mean fluo images recorded all possible recording day combinations
  2934. #################################################################################
  2935. def findOverlayMatchingRoisAllDayCombinations(allCorrDataPerSession, figLocation, allDataRead=None,saveFigure=True):
  2936. nDays = len(allCorrDataPerSession)
  2937. movementValuesPreset = np.zeros((nDays*nDays, 2))
  2938. allOverlayData = {}
  2939. #corrMatrix = np.zeros((nDays,nDays))
  2940. nPair = 0
  2941. for nDayA in range(nDays):
  2942. for nDayB in range(nDays):
  2943. if nDayA != nDayB :
  2944. print(nDayA,nDayB, allCorrDataPerSession[nDayA]['folder'], allCorrDataPerSession[nDayB]['folder'])
  2945. #imgA = allCorrDataPerSession[nDayA][3][0][2]['meanImg']
  2946. imgA = allCorrDataPerSession[nDayA]['caImg']['ops']['meanImg']
  2947. cutLengthsA = removeEmptyColumnAndRows(imgA)
  2948. opsA = allCorrDataPerSession[nDayA]['caImg']['ops'] # allCorrDataPerSession[nDayA][3][0][2]
  2949. statA = allCorrDataPerSession[nDayA]['caImg']['stat'] # allCorrDataPerSession[nDayA][3][0][4]
  2950. imgB = allCorrDataPerSession[nDayB]['caImg']['ops']['meanImg'] # allCorrDataPerSession[nDayB][3][0][2]['meanImg']
  2951. cutLengthsB = removeEmptyColumnAndRows(imgB)
  2952. opsB = allCorrDataPerSession[nDayB]['caImg']['ops'] # allCorrDataPerSession[nDayB][3][0][2]
  2953. statB = allCorrDataPerSession[nDayB]['caImg']['stat'] # allCorrDataPerSession[nDayB][3][0][4]
  2954. if (allDataRead is not None) and (allDataRead[nPair][0]==allCorrDataPerSession[nDayA]['folder']) and (allDataRead[nPair][1]==allCorrDataPerSession[nDayB]['folder']):
  2955. print('warp_matrix for current pair of recordings exists and will be used')
  2956. warp_matrix = allDataRead[nPair][6]
  2957. cc = allDataRead[nPair][7]
  2958. else:
  2959. (warp_matrix,cc) = alignTwoImages(imgA,cutLengthsA,imgB,cutLengthsB,allCorrDataPerSession[nDayA]['folder'],allCorrDataPerSession[nDayB]['folder'],movementValuesPreset[nPair],figSave=saveFigure,figDir=figLocation)
  2960. #corrMatrix[nDayA,nDayB] = cc
  2961. (cleanedIntersectionROIs, intersectionROIsA) = alignROIsCheckOverlap(statA, opsA, statB, opsB, warp_matrix, allCorrDataPerSession[nDayA]['folder'], allCorrDataPerSession[nDayB]['folder'],figSave=saveFigure, figDir=figLocation)
  2962. print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(statA), len(statB), len(cleanedIntersectionROIs))
  2963. #allOverlayData.append([allCorrDataPerSession[nDayA]['folder'], allCorrDataPerSession[nDayB]['folder'], nDayA, nDayB, cutLengthsA, cutLengthsB, warp_matrix, cc,cleanedIntersectionROIs, intersectionROIsA,statA,statB])
  2964. allOverlayData[nPair] = {}
  2965. allOverlayData[nPair]['folderA'] = allCorrDataPerSession[nDayA]['folder'] #0
  2966. allOverlayData[nPair]['folderB'] = allCorrDataPerSession[nDayB]['folder'] #1
  2967. allOverlayData[nPair]['nDayA'] = nDayA #2
  2968. allOverlayData[nPair]['nDayB'] = nDayB #3
  2969. allOverlayData[nPair]['cutLengthsA'] = cutLengthsA #4
  2970. allOverlayData[nPair]['cutLengthsB'] = cutLengthsB #5
  2971. allOverlayData[nPair]['warp_matrix'] = warp_matrix #6
  2972. allOverlayData[nPair]['cc'] = cc #7
  2973. allOverlayData[nPair]['cleanedIntersectionROIs'] = cleanedIntersectionROIs #8
  2974. allOverlayData[nPair]['intersectionROIsA'] = intersectionROIsA #9
  2975. allOverlayData[nPair]['statA'] = statA #10
  2976. allOverlayData[nPair]['statB'] = statB #11
  2977. nPair+=1
  2978. #pdb.set_trace()
  2979. return allOverlayData
  2980. #################################################################################
  2981. # correlates mean fluo images of one day recorded at 910 and 820 nm
  2982. #################################################################################
  2983. def findOverlayMatchingRoisDuringOneDay(allCorrDataPerSession910,allCorrDataPerSession820, figLocation, allDataRead=None,saveFigure=True):
  2984. nDays910 = len(allCorrDataPerSession910)
  2985. nDays820 = len(allCorrDataPerSession820)
  2986. movementValuesPreset = np.zeros((nDays910, 2))
  2987. allAlignData = {}
  2988. #corrMatrix = np.zeros((nDays,nDays))
  2989. maxTimeDelay = 45*60 # maximal 40 min difference btw. 910 and 820 recording
  2990. nPair = 0
  2991. for nDay910 in range(nDays910):
  2992. for nDay820 in range(nDays820):
  2993. if allCorrDataPerSession910[nDay910]['folder'][:-4] == allCorrDataPerSession820[nDay820]['folder'][:-4]:
  2994. #pdb.set_trace()
  2995. timeDiff = np.abs(allCorrDataPerSession910[nDay910]['caImg']['timeStamps'][0,3] - allCorrDataPerSession820[nDay820]['caImg']['timeStamps'][0,3])
  2996. if timeDiff>maxTimeDelay:
  2997. print('PROBLEM in data consistency!')
  2998. print('910 and 820 recordings are separated by %s min' % str(timeDiff/60.))
  2999. pdb.set_trace()
  3000. else:
  3001. print('Delay between 910 and 820 imaging sessions is :', timeDiff/60.,'min')
  3002. print(nDay910,nDay820, allCorrDataPerSession910[nDay910]['folder'], allCorrDataPerSession820[nDay820]['folder'])
  3003. img910 = allCorrDataPerSession910[nDay910]['caImg']['ops']['meanImg'] #allCorrDataPerSession910[nDay910][3][0][2]['meanImg']
  3004. cutLengths910 = removeEmptyColumnAndRows(img910)
  3005. ops910 = allCorrDataPerSession910[nDay910]['caImg']['ops'] # allCorrDataPerSession910[nDay910][3][0][2]
  3006. stat910 = allCorrDataPerSession910[nDay910]['caImg']['stat']
  3007. img820 = allCorrDataPerSession820[nDay820]['caImg']['ops']['meanImg'] #[3][0][2]['meanImg']
  3008. cutLengths820 = removeEmptyColumnAndRows(img820)
  3009. ops820 = allCorrDataPerSession820[nDay820]['caImg']['ops'] #[3][0][2]
  3010. stat820 = allCorrDataPerSession820[nDay820]['caImg']['stat'] #[3][0][4]
  3011. if (allDataRead is not None) and (allDataRead[nPair][0]==allCorrDataPerSession910[nDay910]['folder']) and (allDataRead[nPair][1]==allCorrDataPerSession820[nDay820]['folder']):
  3012. print('warp_matrix for current pair of recordings exists and will be used')
  3013. warp_matrix = allDataRead[nPair][6]
  3014. cc = allDataRead[nPair][7]
  3015. else:
  3016. (warp_matrix,cc) = alignTwoImages(img910,cutLengths910,img820,cutLengths820,allCorrDataPerSession910[nDay910]['folder'],allCorrDataPerSession820[nDay820]['folder'],movementValuesPreset[nPair],figSave=saveFigure,figDir=figLocation)
  3017. #corrMatrix[nDayA,nDayB] = cc
  3018. (cleanedIntersectionROIs, intersectionROIsA) = alignROIsCheckOverlap(stat910, ops910, stat820, ops820, warp_matrix, allCorrDataPerSession910[nDay910]['folder'], allCorrDataPerSession820[nDay820]['folder'],figSave=saveFigure, figDir=figLocation)
  3019. print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(stat910), len(stat820), len(cleanedIntersectionROIs))
  3020. #allAlignData.append([allCorrDataPerSession910[nDay910][0], allCorrDataPerSession820[nDay820][0], nDay910, nDay820, cutLengths910, cutLengths820, warp_matrix, cc,cleanedIntersectionROIs, intersectionROIsA,stat910,stat820])
  3021. allAlignData[nPair] = {}
  3022. allAlignData[nPair]['folder910'] = allCorrDataPerSession910[nDay910]['folder'] #0
  3023. allAlignData[nPair]['folder820'] = allCorrDataPerSession820[nDay820]['folder'] #1
  3024. allAlignData[nPair]['nDay910'] = nDay910 #2
  3025. allAlignData[nPair]['nDay820'] = nDay820 #3
  3026. allAlignData[nPair]['cutLengths910'] = cutLengths910 #4
  3027. allAlignData[nPair]['cutLengths820'] = cutLengths820 #5
  3028. allAlignData[nPair]['warp_matrix'] = warp_matrix #6
  3029. allAlignData[nPair]['cc'] = cc # 7
  3030. allAlignData[nPair]['cleanedIntersectionROIs'] = cleanedIntersectionROIs #8
  3031. allAlignData[nPair]['intersectionROIsA'] = intersectionROIsA #9
  3032. allAlignData[nPair]['statA'] = stat910 #10
  3033. allAlignData[nPair]['statB'] = stat820 #11
  3034. nPair+=1
  3035. return allAlignData
  3036. #################################################################################
  3037. # find ROIs recorded across successive recording days
  3038. #################################################################################
  3039. def findMatchingRoisSuccessivDays(mouse,allCorrDataPerSession,analysisLocation,expDate,figLocation, allDataRead=None):
  3040. # check for sanity
  3041. nDays = len(allCorrDataPerSession)
  3042. # create list of recoridng day indicies
  3043. #recDaysList = [i for i in range(nDays)]
  3044. movementValuesPreset = np.zeros((nDays, 2))
  3045. # movementValuesPreset[0] = np.array([-20,-47])
  3046. # movementValuesPreset[1] = np.array([143, 153])
  3047. # movementValuesPreset[3] = np.array([-1,15])
  3048. allDataStore = []
  3049. for nPair in range(nDays-1):
  3050. nDayA = nPair
  3051. nDayB = nPair + 1
  3052. print(nPair, allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0])
  3053. imgA = allCorrDataPerSession[nDayA][3][0][2]['meanImg']
  3054. cutLengthsA = removeEmptyColumnAndRows(imgA)
  3055. opsA = allCorrDataPerSession[nDayA][3][0][2]
  3056. statA = allCorrDataPerSession[nDayA][3][0][4]
  3057. imgB = allCorrDataPerSession[nDayB][3][0][2]['meanImg']
  3058. cutLengthsB = removeEmptyColumnAndRows(imgB)
  3059. opsB = allCorrDataPerSession[nDayB][3][0][2]
  3060. statB = allCorrDataPerSession[nDayB][3][0][4]
  3061. if (allDataRead is not None) and (allDataRead[nPair][0]==allCorrDataPerSession[nDayA][0]) and (allDataRead[nPair][1]==allCorrDataPerSession[nDayB][0]):
  3062. print('warp_matrix for current pair of recordings exists and will be used')
  3063. warp_matrix = allDataRead[nPair][6]
  3064. cc = allDataRead[nPair][7]
  3065. else:
  3066. (warp_matrix,cc) = alignTwoImages(imgA,cutLengthsA,imgB,cutLengthsB,allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0],movementValuesPreset[nPair],figShow=True,figDir=figLocation)
  3067. (cleanedIntersectionROIs,intersectionROIsA) = alignROIsCheckOverlap(statA,opsA,statB,opsB,warp_matrix,allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0],showFig=True,figDir=figLocation)
  3068. print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(statA), len(statB), len(cleanedIntersectionROIs))
  3069. allDataStore.append([allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0],nDayA,nDayB,cutLengthsA,cutLengthsB,warp_matrix,cc,cleanedIntersectionROIs,intersectionROIsA])
  3070. #intersectingCellsInRefRecording = np.arange(len(statA))
  3071. #for nDay in recDaysList:
  3072. # intersectingCellsInRefRecording = np.intersect1d(intersectingCellsInRefRecording,allData[nDay][5][:,0])
  3073. # print(nDay,allCorrDataPerSession[nDay][0],intersectingCellsInRefRecording)
  3074. #pdb.set_trace()
  3075. return allDataStore
  3076. #################################################################################
  3077. # check which ROIs were recoreded across recordings days
  3078. #################################################################################
  3079. def roisRecordedAllDays(allCorrDataPerSession,allAlignData,alignData910And820,correlationThres):
  3080. def getMatchingPairs(aD):
  3081. ROIpairs = []
  3082. for i in range(len(aD)):
  3083. ROIpairs.append(aD[i][:2])
  3084. ROIpairs = np.asarray(ROIpairs)
  3085. return ROIpairs
  3086. #pdb.set_trace()
  3087. #allCombis = list(itertools.combinations((1,2,3,4,5,6,7,8,9),2))
  3088. nDays = len(allCorrDataPerSession)
  3089. correlationThreshold = correlationThres
  3090. intersectData = {}
  3091. nPair = 0
  3092. for nDayA in range(nDays):
  3093. nRef = 0
  3094. corrDays = []
  3095. # first loop to find indicies which are present in all recordings with good match; note that the idxRemaining will get smaller in each interation over days
  3096. for nDayB in range(nDays):
  3097. if nDayA != nDayB :
  3098. #print(nDayA,nDayB, allCorrDataPerSession[nDayA][0], allCorrDataPerSession[nDayB][0], nRef, nPair)
  3099. if not (nDayA == allAlignData[nPair]['nDayA'] and nDayB == allAlignData[nPair]['nDayB']):
  3100. print(nDayA, nDayB, allAlignData[nPair][2], allAlignData[nPair][3])
  3101. print('sanity check failed! The pairing doesn\'t correspond to the day-pair')
  3102. if nRef==0:
  3103. matchingRoisBefore = getMatchingPairs(allAlignData[nPair]['cleanedIntersectionROIs']) # [8]
  3104. if allAlignData[nPair]['cc']>correlationThreshold:
  3105. corrDays.append([nDayB,allCorrDataPerSession[nDayB]['folder'],nPair]) #[0]
  3106. else:
  3107. matchingRoisAfter = getMatchingPairs(allAlignData[nPair]['cleanedIntersectionROIs']) # [8]
  3108. if allAlignData[nPair]['cc']>correlationThreshold:
  3109. #print(nDayA, nDayB, allCorrDataPerSession[nDayA][0], allCorrDataPerSession[nDayB][0], nRef, nPair)
  3110. idxRemaining = np.intersect1d(matchingRoisBefore[:,0], matchingRoisAfter[:,0]) #
  3111. idxRemainingBefore = [key for key,val in enumerate(matchingRoisBefore[:,0]) if val in idxRemaining]
  3112. idxRemainingAfter = [key for key,val in enumerate(matchingRoisAfter[:,0]) if val in idxRemaining]
  3113. BeforeAlsoAfter = matchingRoisBefore[idxRemainingBefore]
  3114. AfterAlsoBefore = matchingRoisAfter[idxRemainingAfter]
  3115. print('ROIS remaining before and after : ', len(idxRemaining), nDayB, allCorrDataPerSession[nDayB]['folder'])
  3116. matchingRoisBefore = np.copy(BeforeAlsoAfter)
  3117. corrDays.append([nDayB,allCorrDataPerSession[nDayB]['folder'],nPair])
  3118. idxRemainingGood = np.copy(idxRemaining)
  3119. #intersectData.append([nPair,matchingRoisBefore,matchingRoisAfter,idxRemaining,BeforeAlsoAfter,AfterAlsoBefore])
  3120. #print(nPair,matchingRoisBefore,matchingRoisAfter,idxRemaining,BeforeAlsoAfter,AfterAlsoBefore)
  3121. #pdb.set_trace()
  3122. nRef+=1
  3123. nPair+=1
  3124. remainingRoisExists = (True if 'idxRemainingGood' in locals() else False)
  3125. print(nDayA, allCorrDataPerSession[nDayA]['folder'], len(corrDays), (len(idxRemainingGood) if remainingRoisExists else 0), (idxRemainingGood if remainingRoisExists else None), corrDays)
  3126. #intersectData.append([nDayA, allCorrDataPerSession[nDayA]['folder'], len(corrDays), (len(idxRemainingGood) if remainingRoisExists else 0), (idxRemainingGood if remainingRoisExists else None), corrDays])
  3127. intersectData[nDayA] = {}
  3128. intersectData[nDayA]['nDayA'] = nDayA # 0
  3129. intersectData[nDayA]['folder'] = allCorrDataPerSession[nDayA]['folder'] # 1
  3130. intersectData[nDayA]['lenCorrDays'] = len(corrDays) # 2
  3131. intersectData[nDayA]['lenIdxRemainingGood'] = (len(idxRemainingGood) if remainingRoisExists else 0) # 3
  3132. intersectData[nDayA]['idxRemainingGood'] = (idxRemainingGood if remainingRoisExists else None) # 4
  3133. intersectData[nDayA]['corrDays'] = corrDays #5
  3134. # second loop to determine identity of the remaining ROIs in all recordings
  3135. #idxDay = intersectData[idxRef][5][r][0]
  3136. #idxRoi = intersectData[idxRef][5][r][3 + n][1]
  3137. #dayID = intersectData[idxRef][5][r][1]
  3138. #nPairB = 0
  3139. for i in range(len(intersectData)):
  3140. print(i, intersectData[i]['folder'], intersectData[i]['lenCorrDays'],intersectData[i]['lenIdxRemainingGood'])
  3141. idxRemainingGood = intersectData[i]['idxRemainingGood']
  3142. for nDay in range(intersectData[i]['lenCorrDays']):
  3143. #print(allAlignData[intersectData[i][5][2]][2])
  3144. matchingRois = getMatchingPairs(allAlignData[intersectData[i]['corrDays'][nDay][2]]['cleanedIntersectionROIs'])
  3145. #print('match :',matchingRois,idxRemainingGood)
  3146. idxRemaining = [key for key,val in enumerate(matchingRois[:,0]) if val in idxRemainingGood]
  3147. RemainingRois = matchingRois[idxRemaining]
  3148. intersectData[i]['corrDays'][nDay].append(RemainingRois)
  3149. #roiData.append([intersectData[i][0],intersectData[i][1],])
  3150. #pdb.set_trace()
  3151. # another loop to append 820 identities of remaining ROIs
  3152. for i in range(len(intersectData)):
  3153. #print(i, intersectData[i][1], intersectData[i][2], intersectData[i][3])
  3154. idxRemainingGood = intersectData[i]['idxRemainingGood']
  3155. for n in range(len(alignData910And820)):
  3156. if (intersectData[i]['folder'] == alignData910And820[n]['folder820']) and (idxRemainingGood is not None):
  3157. print(i,n,intersectData[i]['folder'],alignData910And820[n]['folder820'])
  3158. matchingRois = getMatchingPairs(alignData910And820[n]['cleanedIntersectionROIs'])
  3159. idxRemaining = [key for key,val in enumerate(matchingRois[:,0]) if val in idxRemainingGood]
  3160. RemainingRois = matchingRois[idxRemaining]
  3161. #intersectData[i].append(RemainingRois)
  3162. intersectData[i]['RemainingRois'] = RemainingRois
  3163. return intersectData
  3164. #################################################################################
  3165. # calculate correlations between ca-imaging, wheel speed and paw speed
  3166. #################################################################################
  3167. def doCorrelationAnalysisLocomotionPeriod(allCorrDataPerSession,allStepData):
  3168. #
  3169. matplotlib.use('TkAgg')
  3170. xPixToUm = 0.79
  3171. yPixToUm = 0.8
  3172. motorizationCorrelations = []
  3173. varExplainedMotorization = []
  3174. maxMotorization = [12, 53] # in sec
  3175. for nDay in range(len(allCorrDataPerSession)):
  3176. print(nDay,allCorrDataPerSession[nDay]['folder'])
  3177. (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP, wheelSpeedDict, pawTracksDict, caTracesDict, slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession,allStepData)
  3178. # ATTENTION : all of the arrays also contain a time array
  3179. # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTimeSame]
  3180. # dims of pawTracksDictInterP : [nSessions][nPaw][2][valuesOverTimeSame]
  3181. # dims of caTracesDictInterP : [nSessions][nRois+1][valuesOverTimeSame]
  3182. # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTime]
  3183. # dims of pawTracksDict : [nSessions][nPaw][2][valuesOverTime]
  3184. # dims of caTracesDict : [nSessions][nRois+1][valuesOverTime]
  3185. # correlations between calcium traces ######################################################################
  3186. stat = allCorrDataPerSession[nDay]['caImg']['stat']#[3][0][4]
  3187. nTrials = len(caTracesDict)
  3188. nRois = np.shape(caTracesDict[0])[0] - 1
  3189. allCoords = []
  3190. for i in range(nRois):
  3191. allCoords.append([i, stat[i]['med'][0], stat[i]['med'][1]]) # first extract the coordinates of all ROIs
  3192. allCoordsSorted = sorted(allCoords, key=lambda x: (x[1], x[2])) # list according to increasing x and y coordinates
  3193. allCoordsSorted = np.asarray(allCoordsSorted)
  3194. combis = list(itertools.combinations(np.array(allCoordsSorted[:, 0], dtype=int), 2)) # use the sorted list to create the combinations
  3195. corrCaTraces = np.zeros((len(combis), nTrials, 9))
  3196. # for moving average windows
  3197. tWindow = 1. # removing changes on the order of 1 s and longer
  3198. ttime = caTracesDict[0][0]
  3199. dt = np.mean(np.diff(ttime))
  3200. Nwindow = int(tWindow / dt + 0.5)
  3201. for i in range(len(combis)):
  3202. # get location information of both ROIs
  3203. xy0 = stat[combis[i][0]]['med']
  3204. xy1 = stat[combis[i][1]]['med']
  3205. # calculate eucleadian distance and x, y distance between cells
  3206. euclDist = np.sqrt(((xy0[1]-xy1[1])*xPixToUm)**2 + ((xy0[0]-xy1[0])*yPixToUm)**2)
  3207. xyDist = ([(xy1[1]-xy0[1])*xPixToUm,(xy0[0]-xy1[0])*yPixToUm]) # note that the the y-value is inverted on purpose; cell order from the top
  3208. #pdb.set_trace()
  3209. for t in range(nTrials):
  3210. mask = (caTracesDict[t][0] >= maxMotorization[0]) & (caTracesDict[t][0] < maxMotorization[1])
  3211. corrTemp = scipy.stats.pearsonr(caTracesDict[t][combis[i][0]+1][mask],caTracesDict[t][combis[i][1]+1][mask])
  3212. caTemp0 = np.convolve(caTracesDict[t][combis[i][0]+1][mask], np.ones((Nwindow,))/Nwindow, mode='same')
  3213. caTemp1 = np.convolve(caTracesDict[t][combis[i][1]+1][mask], np.ones((Nwindow,))/Nwindow, mode='same')
  3214. shortCorrTemp = scipy.stats.pearsonr(caTracesDict[t][combis[i][0]+1][mask]-caTemp0,caTracesDict[t][combis[i][1]+1][mask]-caTemp1)
  3215. corrCaTraces[i,t] = np.array([combis[i][0],combis[i][1],corrTemp[0],corrTemp[1],euclDist,xyDist[0],xyDist[1],shortCorrTemp[0],shortCorrTemp[1]])
  3216. # activity measures of ca traces before vs during walking ######################################################################
  3217. tBaseline = 5
  3218. tMotorization = [12,53]
  3219. activityCaTraces = np.zeros((nRois,nTrials,4))
  3220. for n in range(nRois):
  3221. for t in range(nTrials):
  3222. baselineMask = caTracesDict[t][0]<tBaseline
  3223. activityMask = (caTracesDict[t][0]>=tMotorization[0]) & (caTracesDict[t][0]<tMotorization[1])
  3224. (baseLMean,baseLSTD) = (np.mean(caTracesDict[t][n+1][baselineMask]),np.std(caTracesDict[t][n+1][baselineMask]))
  3225. (actMean, actSTD) = (np.mean(caTracesDict[t][n + 1][activityMask]), np.std(caTracesDict[t][n + 1][activityMask]))
  3226. activityCaTraces[n,t] = np.array([baseLMean,baseLSTD,actMean, actSTD])
  3227. ###################################################################
  3228. # correlation between calcium and paw as well as paw speed
  3229. corrCaWheel = np.zeros((nRois,nTrials, 3))
  3230. corrCaPawTraces = np.zeros((nRois,nTrials, 9))
  3231. for i in range(nRois):
  3232. for t in range(nTrials):
  3233. caMask = (caTracesDictInterP[t][0]>= maxMotorization[0]) & (caTracesDictInterP[t][0] < maxMotorization[1])
  3234. wheelMask = (wheelSpeedDictInterP[t][0]>= maxMotorization[0]) & (wheelSpeedDictInterP[t][0] < maxMotorization[1])
  3235. paw0Mask = (pawTracksDictInterP[t][0][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][0][0] < maxMotorization[1])
  3236. paw1Mask = (pawTracksDictInterP[t][1][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][1][0] < maxMotorization[1])
  3237. paw2Mask = (pawTracksDictInterP[t][2][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][2][0] < maxMotorization[1])
  3238. paw3Mask = (pawTracksDictInterP[t][3][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][3][0] < maxMotorization[1])
  3239. corrWheelTemp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], wheelSpeedDictInterP[t][1][wheelMask])
  3240. corrCaWheel[i][t] = np.array([i,corrWheelTemp[0],corrWheelTemp[1]])
  3241. corrPaw0Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][0][1][paw0Mask])
  3242. corrPaw1Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][1][1][paw1Mask])
  3243. corrPaw2Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][2][1][paw2Mask])
  3244. corrPaw3Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][3][1][paw3Mask])
  3245. corrCaPawTraces[i][t] = np.array([i,corrPaw0Temp[0],corrPaw0Temp[1],corrPaw1Temp[0],corrPaw1Temp[1],corrPaw2Temp[0],corrPaw2Temp[1],corrPaw3Temp[0],corrPaw3Temp[1]])
  3246. ###################################################################
  3247. # correlation btw. PCA components and wheel, paw speeds
  3248. # concatenate calcium, wheel speed and paw speed for PCA
  3249. pawAll = {}
  3250. pawMask = {}
  3251. for t in range(nTrials):
  3252. caMask = (caTracesDictInterP[t][0] >= maxMotorization[0]) & (caTracesDictInterP[t][0] < maxMotorization[1])
  3253. wheelMask = (wheelSpeedDictInterP[t][0] >= maxMotorization[0]) & (wheelSpeedDictInterP[t][0] < maxMotorization[1])
  3254. for i in range(4):
  3255. pawMask[i] = (pawTracksDictInterP[t][i][0] >= maxMotorization[0]) & (pawTracksDictInterP[t][i][0] < maxMotorization[1])
  3256. if t == 0:
  3257. caAll = caTracesDictInterP[t][:,caMask]
  3258. wheelAll = wheelSpeedDictInterP[t][:,wheelMask]
  3259. for i in range(4):
  3260. pawAll[i] = pawTracksDictInterP[t][i][:,pawMask[i]]
  3261. else:
  3262. caAll = np.column_stack((caAll,caTracesDictInterP[t][:,caMask]))
  3263. wheelAll = np.column_stack((wheelAll,wheelSpeedDictInterP[t][:,wheelMask]))
  3264. for i in range(4):
  3265. pawAll[i] = np.column_stack((pawAll[i],pawTracksDictInterP[t][i][:,pawMask[i]]))
  3266. #pdb.set_trace()
  3267. pcaComponents = 5
  3268. print('doing PCA ...')
  3269. X = np.transpose(caAll[1:])
  3270. pca = PCA(n_components=pcaComponents)
  3271. pca.fit(X)
  3272. X_pca = pca.transform(X)
  3273. #print(pca.components_)
  3274. pcaCorrs = np.zeros((pcaComponents,11))
  3275. varExplainedMotorization.append(pca.explained_variance_ratio_)
  3276. print(pca.explained_variance_ratio_)
  3277. for i in range(pcaComponents):
  3278. corrWheelTemp = scipy.stats.pearsonr(X_pca[:,i], wheelAll[1])
  3279. corrPaw0Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[0][1])
  3280. corrPaw1Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[1][1])
  3281. corrPaw2Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[2][1])
  3282. corrPaw3Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[3][1])
  3283. pcaCorrs[i] = ([i,corrWheelTemp[0],corrWheelTemp[1],corrPaw0Temp[0],corrPaw0Temp[1],corrPaw1Temp[0],corrPaw1Temp[1],corrPaw2Temp[0],corrPaw2Temp[1],corrPaw3Temp[0],corrPaw3Temp[1]])
  3284. ###################################################################
  3285. motorizationCorrelations.append([nTrials,corrCaTraces,corrCaWheel,corrCaPawTraces,pcaCorrs,activityCaTraces])
  3286. return (motorizationCorrelations,varExplainedMotorization)
  3287. #################################################################################
  3288. # chose day of reference from FOV alignment data
  3289. #################################################################################
  3290. def getRefIdxPerMouse(mouse):
  3291. refIdxMouseDict = {'210214_m12':7,
  3292. '210214_m13':4,
  3293. '210214_m14':5,
  3294. '210214_m15':6,
  3295. '210214_m17':6,
  3296. '210214_m18':9,
  3297. '210214_m19':5,
  3298. '210214_m20':3,
  3299. '210122_f83':2,
  3300. '210122_f84':6,
  3301. '210120_m85':1,
  3302. '210120_m86':0,
  3303. }
  3304. idxRef = refIdxMouseDict[mouse]
  3305. return idxRef
  3306. #################################################################################
  3307. # chose day of reference from FOV alignment data
  3308. ##########################################################

dataAnalysis.py at commit 7217d46, under gpl · at the source

Overview

Authors: Andry Andrianarivelo1, Heike Stein2,3, Jeremy Gabillet1, Clarisse Batifol1, Abdelali Jalil1, N Alex Cayco Gajic2, Michael Graupner1
  1. Université Paris Cité, CNRS, Saints-Pères Paris Institute for the Neurosciences, Paris, France
  2. Laboratoire de Neurosciences Cognitives et Computationnelles, INSERM U960, Département d’Études Cognitives, École Normale Supérieure, PSL University, Paris, France
  3. Institut des Systèmes Intelligents et de Robotique, Sorbonne Université, CNRS, Paris, France
Journal: Nature communications, volume 17, issue 1, article 8861
Dates: received 3 May 2024; accepted 12 June 2026; published online 20 July 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41467-026-74823-1 · PMID 42476976 · PMCID PMC13500626 · OpenAlex W7169770214
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: mouse (organism), cellular / molecular (subfield)
Methods: Spectral & time-frequency, Statistics, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Connectivity, fMRI & imaging, Single-unit activity, calcium imaging
Keywords: Cerebellum, Cellular neuroscience
MeSH: Cerebellum*, Learning*, Locomotion*, Animals, Interneurons, Male, Mice, Mice, Inbred C57BL, Optogenetics, Purkinje Cells (* major topic)
Topic: Vestibular and auditory disorders (Neurology, Neuroscience), according to OpenAlex
Citations: cited by 1 paper (Europe PMC); 58 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repository

Its files are read in the Code ↔ Paper reader above, with 23 matches between paragraphs and lines of code.

mgraupe/LocoReach-analysis-2026

License: gpl
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 7217d46bf438e44d9d4f735d109cab8e4d9f6d2a, 25 April 2026
Languages: Python (54)
Size: 61 files, 54 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (48 files), Matplotlib (43 files), pandas (24 files), SciPy (15 files), scikit-learn (9 files), statsmodels (7 files), OpenCV (6 files), tifffile (6 files), h5py (5 files), seaborn (4 files), UMAP (3 files), Pingouin (2 files), scikit-posthocs (2 files), scikit-image (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
55 files

Code availability statement

The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41467-026-74823-1.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 54 scripts, each with its path and the digest of its content;
  • 23 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data availability statement

The paper has a data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41467-026-74823-1.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 7 authors, 2 keywords, 10 MeSH terms, 55 references.

Cite

This paper

Andrianarivelo, A., Stein, H., Gabillet, J., Batifol, C., Jalil, A., Cayco Gajic, N. A., & Graupner, M. (2026). Cerebellar activity is triggered by reach endpoint during learning of a complex locomotor task. Nature communications, 17(1), 8861. https://doi.org/10.1038/s41467-026-74823-1

BibTeX

@article{andrianarivelo2026cerebellar,
author = {Andrianarivelo, Andry and Stein, Heike and Gabillet, Jeremy and Batifol, Clarisse and Jalil, Abdelali and Cayco Gajic, N Alex and Graupner, Michael},
title = {{Cerebellar activity is triggered by reach endpoint during learning of a complex locomotor task}},
journal = {Nature communications},
year = {2026},
month = jul,
volume = {17},
number = {1},
pages = {8861},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-74823-1},
url = {https://doi.org/10.1038/s41467-026-74823-1},
pmid = {42476976},
pmcid = {PMC13500626}
}

RIS

TY - JOUR
AU - Andrianarivelo, Andry
AU - Stein, Heike
AU - Gabillet, Jeremy
AU - Batifol, Clarisse
AU - Jalil, Abdelali
AU - Cayco Gajic, N Alex
AU - Graupner, Michael
TI - Cerebellar activity is triggered by reach endpoint during learning of a complex locomotor task
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/07/20
VL - 17
IS - 1
SP - 8861
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-74823-1
UR - https://doi.org/10.1038/s41467-026-74823-1
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-74823-1",
"type": "article-journal",
"title": "Cerebellar activity is triggered by reach endpoint during learning of a complex locomotor task",
"container-title": "Nature communications",
"author": [
{
"family": "Andrianarivelo",
"given": "Andry"
},
{
"family": "Stein",
"given": "Heike"
},
{
"family": "Gabillet",
"given": "Jeremy"
},
{
"family": "Batifol",
"given": "Clarisse"
},
{
"family": "Jalil",
"given": "Abdelali"
},
{
"family": "Cayco Gajic",
"given": "N Alex"
},
{
"family": "Graupner",
"given": "Michael"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "8861",
"DOI": "10.1038/s41467-026-74823-1",
"PMID": "42476976",
"PMCID": "PMC13500626",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-74823-1",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
20
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi: [code]
Real-time closed-loop feedback system for mouse mesoscale cortical signal and movement control
Journal: eLife
In common: Pingouin, tifffile, OpenCV, 9 other tools, mouse, 1 reference
[2] doi:10.1016/j.isci.2026.116825 [code]
Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.
Journal: iScience
In common: Pingouin, UMAP, OpenCV, 9 other tools, mouse, 1 reference
[3] doi:10.7554/elife.109717 [code]
Retrosplenial cortex enables context-dependent goal-directed sensorimotor transformation.
Journal: eLife
In common: tifffile, OpenCV, scikit-image, 8 other tools, mouse, 2 references
[4] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: Pingouin, OpenCV, scikit-image, 8 other tools, mouse, 1 reference
[5] doi:10.1016/j.isci.2026.117375 [code]
Motor priming is associated with widespread recruitment into neural ensembles and more rapid ensemble transitions.
Journal: iScience
In common: scikit-posthocs, Pingouin, OpenCV, 8 other tools
[6] doi:10.1038/s41593-026-02362-5 [code]
Replay of procedural memory is independent of the hippocampus.
Journal: Nature neuroscience
In common: scikit-posthocs, Pingouin, h5py, 7 other tools, mouse, 1 reference
[7] doi:10.1016/j.xcrm.2026.102766 [code]
A longitudinal single-cell and spatial multiomic atlas of pediatric high-grade glioma.
Journal: Cell reports. Medicine
In common: tifffile, UMAP, OpenCV, 8 other tools, cellular / molecular
[8] doi:10.1002/advs.202520220 [code]
Learnable Diffusion Framework for Mouse V1 Neural Decoding.
Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)
In common: UMAP, OpenCV, scikit-image, 7 other tools, mouse, 1 reference
[9] doi:10.1038/s41586-026-10323-y [code]
Genetically encoded assembly recorder temporally resolves cellular history.
Journal: Nature
In common: scikit-posthocs, tifffile, OpenCV, 7 other tools, mouse, cellular / molecular
[10] doi:10.1038/s41593-026-02314-z [code]
Low-dimensional population dynamics in the brainstem gate REM sleep.
Journal: Nature neuroscience
In common: Pingouin, OpenCV, h5py, 7 other tools, mouse, 1 reference

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.