Cerebellar activity is triggered by reach endpoint during learning of a complex locomotor task.
The 23 matches
- [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] § Methods › PSTH analysis ↔ getPsortWalkingActivityAndPawTracesCalculatePSTH.py, lines 151–231 · score 0.82 · 0–20, 40–60, 60–80, 80–100, 20–40, swing duration
- [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] § 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] § 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] § 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] § Methods › LocoReach setup ↔ extractSwingStancePhases.py, lines 79–158 · score 0.74 · rotary encoder, behavioral videos, video frames, ACQ4, DAQ, Devices
- [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] § Methods › LocoReach setup ↔ getRawBehaviorImagesSaveVideo.py, lines 56–81 · score 0.69 · rotary encoder, behavioral videos, video frames, Devices, camera, LEDs
- [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] § 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] § 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] § 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] § 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] § Methods › Event kernels and linear models ↔ tools/dataAnalysis.py, lines 2223–2303 · score 0.60 · scikit-learn, linear models, Lasso, Ridge, L1, L2
- [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] § Methods › Event kernels and linear models ↔ tools/dataAnalysis.py, lines 2163–2222 · score 0.58 · L2 regularized, Cross validated, L1, model, regressors, variable
- [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] § 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] § 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] § Methods › Statistics ↔ tools/createGroupVisualizations.py, lines 1–88 · score 0.55 · statsmodels, linear models, Scipy, Pearson, ANOVA, squares
- [22] § Methods › PSTH analysis ↔ tools/dataAnalysis.py, lines 335–405 · score 0.55 · original spike train, Gaussian, histograms, firing rate, bins, window
- [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
- import time
- import numpy as np
- import sys
- import os
- import scipy, scipy.io
- import matplotlib.pyplot as plt
- import matplotlib.patches as patches
- import tifffile as tiff
- from scipy import io
- import pdb
- import scipy.ndimage
- import itertools
- import pandas as pd
- from scipy.interpolate import interp1d
- from sklearn.decomposition import PCA
- import cv2
- from scipy import signal
- from scipy.signal import find_peaks
- import pickle
- import random
- from statsmodels.stats.anova import anova_single
- import scikits.bootstrap as boot
- from scipy import ndimage
- from matplotlib import rcParams
- import matplotlib.pyplot as plt
- import matplotlib.gridspec as gridspec
- import matplotlib.cm as cm
- import matplotlib
- import multiprocessing as mp
- from joblib import Parallel, delayed
- import scipy.stats as stats
- from tools.pyqtgraph.Qt import QtGui, QtCore
- import tools.pyqtgraph as pg
- matplotlib.use('TkAgg') # WxAgg
- from array import array
- import scipy.interpolate as interpolate
- import scipy.optimize as optimize
- import array as arr
- from numpy import trapz
- import numpy as np
- import scipy.optimize
- import matplotlib.pyplot as plt
- from sklearn.cluster import KMeans
- numcores = mp.cpu_count()-1
- from sklearn import linear_model
- from sklearn import preprocessing
- from sklearn.model_selection import GroupKFold
- from sklearn.model_selection import KFold
- from sklearn.model_selection import cross_val_score
- from sklearn.model_selection import RepeatedKFold
- from scipy.stats import vonmises_line
- def getSpeed(angles,times,circumsphere,minSpacing):
- angleJumps = angles[np.concatenate((([True]),np.diff(angles)!=0.))] # find angles at the points where the angle value changes
- timePoints = times[np.concatenate((([True]),np.diff(angles)!=0.))] # find times at the points where the angle value changes
- #
- #
- dt = np.diff(timePoints) # delta t values of the time array
- dtMultipleSpace = dt//minSpacing # how many times does the minSpacing fit in the gaps
- dtMultipleSpace = dtMultipleSpace[dtMultipleSpace>1] # gap must be as big as 2 times the spacing to add a point
- # note that gap must be larger than 2 times the spacing
- startGap = np.arange(len(timePoints))[np.hstack(((dt>2.*minSpacing),np.array(False)))] # index values of the start of the gap
- endGap = np.arange(len(timePoints))[np.hstack((np.array(False),(dt>2.*minSpacing)))] # index values of the end of the gap
- newSpacingValue = (timePoints[endGap] - timePoints[startGap]) / dtMultipleSpace
- newTvalues = []
- newAvalues = []
- for i in range(len(dtMultipleSpace)):
- newTvalues.extend(timePoints[startGap][i] + newSpacingValue[i] * np.arange(1, dtMultipleSpace[i]))
- newAvalues.extend(np.repeat(angleJumps[startGap][i], (dtMultipleSpace[i] - 1)))
- timesNew = np.hstack((timePoints,np.asarray(newTvalues)))
- anglesNew = np.hstack((angleJumps,np.asarray(newAvalues)))
- both = np.row_stack((timesNew,anglesNew))
- bothSorted = both[:,both[0].argsort()]
- angularSpeed = np.diff(bothSorted[1])/np.diff(bothSorted[0])
- angularSpeedM = (angularSpeed[1:]+angularSpeed[:-1])/2.
- linearSpeed = angularSpeedM*circumsphere/360.
- speedTimes = bothSorted[0][1:-1]
- #pdb.set_trace()
- angularSpeedSmooth = scipy.signal.medfilt(angularSpeed,kernel_size=9)
- linearSpeedSmooth = scipy.signal.medfilt(linearSpeed,kernel_size=9)
- return (angularSpeedSmooth,linearSpeedSmooth,speedTimes,angularSpeed,linearSpeed)
- def crosscorr(deltat, y0, y1, correlationRange=1.5, fast=False):
- """
- home-written routine to calcualte cross-correlation between two contiuous traces
- new version from February 9th, 2011
- """
- if len(y0) != len(y1):
- print('Data to be correlated has different dimensions!')
- sys.exit(1)
- y0mean = y0.mean()
- y1mean = y1.mean()
- y0sd = y0.std()
- y1sd = y1.std()
- if y0sd != 0 and y1sd != 0:
- y0norm = (y0 - y0mean) / y0sd
- y1norm = (y1 - y1mean) / y1sd
- else:
- y0norm = y0 - y0mean
- y1norm = y1 - y1mean
- # defined range calculation of cross-correlation
- # value is specified in main routine
- # deltat = 0.9
- pointnumber1 = len(y0)
- ncorrrange = np.ceil(correlationRange / deltat)
- corrrange = np.arange(2 * ncorrrange + 1) - ncorrrange
- ycorr = np.zeros(len(corrrange))
- # print corrrange
- if fast:
- pass
- else:
- for n in corrrange:
- corrpairs = pointnumber1 - abs(n)
- # ccc = arange(corrpairs)
- # print n
- if n < 0:
- y1mod = np.hstack((y1norm[int(-abs(n)):], y1norm[:-int(abs(n))]))
- ycorr[int(n + ncorrrange)] = (np.add.reduce(y0norm * y1mod)) / (float(pointnumber1))
- # if n > -10 :
- # print n, ncorrrange, n+ncorrrange, ycorr[n+ncorrrange], float(pointnumber1)
- elif n == 0:
- ycorr[int(n + ncorrrange)] = (np.add.reduce(y0norm * y1norm)) / (float(pointnumber1))
- # print n, ncorrrange, n+ncorrrange, ycorr[n+ncorrrange], (float(pointnumber1-1))
- elif n > 0:
- y1mod = np.hstack((y1norm[int(abs(n)):], y1norm[:int(abs(n))]))
- ycorr[int(n + ncorrrange)] = (np.add.reduce(y0norm * y1mod)) / (float(pointnumber1))
- # if n < 10 :
- # print n, ncorrrange, n+ncorrrange, ycorr[n+ncorrrange], float(pointnumber1-1)
- else:
- print('Problem!')
- exit(1)
- # print n , ycorr[n+ncorrrange]
- float_corrrange = np.array([float(i) for i in corrrange])
- xcorr = float_corrrange * deltat
- normcorr = np.column_stack((xcorr, ycorr))
- return normcorr
- ############################################################
- ## high-pass filter from http://nullege.com/codes/show/[email hidden]-0.3.3@obspy@[email hidden]
- ############################################################
- def highpass(data, freq, df=200, corners=4, zerophase=False):
- """
- Butterworth-Highpass Filter.
- Filter data removing data below certain frequency freq using corners.
- :param data: Data to filter, type numpy.ndarray.
- :param freq: Filter corner frequency.
- :param df: Sampling rate in Hz; Default 200.
- :param corners: Filter corners. Note: This is twice the value of PITSA's
- filter sections
- :param zerophase: If True, apply filter once forwards and once backwards.
- This results in twice the number of corners but zero phase shift in
- the resulting filtered trace.
- :return: Filtered data.
- """
- fe = 0.5 * df
- [b, a] = iirfilter(corners, freq / fe, btype='highpass', ftype='butter', output='ba')
- if zerophase:
- firstpass = lfilter(b, a, data)
- return lfilter(b, a, firstpass[::-1])[::-1]
- else:
- return lfilter(b, a, data)
- ############################################################
- ## high-pass filter from http://stackoverflow.com/questions/12093594/how-to-implement-band-pass-butterworth-filter-with-scipy-signal-butter
- ############################################################
- def butter_highpass(interval, sampling_rate, cutoff, order=5):
- nyq = sampling_rate * 0.5
- stopfreq = float(cutoff)
- cornerfreq = 0.4 * stopfreq # (?)
- ws = cornerfreq / nyq
- wp = stopfreq / nyq
- # for bandpass:
- # wp = [0.2, 0.5], ws = [0.1, 0.6]
- N, wn = scipy.signal.buttord(wp, ws, 3, 16) # (?)
- # for hardcoded order:
- # N = order
- b, a = scipy.signal.butter(N, wn, btype='high') # should 'high' be here for bandpass?
- sf = scipy.signal.lfilter(b, a, interval)
- return sf
- ##################################################################
- ## high-pass filters the ephys recording and extracts spikes through thresholding
- ##################################################################
- def extractSpikes(eData, eTime, stim=False):
- highpassfreq = 150. # Hz
- spikecountwindow = 0.05 # in sec
- binWidth = 1.E-3 # in sec
- stimRinging = 0.002
- dt = np.mean(eTime[1:] - eTime[:-1])
- rate = 1. / dt
- # set binned array for convolution
- binWidth = 1.E-3 # in sec
- tbins = np.linspace(0., len(eData) * dt, int(len(eData) * dt / binWidth) + 1)
- nspikecountwindow = spikecountwindow / binWidth
- ############################################
- # create new group in hdf5 file
- #grp_spikes = self.analyzed_data.require_group('spiking_data')
- detectSpikes = True
- #if ('spikeTreshold' in grp_spikes.keys()) and ('artifactTreshold' in grp_spikes.keys()):
- # input_ = raw_input('Spike and Artifact detection thresholds exist already. Do you want to re-detect spikes? (\'y\', or any other key for no) : ')
- # if input_ != 'y':
- # detectSpikes = False
- if detectSpikes:
- # get time of the stimuls
- if stim:
- # in case of external stimulation: exclude period of stimuli
- stimuli = self.analyzed_data['stimulation_data/stimulus_times'].value
- startStim = np.array(stimuli / dt, dtype=int)
- endStim = int(stimuli[-1] / dt) + int(stimRinging / dt) # eDataReplaced = copy(eData)
- # high-pass filter recording #################################
- # eDataHP = self.analysisTools.highpass(eData,highpassfreq,rate,corners=4,zerophase=True)
- eDataHP = butter_highpass(eData, rate, highpassfreq, order=4)
- #self.h5pyTools.createOverwriteDS(grp_spikes, 'ephys_data_high-pass', eDataHP)
- # detect spikes ################################################
- app = QtGui.QApplication([])
- win = pg.GraphicsWindow(title="Data plotting")
- win.resize(1800, 600)
- win.setWindowTitle('high-pass filtered recording')
- label = pg.LabelItem(justify='right')
- win.addItem(label)
- pg.setConfigOptions(antialias=True)
- x2 = np.linspace(-100, 100, 1000)
- data2 = np.sin(x2) / x2
- p8 = win.addPlot(row=1, col=0, title="set threshold for spike detection with mouse click")
- p8.plot(eDataHP, pen=(255, 255, 255, 200))
- # lr = pg.LinearRegionItem([400,700])
- # lr.setZValue(-10)
- # p8.addItem(lr)
- vLine = pg.InfiniteLine(angle=90, movable=False)
- hLine = pg.InfiniteLine(angle=0, movable=False)
- hLineSpikes = pg.InfiniteLine(angle=0, pen=pg.mkPen(0, 255, 0), movable=False)
- hLineArtifacts = pg.InfiniteLine(angle=0, pen=pg.mkPen(255, 0, 0), movable=False)
- p8.addItem(vLine, ignoreBounds=True)
- p8.addItem(hLine, ignoreBounds=True)
- p8.addItem(hLineSpikes, ignoreBounds=True)
- p8.addItem(hLineArtifacts, ignoreBounds=True)
- vb = p8.vb
- # detectionTreshold = empty(0)
- def detectSpikeTimes(tresh):
- global detectionTreshold
- excursion = eDataHP < tresh # threshold ephys trace
- excursionInt = np.array(excursion, dtype=int) # convert boolean array into array of zeros and ones
- diff = excursionInt[1:] - excursionInt[:-1] # calculate difference
- spikeStart = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == 1))] # a difference of one is the start of a spike
- spikeEnd = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == -1))] # a difference of -1 is the spike end
- if (spikeEnd[0] - spikeStart[0]) < 0.: # if trace starts below threshold
- spikeEnd = spikeEnd[1:]
- if (spikeEnd[-1] - spikeStart[-1]) < 0.: # if trace ends below threshold
- spikeStart = spikeStart[:-1]
- if len(spikeStart) != len(spikeEnd): # unequal lenght of starts and ends is a problem of course
- print('problem in length of spikeStart and spikeEnd')
- sys.exit(1)
- spikeT = []
- for i in range(len(spikeStart)):
- if (spikeEnd[i] - spikeStart[i]) > 10: # ignore if difference between end and start is smaller than 15 points
- nMin = np.argmin(eDataHP[spikeStart[i]:spikeEnd[i]]) + spikeStart[i]
- spikeT.append(nMin)
- # detectionTreshold = tresh
- return spikeT
- def mouseMoved(evt):
- pos = evt[0] ## using signal proxy turns original arguments into a tuple
- if p8.sceneBoundingRect().contains(pos):
- mousePoint = vb.mapSceneToView(pos)
- index = int(mousePoint.x())
- if index > 0 and index < len(eDataHP):
- label.setText("<span style='font-size: 12pt'>x=%0.1f, <span style='font-size: 12pt'>y=%s</span>" % (mousePoint.x(), eDataHP[index]))
- vLine.setPos(mousePoint.x())
- hLine.setPos(mousePoint.y())
- pointSpikes = [0]
- pointArtifacts = [0]
- sSpikes = pg.ScatterPlotItem(size=10, pen=pg.mkPen(None), brush=pg.mkBrush(0, 255, 0))
- sSpikes.addPoints(x=pointSpikes, y=len(pointSpikes) * [0])
- p8.addItem(sSpikes)
- sArtifacts = pg.ScatterPlotItem(size=10, pen=pg.mkPen(None), brush=pg.mkBrush(255, 0, 0))
- sArtifacts.addPoints(x=pointArtifacts, y=len(pointArtifacts) * [0])
- p8.addItem(sArtifacts)
- def mouseClickedSpikes(evt):
- posClick = evt.pos()
- if p8.sceneBoundingRect().contains(posClick):
- mousePointC = vb.mapSceneToView(posClick)
- hLineSpikes.setPos(mousePointC.y())
- threshold = mousePointC.y()
- pointSpikes = detectSpikeTimes(threshold)
- # print 'spikes clicked:',pointSpikes
- sSpikes.setData(x=pointSpikes, y=eDataHP[pointSpikes])
- def mouseClickedArtifacts(evt):
- posClick = evt.pos()
- if p8.sceneBoundingRect().contains(posClick):
- mousePointC = vb.mapSceneToView(posClick)
- hLineArtifacts.setPos(mousePointC.y())
- # print type(evt)
- threshold = mousePointC.y()
- pointArtifacts = detectSpikeTimes(threshold)
- # sArtifacts.clear()
- sSpikes.setData(x=spikesRaw[0], y=spikesRaw[1])
- sArtifacts.setData(x=pointArtifacts, y=eDataHP[pointArtifacts])
- # first graphical dialog to set spike treshold
- proxy = pg.SignalProxy(p8.scene().sigMouseMoved, rateLimit=60, slot=mouseMoved)
- p8.scene().sigMouseClicked.connect(mouseClickedSpikes)
- pdb.set_trace() # input_ = input("Chose spike treshold in graphical window. Press any button to continue.")
- spikesRaw = sSpikes.getData()
- spikeTreshold = hLineSpikes.getPos()[1] # copy(detectionTreshold)
- # second graphical dialog to set treshold for artifacts
- p8.scene().sigMouseClicked.disconnect(mouseClickedSpikes)
- p8.scene().sigMouseClicked.connect(mouseClickedArtifacts)
- pdb.set_trace()
- # input_ = input("Chose artifact treshold in graphical window. Press any button to continue.")
- falseSpikes = sArtifacts.getData()
- artifactTreshold = hLineArtifacts.getPos()[1]
- while True:
- input_ = input("Enter pairs of indicies of regions to excluce from spike detection (e.g. [[0,700],[5450,5560]]). Press any number if None. : ")
- try:
- aaa = len(input_)
- except:
- print('No regions to exclude specified.')
- exclusionBorders = None
- break
- else:
- print('recorded')
- exclusionBorders = input_
- break
- while True:
- 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. : ")
- try:
- type(input2_)
- except:
- print('No regions to exclude specified.')
- artifactLength = None
- break
- else:
- artifactLength = input2_
- break
- #
- # pdb.set_trace()
- lspikes = spikesRaw[0].tolist()
- lartif = falseSpikes[0].tolist()
- spikes0 = [x for x in lspikes if x not in lartif]
- # add spike artifacts to regions to remove
- if artifactLength:
- if exclusionBorders == None:
- exclusionBorders = []
- if stim:
- for i in range(len(startStim)):
- exclusionBorders.append([startStim[i], startStim[i] + artifactLength])
- # remove spikes which fall in to regions to exclude
- if exclusionBorders:
- spikes1 = list(spikes0)
- for n in range(len(exclusionBorders)):
- spikes1 = [x for x in spikes1 if not (x > exclusionBorders[n][0] and x < exclusionBorders[n][1])]
- spikeTimes = eTime.value[np.array(spikes1, dtype=int)]
- #firingRate = brian.firing_rate(spikeTimes)
- #cv = brian.CV(spikeTimes)
- # pdb.set_trace()
- ######################################################
- # convolv original spike trains with Gaussian kernels
- binnedspikes, _ = np.histogram(spikeTimes, tbins)
- spikesconv = scipy.ndimage.filters.gaussian_filter1d(np.array(binnedspikes, float), sigma=nspikecountwindow)
- # convert the convolved spike trains to units of spikes/sec
- spikesconv *= 1. / binWidth
- # save data
- #self.h5pyTools.createOverwriteDS(grp_spikes, 'spikeTreshold', array([spikeTreshold]))
- #self.h5pyTools.createOverwriteDS(grp_spikes, 'artifactTreshold', array([artifactTreshold]))
- #self.h5pyTools.createOverwriteDS(grp_spikes, 'spikes', spikeTimes)
- #self.h5pyTools.createOverwriteDS(grp_spikes, 'firing_rate_evolution', spikesconv, ['dt', binWidth])
- #self.h5pyTools.createOverwriteDS(grp_spikes, 'firing_rate', array([firingRate]))
- #self.h5pyTools.createOverwriteDS(grp_spikes, 'CV', array([cv]))
- def NormalizeData(data):
- return (data - np.min(data)) / (np.max(data) - np.min(data))
- #################################################################################
- # detect spikes in ephys trace
- #################################################################################
- def detectSpikeTimes(tresh,eDataHP,ephysTimes,positive=True,plot=False):
- #global detectionTreshold
- while True:
- if positive :
- excursion = eDataHP > tresh # threshold ephys trace
- else:
- excursion = eDataHP < tresh
- excursionInt = np.array(excursion, dtype=int) # convert boolean array into array of zeros and ones
- diff = excursionInt[1:] - excursionInt[:-1] # calculate difference
- spikeStart = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == 1))] # a difference of one is the start of a spike
- spikeEnd = np.arange(len(eDataHP))[np.concatenate((np.array([False]), diff == -1))] # a difference of -1 is the spike end
- if len(spikeEnd)>0 and len(spikeStart)>0:
- if (spikeEnd[0] - spikeStart[0]) < 0.: # if trace starts below threshold
- spikeEnd = spikeEnd[1:]
- if (spikeEnd[-1] - spikeStart[-1]) < 0.: # if trace ends below threshold
- spikeStart = spikeStart[:-1]
- if len(spikeStart) != len(spikeEnd): # unequal lenght of starts and ends is a problem of course
- print('problem in length of spikeStart and spikeEnd')
- sys.exit(1)
- spikeT = []
- spikeStart = spikeStart[spikeStart>100]
- #for i in range(len(spikeStart)):
- # #if (spikeEnd[i] - spikeStart[i]) > 10: # ignore if difference between end and start is smaller than 15 points
- # nMin = spikeStart[i] #np.argmin(eDataHP[spikeStart[i]:spikeEnd[i]]) + spikeStart[i]
- # spikeT.append(nMin)
- # detectionTreshold = tresh
- #pdb.set_trace()
- spikeTimes = ephysTimes[spikeStart]
- if plot:
- fig = plt.figure(figsize=(12,8))
- ax = fig.add_subplot(111)
- ax.plot(ephysTimes,eDataHP)
- ax.plot(ephysTimes[spikeStart], eDataHP[spikeStart],'.')
- ax.axhline(y=tresh, ls='--', c='0.5')
- plt.show()
- if plot:
- print('Is threshold ok? ->No : type new threshold value ->Yes : press Enter')
- recInput = input()
- if recInput == "":
- break
- else:
- tresh = float(recInput)
- print('new threshold : ', tresh)
- else:
- break
- #recInputIdx = [int(i) for i in recInput.split(',')]
- return (spikeTimes,spikeStart,tresh)
- #################################################################################
- # maps an abritray input array to the entire range of X-bit encoding
- #################################################################################
- def mapToXbit(inputArray,xBitEncoding):
- oldMin = np.min(inputArray)
- oldMax = np.max(inputArray)
- newMin = 0.
- newMax = 2**xBitEncoding-1.
- normXBit = newMin + (inputArray - oldMin) * newMax / (oldMax - oldMin)
- normXBitInt = np.array(normXBit, dtype=int)
- return normXBitInt
- #################################################################################
- # maps an abritray input array to the entire range of X-bit encoding
- #################################################################################
- # dataAnalysis.determineFrameTimes(exposureArray[0],arrayTimes,frames)
- def determineFrameTimes(exposureArray,arrayTimes,frames,rec=None):
- display = False
- #pdb.set_trace()
- numberOfFrames = len(frames)
- exposure = exposureArray > 20 # threshold trace
- exposureInt = np.array(exposure, dtype=int) # convert boolean array into array of zeros and ones
- difference = np.diff(exposureInt) # calculate difference
- expStart = np.arange(len(exposureArray))[np.concatenate((np.array([False]), difference == 1))] # a difference of one is the start of a spike
- expEnd = np.arange(len(exposureArray))[np.concatenate((np.array([False]), difference == -1))] # a difference of -1 is the spike end
- if (expEnd[0] - expStart[0]) < 0.: # if trace starts above threshold
- expEnd = expEnd[1:]
- if (expEnd[-1] - expStart[-1]) < 0.: # if trace ends above threshold
- expStart = expStart[:-1]
- frameDuration = expEnd - expStart
- midExposure = (expStart + expEnd)/2
- expStartTime = arrayTimes[expStart.astype(int)]
- expEndTime = arrayTimes[expEnd.astype(int)]
- #framesIdxDuringRec = np.array(len(softFrameTimes))[(arrayTimes[expEnd[0]]+0.002) < softFrameTimes]
- #framesIdxDuringRec = framesIdxDuringRec[:len(expStart)]
- if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
- recordedFrames = frames[3:(len(midExposure) + 3)]
- elif arrayTimes[int(midExposure[0])]<0.003:
- recordedFrames = frames[2:(len(midExposure) + 2)]
- else:
- recordedFrames = frames[:len(midExposure)]
- print('number of tot. frames, recorded frames, exposures start, end :',numberOfFrames,len(recordedFrames), len(expStart), len(expEnd))
- if display:
- ledON = np.zeros(len(exposureArray))
- for i in range(11):
- ledON[((i*1.)<=arrayTimes) & ((i*1.+0.2)>arrayTimes)] = 1.
- ledON[29.<=arrayTimes] = 1.
- data = np.loadtxt('/home/mgraupe/2019.04.01_000-%s.csv' % (rec[-3:]),delimiter=',',skiprows=1,usecols=(0,1))
- print(len(data))
- plt.plot(arrayTimes,exposureArray/32.)
- plt.plot(arrayTimes,ledON)
- print('fist frame at %s sec' % arrayTimes[int(midExposure[0])],end='')
- if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
- plt.plot(arrayTimes[midExposure.astype(int)], (data[3:(len(midExposure) + 3), 1] - 148.6) / 105.4, 'o-')
- print(3)
- elif arrayTimes[int(midExposure[0])]<0.003:
- plt.plot(arrayTimes[midExposure.astype(int)], (data[2:(len(midExposure) + 2), 1] - 148.6) / 105.4, 'o-')
- print(2)
- else:
- plt.plot(arrayTimes[midExposure.astype(int)], (data[:len(midExposure), 1] - 148.6) / 105.4, 'o-')
- print(0)
- #plt.plot(softFrameTimes[data[:-6,0].astype(np.int)+6],(data[:-6,1]-148.6)/105.4)
- #plt.plot(softFrameTimes,np.ones(len(softFrameTimes)),'|')
- plt.show()
- pdb.set_trace()
- return (expStartTime,expEndTime,recordedFrames)
- #################################################################################
- def generatePlotWithSTD(data,std=[2,3,4],names = None):
- matplotlib.use('TkAgg')
- nData = len(data)
- fig = plt.figure(figsize=(15,15))
- for i in range(nData):
- ax0 = fig.add_subplot(nData,1,i+1)
- if names is not None:
- ax0.set_title('%s' % names[i])
- STD = np.std(data[i])
- MM = np.mean(data[i])
- for n in range(len(std)):
- ax0.axhline(y=MM+std[n]*STD,ls='--',c=plt.cm.RdYlBu(n/len(std)),label='%s STD' % std[n])
- ax0.axhline(y=MM-std[n]*STD, ls='--', c=plt.cm.RdYlBu(n/len(std)))
- ax0.axhline(y=MM,c='C1',label='mean')
- ax0.plot(data[i],c='C0')
- if i==2:
- ax0.plot(np.abs(data[i]), c='C4')
- ax0.legend()
- plt.show()
- #################################################################################
- # maps an abritray input array to the entire range of X-bit encoding
- #################################################################################
- def determineFramesToExclude(frames,probIdx):
- listOfFramesToExclude = []
- canBeUsed = True
- # first let's decide on how many LED's (if any) are present in the FOV
- for i in range(len(probIdx)):
- currIdx = probIdx[i]
- continueDetectLoop = True
- while continueDetectLoop:
- print('checking idx, started at :', currIdx, probIdx[i])
- #frame8bit = np.array(np.transpose(frames[currIdx]), dtype=np.uint8)
- img = cv2.cvtColor(frames[currIdx], cv2.COLOR_GRAY2BGR)
- # rungs = []
- #imgPure = img.copy()
- cv2.imshow("PureImage", img)
- 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 :')
- PressedKey = cv2.waitKey(0)
- print(PressedKey)
- if PressedKey == 81: # left arrow key
- currIdx -=1
- elif PressedKey == 83: # right arrow key
- currIdx +=1
- elif PressedKey == 101: # y key
- print('%s added to exclude list' %currIdx)
- listOfFramesToExclude.append(currIdx)
- elif PressedKey == 114: # e key
- print('%s removed from exclude list' % currIdx)
- listOfFramesToExclude.remove(currIdx)
- elif PressedKey == 111 : # o key
- nIdx = input('specify a new idx to check :')
- currIdx = int(nIdx)
- elif PressedKey == 102: # f key
- continueDetectLoop = False
- elif PressedKey == 120: # x key
- canBeUsed = False
- break
- else:
- print('Key not recognized, try again.')
- print('current exclude list :',listOfFramesToExclude)
- if not canBeUsed:
- break
- cv2.destroyWindow("PureImage") # only destroy window at the end of the exploration
- lofEx = list(dict.fromkeys(listOfFramesToExclude)) # removes duplicates
- lofEx.sort()
- print('starting list of indexes :', probIdx)
- print('indexes to exclude :', lofEx)
- lofEx = np.asarray(lofEx,dtype=int)
- #pdb.set_trace()
- return (lofEx, canBeUsed)
- #################################################################################
- # maps an abritray input array to the entire range of X-bit encoding
- #################################################################################
- # ([ledTraces,ledCoordinates,frames,softFrameTimes,imageMetaInfo],[exposureDAQArray,exposureDAQArrayTimes],[ledDAQControlArray, ledDAQControlArrayTimes],verbose=True)
- def determineErroneousFrames(frames):
- # first threshold metrics of the movie to detect and exclude erroneous frames with horizontal lines, flash-back frames #########################
- frameDiff = []
- lineDiff = []
- print('calculating frame and line diffs ... ',end='')
- for i in range(len(frames)):
- if i>0:
- frameDiffAllPix = cv2.absdiff(frames[i],frames[i-1])
- fD = np.average(frameDiffAllPix)
- frameDiff.append(fD)
- lineDiffAllLines = cv2.absdiff(frames[i][:,1:],frames[i][:,:-1])
- lD = np.average(lineDiffAllLines,axis=0)
- lineDiff.append(lD)
- #pdb.set_trace()
- print('done!')
- frameDiff = np.asarray(frameDiff)
- lineDiff = np.asarray(lineDiff)
- lineDiffSum = np.sum(lineDiff,axis=1)
- frameDiffDiff = np.diff(frameDiff)
- generatePlotWithSTD([lineDiffSum,frameDiff,frameDiffDiff],std=[3,3.5,4],names=['lineDiffSum','frameDiff','diff of FrameDiff'])
- # trick to display the above image
- #frame8bit = np.array(np.transpose(frames[0]), dtype=np.uint8)
- img = cv2.cvtColor(frames[0], cv2.COLOR_GRAY2BGR)
- cv2.imshow('HoldImage',img)
- cv2.waitKey(0) #cv2.imshow()
- cv2.destroyWindow('HoldImage')
- 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 : ")
- threshold = [float(i) for i in thresholdingInput.split()]
- print('choice :', threshold)
- #pdb.set_trace()
- if threshold[0] == 5.:
- idxToExclude = np.array([], dtype=np.int64)
- canBeUsed = True
- plt.close('all')
- return(idxToExclude,canBeUsed)
- elif threshold[0] == 1.:
- thresholded = lineDiffSum > np.mean(lineDiffSum) + np.std(lineDiffSum)*threshold[1]
- outlierIdx = np.arange(len(lineDiffSum))[thresholded] # use indices taking into account missed frames
- elif threshold[0] == 2.:
- thresholded = frameDiff > np.mean(frameDiff) + np.std(frameDiff) * threshold[1]
- outlierIdx = np.arange(len(frameDiff))[thresholded] # use indices taking into account missed frames
- 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
- elif threshold[0] == 3.:
- thresholded = np.abs(frameDiffDiff) > np.mean(frameDiffDiff) + np.std(frameDiffDiff)*threshold[1]
- outlierIdx = np.arange(len(frameDiffDiff))[thresholded] # use indices taking into account missed frames
- 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
- elif threshold[0] == 4.:
- canBeUsed = False
- idxToExclude = np.array([], dtype=np.int64)
- plt.close('all')
- return (idxToExclude, canBeUsed)
- print('length and identity of possible erronous frames :' , len(outlierIdx), outlierIdx)
- (idxExclude,canBeUsed) = determineFramesToExclude(frames,outlierIdx)
- #excludeMask = np.ones(len(ledVideoRoi[2]),dtype=bool)
- # add indicies for equivalent frames
- sameFrames = (frameDiff == 0)
- sameFrameIdx = np.arange(len(frameDiff))[sameFrames]
- sameFrameIdx += 1
- print('same frames were recorded here :', sameFrameIdx)
- idxToExclude = np.sort(np.concatenate((sameFrameIdx, idxExclude)))
- #pdb.set_trace()
- #excludeMask[idxToExclude] = False
- plt.close('all')
- return (idxToExclude,canBeUsed)
- #################################################################################
- # maps an abritray input array to the entire range of X-bit encoding
- #################################################################################
- # ([ledTraces,ledCoordinates,frames,softFrameTimes,imageMetaInfo,idxToExclude],[exposureDAQArray,exposureDAQArrayTimes],[ledDAQControlArray, ledDAQControlArrayTimes],verbose=True)
- def determineFrameTimesBasedOnLED(ledVideoRoi, cameraExposure, ledDAQc, pc, verbose=False, tail=False,manualThreshold=False):
- ##############################################################################################################
- # auxiliary function to convert bimodal trace into boolean array
- def traceToBinary(trace,threshold=None):
- rescaledTrace = (trace - np.min(trace)) / (np.max(trace) - np.min(trace))
- if threshold is None:
- rescaledTraceBin = rescaledTrace > 0.3
- else:
- rescaledTraceBin = rescaledTrace > threshold
- return (rescaledTrace,rescaledTraceBin)
- ##############################################################################################################
- def traceToBinaryForChangingMaxMin(trace,threshold=None):
- maxTrace = ndimage.maximum_filter(trace, size=5*2)
- minTrace = ndimage.minimum_filter(trace, size=5*2)
- rescaledTrace = (trace - minTrace) / (maxTrace - minTrace)
- if threshold is None:
- rescaledTraceBin = rescaledTrace > 0.3
- else:
- rescaledTraceBin = rescaledTrace > threshold
- return (rescaledTrace,rescaledTraceBin)
- ##############################################################################################################
- # maps LED daq control trace to boolean array ################################################################
- # TODO this number is zero on the behavior setup and 4 here
- if pc == 'behaviorPC':
- LEDcontrolIdx = 0 # which trace of the DAQ recording is linked to the !!!
- elif pc == '2photonPC':
- LEDcontrolIdx = 4
- else:
- print('Make sure the computer of the recording is specified.')
- ledDAQcontrolBin = traceToBinary(ledDAQc[0][LEDcontrolIdx])[1] # here the threshold is not important as the trace is binary to start out with
- # convert LED roi traces from video to boolean arrays
- ledVideoRoiBins = []
- ledVideoRoiRescaled = []
- allLEDVideoRoiValues = []
- # determine threshold [ledTraces,ledCoordinates,frames,softFrameTimes,imageMetaInfo,idxToExclude]
- # tail covering the LEDs for some
- if tail:
- matplotlib.use('TkAgg')
- print(' in tail ...')
- anticipateCorrectValues = True
- for i in range(ledVideoRoi[1][0]):
- plt.plot(ledVideoRoi[0][i],'o-',ms=2,label='%s' % i)
- plt.legend(loc=1)
- plt.show()
- if anticipateCorrectValues:
- inputA = input('Index until which the recording is not affected by the tail (integer; type 0 if recording is ok) :')
- #inputA=350
- untilOKidx = int(inputA)
- if untilOKidx != 0:
- period = [7, 7, 7, 5]
- for i in range(ledVideoRoi[1][0]):
- maxVal = np.max(ledVideoRoi[0][i][20:untilOKidx])
- minVal = np.min(ledVideoRoi[0][i][20:untilOKidx])
- for n in range(period[i]):
- # repeatValue(ledVideoRoi[0][i][(untilOKidx+n):],7)
- isHigh = [True if abs(ledVideoRoi[0][i][(untilOKidx + n)] - maxVal) < abs(ledVideoRoi[0][i][(untilOKidx + n)] - minVal) else False]
- if isHigh:
- ledVideoRoi[0][i][(untilOKidx + n):][::period[i]] = ledVideoRoi[0][i][(untilOKidx + n)]
- else:
- ledVideoRoi[0][i][(untilOKidx + n):][::period[i]] = ledVideoRoi[0][i][(untilOKidx + n)] # ledVideoRoi[0][i][]
- else:
- maxV = [254,251,250,213]
- minV = [200,174,217,147]
- idxMaxV = [[8580,8582,8583,8585],
- [],
- [8589],
- []]
- idxMinV = [[8581,8584,8586,8588],
- [8579,8580],
- [8590,8591],
- [8584,8585,8586,8589,8590,8591]]
- for i in range(4):
- for n in idxMaxV[i]:
- ledVideoRoi[0][i][n] = maxV[i]
- for m in idxMinV[i]:
- ledVideoRoi[0][i][m] = minV[i]
- # fig = plt.figure()
- # for i in range(ledVideoRoi[1][0]):
- # #ax = fig.add_subplot(3,1,i)
- # plt.plot(ledVideoRoi[0][i],'o-',ms=2,label='%s' % i)
- # #ax.set_xlim(8)
- # plt.legend(loc=1)
- # plt.show()
- # pdb.set_trace()
- ###########
- for i in range(ledVideoRoi[1][0]):
- allLEDVideoRoiValues.extend(traceToBinaryForChangingMaxMin(ledVideoRoi[0][i])[0]) # rescale all values to [0,1] and stack them
- allLEDVideoRoiValues = np.sort(np.asarray(allLEDVideoRoiValues)) # convert to array and sort
- luminocityDifferences = np.diff(allLEDVideoRoiValues)
- idxMaxDiff = np.argmax(luminocityDifferences)
- LEDVideoThreshold = allLEDVideoRoiValues[idxMaxDiff] + (allLEDVideoRoiValues[idxMaxDiff+1] - allLEDVideoRoiValues[idxMaxDiff])/2.
- if pc == 'behaviorPC':
- #illumLEDcontrolThreshold = LEDVideoThreshold**4.49185827 # mapping, i.e. exponent, from tools/fitOfIlluminationValues
- illumLEDcontrolThreshold = LEDVideoThreshold**18.37008924
- elif pc == '2photonPC':
- illumLEDcontrolThreshold = LEDVideoThreshold**2.61290794 # 2pinvivo
- # if (illumLEDcontrolThreshold) < 0.08 or manualThreshold:
- # print('thresholds before: ',LEDVideoThreshold, illumLEDcontrolThreshold)
- # #print('LED threshold extremly low! Fixed by setting both threshold to 0.8.')
- # fig = plt.figure(figsize=(10,10))
- # ax = fig.add_subplot(111)
- # #ax.plot(np.ones(len(allLEDVideoRoiValues)),allLEDVideoRoiValues,'.',ms=0.5)
- # ax.axhline(y=LEDVideoThreshold,ls='--',c='C0')
- # ax.plot(allLEDVideoRoiValues,'.',ms=0.5,c='C0')
- # #plt.plot(np.ones(len(ledDAQcontrolBin)),ledDAQcontrolBin,'.')
- # #ax.plot(ledDAQcontrolBin,'.',ms=0.5)
- # plt.show()
- # # thresholdInput = '0.8,0.8'
- # # thresholdInput = ''
- # thresholdInput = input('Provide alternative thresholds (e.g. 0.8,0.7), otherwise press enter to keep current thresholds : ')
- # if not thresholdInput == '':
- # newThresholds = [float(i) for i in thresholdInput.split(',')]
- # LEDVideoThreshold = newThresholds[0]
- # illumLEDcontrolThreshold = newThresholds[1]
- # find start and end of camera exposure period ################################################################
- exposureInt = np.array(cameraExposure[0][0], dtype=int) # convert boolean array into array of zeros and ones
- difference = np.diff(exposureInt) # calculate difference
- expStart = np.arange(len(exposureInt))[np.concatenate((np.array([False]), difference > 0))] # a difference of one is the start of the exposure
- expEnd = np.arange(len(exposureInt))[np.concatenate((np.array([False]), difference < 0))] # a difference of -1 is the end the exposure period
- #pdb.set_trace()
- if (expEnd[0] - expStart[0]) < 0.: # if trace starts above threshold
- print('exposure at start of recording')
- expEnd = expEnd[1:]
- exposureAtStart = True
- else:
- exposureAtStart = False
- if (expEnd[-1] - expStart[-1]) < 0.: # if trace ends above threshold
- print('exposure during end of recording')
- exposureAtEnd = True
- expStart = expStart[:-1]
- else:
- exposureAtEnd = False
- expStart = expStart.astype(int)
- expEnd = expEnd.astype(int)
- expStartTime = cameraExposure[1][expStart] # everything was based on indicies up to this point : here indicies -> time
- expEndTime = cameraExposure[1][expEnd] # everything was based on indicies up to this point : here indicies -> time
- frameDuration = expEndTime - expStartTime
- print('first frame started at ', expStartTime[0]*1000., 'ms' )
- ## based on exposure start-stop, how bright should the DAQ LED signal be ##########################################
- startEndExposureTime = np.column_stack((expStartTime, expEndTime))
- startEndExposurepIdx = np.column_stack((expStart,expEnd)) # create a 2-column array with 1st column containing start and 2nd column containing end index
- illumLEDcontrol = [np.mean(ledDAQcontrolBin[b[0]:b[1]]) for b in startEndExposurepIdx] # extract MEAN illumination value - from LED control trace - during exposure period
- illumLEDcontrol = np.asarray(illumLEDcontrol)
- adjusted = False
- if manualThreshold:
- while True:
- if (illumLEDcontrolThreshold) < 0.08 or manualThreshold:
- sortedIllumLEDcontrol = np.sort(illumLEDcontrol)
- fig = plt.figure(figsize=(10,10))
- ax = fig.add_subplot(111)
- #ax.plot(np.ones(len(allLEDVideoRoiValues)),allLEDVideoRoiValues,'.',ms=0.5)
- print('thresholds before: ',LEDVideoThreshold, illumLEDcontrolThreshold)
- if (illumLEDcontrolThreshold < 0.08) and not (adjusted):
- illumLEDcontrolThreshold = 0.2
- adjusted = True # do it only once
- print('adjusted thresholds: ', LEDVideoThreshold, illumLEDcontrolThreshold)
- print('video (up, down, down fraction) : ', np.sum(sortedIllumLEDcontrol > LEDVideoThreshold), np.sum(sortedIllumLEDcontrol < LEDVideoThreshold), np.sum(sortedIllumLEDcontrol < LEDVideoThreshold)/len(sortedIllumLEDcontrol))
- print('illum (up, down, down fraction): ', np.sum(allLEDVideoRoiValues>illumLEDcontrolThreshold), np.sum(allLEDVideoRoiValues<illumLEDcontrolThreshold), np.sum(allLEDVideoRoiValues<illumLEDcontrolThreshold)/len(allLEDVideoRoiValues))
- ax.axhline(y=illumLEDcontrolThreshold,ls='--',c='C0')
- ax.plot(np.linspace(0,1,len(sortedIllumLEDcontrol)),sortedIllumLEDcontrol,'.',ms=0.5,c='C0')
- ax.axhline(y=LEDVideoThreshold, ls='--', c='C1')
- ax.plot(np.linspace(0,1,len(allLEDVideoRoiValues)),allLEDVideoRoiValues, '.', ms=0.5,c='C1')
- #plt.plot(np.ones(len(ledDAQcontrolBin)),ledDAQcontrolBin,'.')
- #ax.plot(ledDAQcontrolBin,'.',ms=0.5)
- plt.show()
- # thresholdInput = '0.8,0.8'
- # thresholdInput = ''
- thresholdInput = input('Provide alternative thresholds (e.g. 0.8,0.2), otherwise press enter exit loop : ')
- if not thresholdInput == '':
- newThresholds = [float(i) for i in thresholdInput.split(',')]
- LEDVideoThreshold = newThresholds[0]
- illumLEDcontrolThreshold = newThresholds[1]
- else:
- break
- else:
- LEDVideoThreshold = LEDVideoThreshold
- illumLEDcontrolThreshold = illumLEDcontrolThreshold
- # print('press any key to check/redefine thresholds; exit loop with space or enter:')
- # PressedKey = cv2.waitKey(0)
- # if PressedKey == 13 or PressedKey == 32: # Enter or Space
- # break
- # else:
- # pass
- print('thresholds final: ',LEDVideoThreshold, illumLEDcontrolThreshold)
- # pdb.set_trace()
- # LEDVideoThreshold = 0.8
- # illumLEDcontrolThreshold = 0.8
- #print('adjusted thresholds : ', LEDVideoThreshold, illumLEDcontrolThreshold)
- (illumLEDcontrolrescaled, illumLEDcontrolBin) = traceToBinary(illumLEDcontrol, threshold=illumLEDcontrolThreshold) # 0.2 and 0.15 before
- # pdb.set_trace()
- # threshold and convert to binary
- for i in range(ledVideoRoi[1][0]):
- ledVideoRoiBins.append(traceToBinary(ledVideoRoi[0][i],threshold=LEDVideoThreshold)[1]) # 0.6 before 0.4
- ledVideoRoiRescaled.append(traceToBinary(ledVideoRoi[0][i])[0])
- #plt.plot(allLEDVideoRoiValues,illumLEDcontrol, '.', ms=0.5)
- #plt.show()
- #plt.plot(illumLEDcontrol)
- ## loop over frame numbers and extract binary number shown by leds ################################################
- nFrames = len(ledVideoRoiBins[0])
- 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]])
- recordedFrames = 0
- frameCount = []
- binNumberInFrame = np.column_stack((ledVideoRoiBins[0],ledVideoRoiBins[1],ledVideoRoiBins[2]))
- frameNBefore = 0
- oldI = -1
- exceptionsInFrameCount = []
- idxToExclude = ledVideoRoi[5]
- for i in range(nFrames):
- if i not in idxToExclude:
- 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
- matchFrameN = np.arange(len(binNumbers))[matchBool][0] # converts the boolean list into the index corresponding to the match
- frameDiff = matchFrameN - frameNBefore # difference in count to previous frame
- if frameDiff < 0: # else : negative difference indicates that the counter restarted
- frameDiff+=7
- if (frameDiff != 1) and (frameDiff != -6):
- print(i,oldI,i-oldI,matchFrameN,frameNBefore,frameDiff,binNumberInFrame[i],binNumberInFrame[i-1])
- exceptionsInFrameCount.append([i,oldI,i-oldI,matchFrameN,frameNBefore,frameDiff,binNumberInFrame[i],binNumberInFrame[i-1]])
- if matchFrameN == 0: # counter will start at 0 and possibly go back to zero after end of recording
- if (i>70) and (i<(nFrames-10)):
- print(i,oldI,matchFrameN,frameNBefore,frameDiff,binNumberInFrame[i],binNumberInFrame[i-1])
- print('strange, zero frame in the middle of recording')
- pdb.set_trace()
- #frameDiff = 0
- frameCount.append([i,matchFrameN,frameDiff,int(ledVideoRoiBins[3][i]),oldI])
- frameNBefore = matchFrameN
- oldI = i
- frameCount = np.asarray(frameCount,dtype=int) # convert list to integer array
- idxRecordedFrames = np.cumsum(frameCount[:,2]) # use the frame differences to generate new index corresponding to video recording
- idxCounting = np.argwhere(idxRecordedFrames>0) # start and end index with first and last frame recording the counter
- 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
- #pdb.set_trace()
- if exposureAtStart: # remove first frame if exposure was active during start of recording, i.e., at t = 0 s
- idxFramesDuringRecording = idxFramesDuringRecording[1:] - 1
- if exposureAtEnd:
- idxFramesDuringRecording = idxFramesDuringRecording[:-1]
- #pdb.set_trace()
- idxMissingFrames = np.delete(np.arange(idxFramesDuringRecording[-1]+1),idxFramesDuringRecording)
- #idxTestMask = idxFramesDuringRecording < len(illumLEDcontrolBin) # index should not exceed length of array
- #illum = illumLEDcontrolBin[idxFramesDuringRecording[idxTestMask]]
- ## the excluded frames - based on distortions - need to be removed from the video sequence
- videoRoi = ledVideoRoiBins[3]
- mask = np.ones(len(videoRoi),dtype=bool)
- mask[idxToExclude] = False
- #pdb.set_trace()
- ## first loop to align the START of the frame recording - in illumLEDcontrolBin - with the video recording
- if any(idxToExclude < 20):
- print('Early frames to exclude. Problem!')
- pdb.set_trace()
- else:
- tmpIdx = np.where(videoRoi==True) # Index of the first ON frame for the 4th LED
- idxFirstFrameRec = tmpIdx[0][0] # extract index of first frame during recording
- if exposureAtStart:
- idxFirstFrameRec+=1 # increase that
- for j in range(70):
- videoRoiWOEX = videoRoi[mask][j:]
- if np.all(videoRoiWOEX[:20] == illumLEDcontrolBin[:20]): # note that illumLEDcontrolBin already accounts for a recording during stat of rec, this frame is removed
- missedFramesBegin = j
- break
- try:
- a=missedFramesBegin
- #a = lllll
- except:
- print("bad alignement !!!!!!!!!!!!!!!!!!!")
- print(videoRoi[mask][:20],illumLEDcontrolBin[:20])
- plt.plot(illumLEDcontrolrescaled,'.',ms=0.5)
- plt.plot(ledVideoRoiRescaled[3],'.',ms=0.5)
- plt.show()
- pdb.set_trace()
- if (idxFirstFrameRec == missedFramesBegin) or (missedFramesBegin == 0):
- print('Number of frames recorded before first full exposed frame during recording :', missedFramesBegin, idxFirstFrameRec, illumLEDcontrolBin[:20],videoRoi[:20] )
- videoRoiWOEX = videoRoi[mask][missedFramesBegin:]
- elif idxFirstFrameRec == (missedFramesBegin+1):
- missedFramesBegin+=1
- print('Number of frames recorded before first full exposed frame during recording (increased by one):', missedFramesBegin, idxFirstFrameRec, illumLEDcontrolBin[:20],videoRoi[:20])
- videoRoiWOEX = videoRoi[mask][missedFramesBegin:]
- else:
- print('Number of frames recorded before first full exposed frame during recording :', missedFramesBegin, idxFirstFrameRec, illumLEDcontrolBin[:20],videoRoi[:20] )
- print('Problem with determining index of first recorded frame.')
- pdb.set_trace()
- #pdb.set_trace()
- ## second loop in order to align the
- shiftDifference = []
- lengthOfIllumLEDcontrol = len(illumLEDcontrolBin)
- lengthOfROIinVideo = len(videoRoiWOEX)
- lengthOfIdxCount = idxFramesDuringRecording[-1] + 1
- print('length of illumLEDcontrolBin and videoRoiWOEX and idxFramesDuringRecording[-1] : ', lengthOfIllumLEDcontrol, lengthOfROIinVideo, lengthOfIdxCount)
- for i in range(-10,11,1): # loop to shift the mask over
- idxTemp = idxFramesDuringRecording + i # shift the array by increasing the indicies by a certain number
- 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
- #idxIllum = idx[idx<lengthOfIllumLEDcontrol] # indicies should not be larger than the length of the illumLEDcontrolBin array
- illum = illumLEDcontrolBin[idxIllum] # illumination at these indicies
- # idxMissing = np.delete(np.arange(idxIllum[-1]), idxIllum) #[i:]
- # idxMissing = np.delete(np.arange(idxFramesDuringRecording[-1]), idxIllum) # [i:]
- idxMissing = np.delete(np.arange(lengthOfIllumLEDcontrol), idxIllum)
- NidxRemovedAtExtremities = idxIllum[0] + ((lengthOfIllumLEDcontrol-1) - idxIllum[-1]) # counts number of frames missing in the beginning and end
- NidxRemovedAtExtremities -= np.sum((idxMissingFrames<idxIllum[0]) | (idxMissingFrames>idxIllum[-1])) # reduce if missing frames are in the extrimities
- #if i<0:
- # frameOverlap = [0 if ((lengthOfIllumLEDcontrol+np.abs(i)+1)<(lengthOfROIinVideo+len(idxMissingFrames))) else ((lengthOfROIinVideo+len(idxMissingFrames)) - (lengthOfIllumLEDcontrol+np.abs(i)+1))]
- #elif i>=0:
- # frameOverlap = [i if (lengthOfIllumLEDcontrol<(lengthOfROIinVideo+len(idxMissingFrames))) else (lengthOfIllumLEDcontrol-(lengthOfROIinVideo+len(idxMissingFrames)+i+1))]
- #print('overlap :',frameOverlap[0])
- #print('test',len(np.intersect1d(idxMissing,idxMissingFrames)), idxMissingFrames, idxMissing,(-len(idxMissing)),NidxRemovedAtExtremities)
- compareIdx = len(np.intersect1d(idxMissing,idxMissingFrames)) - len(idxMissing) + np.abs(NidxRemovedAtExtremities) # abs(i)
- len0 = len(illum)
- len1 = lengthOfROIinVideo
- if len0 < len1:
- compare = np.sum(np.equal(illum,videoRoiWOEX[:len0]))
- versch = compare - len0
- totLength = len0
- largeLength = len1
- else:
- compare = np.sum(np.equal(illum[:len1],videoRoiWOEX))
- versch = compare - len1
- totLength = len1
- largeLength = len0
- #pdb.set_trace()
- shiftDifference.append([i, versch, totLength, compareIdx,NidxRemovedAtExtremities])
- print(i, versch, totLength, compareIdx, NidxRemovedAtExtremities, idxMissing, idxMissingFrames)
- #if i >=0 :
- #pdb.set_trace()
- #compare = np.equal()
- shiftDifference = np.asarray(shiftDifference)
- if len(idxToExclude)==0: # without erronous frames ...
- shiftDifference = shiftDifference[shiftDifference[:,0]>=0] # ... use only the shifts which are larger than zero
- shiftToZero = shiftDifference[:,0][(shiftDifference[:,1]==0) & (shiftDifference[:,3]==0)]
- finalLength = shiftDifference[:,2][(shiftDifference[:,1]==0) & (shiftDifference[:,3]==0)]
- if len(shiftToZero)>1 or len(shiftToZero)==0:
- if len(shiftToZero)==0:
- idxTemp = idxFramesDuringRecording + 0
- idx = idxTemp[idxTemp >= 0]
- idxIllum = idx[idx < len(illumLEDcontrolBin)]
- # pdb.set_trace()
- shortest = [len(ledVideoRoiRescaled[3][mask][missedFramesBegin:]) if len(ledVideoRoiRescaled[3][mask][missedFramesBegin:]) < len(illumLEDcontrolrescaled[idxIllum]) else len(
- illumLEDcontrolrescaled[idxIllum])]
- plt.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]], illumLEDcontrolrescaled[idxIllum][:shortest[0]], 'o', ms=1)
- plt.show()
- pdb.set_trace()
- fig = plt.figure(figsize=(20,10))
- ax = fig.add_subplot(111)
- print('difference at :',)
- ax.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]], 'o-',ms=0.5,lw=0.3)
- ax.plot(illumLEDcontrolrescaled[idxIllum][:shortest[0]], 'o-',ms=0.5,lw=0.3)
- plt.show()
- #plt.clf()
- fig = plt.figure(figsize=(20, 10))
- ax = fig.add_subplot(111)
- ax.axhline(y=LEDVideoThreshold,c='C0',ls='--')
- ax.axhline(y=illumLEDcontrolThreshold,c='C1',ls='--')
- ledVideo = ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]]
- ledIllumDAQ = illumLEDcontrolrescaled[idxIllum][:shortest[0]]
- ax.plot(np.arange(len(ledVideo))[ledVideo > LEDVideoThreshold],ledVideo[ledVideo > LEDVideoThreshold], 'v', c='C0', ms=2)
- ax.plot(np.arange(len(ledVideo))[ledVideo < LEDVideoThreshold],ledVideo[ledVideo < LEDVideoThreshold], 'o', c='C0', ms=2)
- ax.plot(np.arange(len(ledIllumDAQ))[ledIllumDAQ > illumLEDcontrolThreshold],ledIllumDAQ[ledIllumDAQ>illumLEDcontrolThreshold], 'v',c='C1', ms=2)
- ax.plot(np.arange(len(ledIllumDAQ))[ledIllumDAQ < illumLEDcontrolThreshold],ledIllumDAQ[ledIllumDAQ<illumLEDcontrolThreshold], 'o', c='C1', ms=2)
- plt.show()
- pdb.set_trace()
- elif shiftToZero[1] == (shiftToZero[0]+5):
- print('Multiple shifts to zero, so multiple perfect overlays exist. First overlay with shift %s will be used.' % shiftToZero[0])
- pass
- else:
- print('Problem! More than one shift led to perfect overlay!')
- #np.arange(np.diff(idxRecordedFrames)>1)
- #
- idxTemp = idxFramesDuringRecording + 0
- idx = idxTemp[idxTemp>=0]
- idxIllum = idx[idx<len(illumLEDcontrolBin)]
- #pdb.set_trace()
- shortest = [len(ledVideoRoiRescaled[3][mask][missedFramesBegin:]) if len(ledVideoRoiRescaled[3][mask][missedFramesBegin:])<len(illumLEDcontrolrescaled[idxIllum]) else len(illumLEDcontrolrescaled[idxIllum])]
- plt.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]],illumLEDcontrolrescaled[idxIllum][:shortest[0]],'o',ms=1)
- plt.show()
- pdb.set_trace()
- plt.plot(ledVideoRoiRescaled[3][mask][missedFramesBegin:][:shortest[0]],'o-')
- plt.plot(illumLEDcontrolrescaled[idxIllum][:shortest[0]],'o-')
- plt.show()
- pdb.set_trace()
- finalShiftToZero = shiftToZero[0]
- print('final shift to zero and final length : ', finalShiftToZero,finalLength[0])
- idxTemp = idxFramesDuringRecording + finalShiftToZero
- idx = idxTemp[idxTemp>=0]
- idxIllumFinal = idx[idx<len(illumLEDcontrolBin)][:finalLength[0]]
- compareIllumination = False
- if compareIllumination:
- illum = illumLEDcontrolrescaled[idxIllumFinal]
- videoROI = ledVideoRoiRescaled[3][mask][missedFramesBegin:]
- shortest = [len(videoROI) if len(videoROI) < len(illum) else len(illum)]
- bothCombined = np.column_stack((videoROI[:shortest[0]],illum[:shortest[0]]))
- plt.plot(videoROI[:shortest[0]], illum[:shortest[0]], 'o')
- plt.show()
- pdb.set_trace()
- try :
- illumValues = pickle.load( open('illuminatoinValues.p', 'rb' ) )
- except :
- illumValues = bothCombined
- else:
- illumValues = np.row_stack((illumValues,bothCombined))
- pickle.dump(illumValues, open('illuminatoinValues.p', 'wb'))
- frameTimes = startEndExposureTime[idxIllumFinal]
- frameStartStopIdx = startEndExposurepIdx[idxIllumFinal]
- videoIdx = np.arange(len(ledVideoRoiBins[3]))[mask][missedFramesBegin:][:finalLength[0]]
- #recFrames = ledVideoRoi[2][videoIdx]
- ddd = np.diff(idxIllumFinal)
- print('Total number of dropped and excluded frames : ', np.sum(ddd-1), 'out of',len(ledVideoRoi[2]),'frame in total.')
- print('Excluded frames :', len(idxToExclude))
- print('Dropped frames :', np.sum(ddd-1)-len(idxToExclude))
- frameSummary = np.array([len(ledVideoRoi[2]),np.sum(ddd-1),len(idxToExclude), np.sum(ddd-1)-len(idxToExclude)])
- #pdb.set_trace()
- return (idxIllumFinal,frameTimes,frameStartStopIdx,videoIdx,frameSummary)
- ##############################################################################################################################
- #pdb.set_trace()
- for i in range(10):
- #print(i)
- #idxTest = idxRecordedFramesCleaned[1:-3] - 1
- #compare = ledVideoRoiBins[3][2:-3] == illumLEDcontrolBin[idxTest]
- #shortestLength = [len(illum) if (0<(len(ledVideoRoiBins)-(len(illum)+i))) else ]
- videoRoi = ledVideoRoiBins[3][i]
- if len(videoRoi) > len(illum):
- #shortestLength = len(illum)
- #else:
- #shortestLength = len(videoRoi)
- print('problem in length relations')
- pdb.set_trace()
- compare = illum[:len(videoRoi)] == videoRoi
- differences = np.sum(np.invert(compare))
- print('number of differences :', i, differences,i,len(illum)+i,len(videoRoi))
- shifting.append([i,differences,i,len(illum)+i,len(videoRoi)])
- shifting = np.asarray(shifting)
- correctShift = np.argwhere(shifting[:,1]==0)
- if len(correctShift) == 0:
- print('No perfect overlay has been found')
- print(shifting)
- pdb.set_trace()
- elif len(correctShift)>1:
- print('Multiple corret overlays have been found. Suspicious!')
- print(shifting)
- pdb.set_trace()
- elif len(correctShift) == 1:
- rightShift = shifting[correctShift[0][0]]
- print('The correct shift is ', rightShift)
- print('Number of recorded videos :', len(ledVideoRoi[2][rightShift[2]:rightShift[3]]))
- print('Number of associated time points :', len(startEndExposurepIdx[idxRecordedFramesCleaned[idxTestMask]][:rightShift[4]]))
- ddd = np.diff(idxRecordedFramesCleaned[idxTestMask][:rightShift[4]])
- 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])
- idxVideo = np.arange(rightShift[2],rightShift[3])
- idxTimePoints = idxRecordedFramesCleaned[idxTestMask][:rightShift[4]]
- #pdb.set_trace()
- return (idxVideo,idxTimePoints,startEndExposureTime,startEndExposurepIdx,rightShift)
- if compare == 0:
- #plt.plot(ledVideoRoi[0][3][2:-3], 'o-', label='ledVideoRoi')
- ii = 3
- plt.plot(ledVideoRoiRescaled[3][ii:len(illum)+ii], 'o-', label='ledVideoRoiRescaled')
- plt.plot(ledVideoRoiBins[3][ii:len(illum)+ii],'o-',label='ledVideoRoiBins')
- plt.plot(illumLEDcontrol[idxRecordedFramesCleaned[idxTestMask]],'o-',label='illumLEDcontrol')
- plt.plot(illumLEDcontrolBin[idxRecordedFramesCleaned[idxTestMask]], 'o-', label='illumLEDcontrolBin')
- plt.legend()
- plt.show()
- pdb.set_trace()
- #compare = illumLEDcontrol[idxTest] ==
- #idxTest = idxRecordedFramesCleaned[1:-1]-1
- totLength = len(illumLEDcontrol[idxTest])
- ret = np.array_equal(illumLEDcontrol[idxRecordedFramesCleaned],ledVideoRoiBins[3][i:(totLength+i)])
- print(i,ret)
- pdb.set_trace()
- pdb.set_trace()
- if len(illumLEDcontrol) <= len(ledVIDEOroi):
- illuminationLonger = True
- ledVIDEOroiMask = np.arange(len(ledVIDEOroi)) < len(illumination)
- illuminationMask = np.arange(len(illumination)) < len(illumination)
- cc = crosscorr(1,illumination,ledVIDEOroi[ledVIDEOroiMask],20) # calculate cross-correlation between LED in video and LED from DAQ array
- else:
- illumniationLonger = False
- ledVIDEOroiMask = np.arange(len(ledVIDEOroi)) < len(ledVIDEOroi)
- illuminationMask = np.arange(len(illumination)) < len(ledVIDEOroi)
- cc = crosscorr(1, illumination[illuminationMask], ledVIDEOroi, 3) # calculate cross-correlation between LED in video and LED from DAQ array
- peaks = find_peaks(cc[:,1],height=0)
- if len(peaks[0]) > 1:
- print('MULTIPLE peaks found in cross-correlogram between LED brigthness and DAQ array')
- pdb.set_trace()
- elif len(peaks[0]) == 0:
- print('NO peaks were found in cross-correlogram between LED brigthness and DAQ array')
- pdb.set_trace()
- else:
- pdb.set_trace()
- shift = cc[:,0][peaks[0][0]]
- shiftInt = int(shift)
- print('video trace has to be shifted by (float and int number) ', shift, shiftInt)
- #print(len(ledVIDEOroi),len(illumination))
- #pdb.set_trace()
- if verbose:
- if shiftInt >= 0:
- plt.plot(ledVIDEOroi[ledVIDEOroiMask][shiftInt:],'o-',ms=1,label='Video roi (shifted)')
- else:
- plt.plot(ledVIDEOroi[ledVIDEOroiMask][:shiftInt], 'o-', ms=1, label='Video roi (shifted)')
- plt.plot(illumination[illuminationMask],'o-',ms=1,label='from LED daq control')
- plt.legend()
- plt.show()
- frameIdx = np.arange(len(ledVIDEOroi))
- recordedFramesIdx = frameIdx[shiftInt:(len(illumination)+shiftInt)]
- #pdb.set_trace()
- return (startEndExpTime,startEndExpIdx,recordedFramesIdx)
- ##### end of current implementation ##############################################################################################
- #framesIdxDuringRec = np.array(len(softFrameTimes))[(arrayTimes[expEnd[0]]+0.002) < softFrameTimes]
- #framesIdxDuringRec = framesIdxDuringRec[:len(expStart)]
- if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
- recordedFrames = frames[3:(len(midExposure) + 3)]
- elif arrayTimes[int(midExposure[0])]<0.003:
- recordedFrames = frames[2:(len(midExposure) + 2)]
- else:
- recordedFrames = frames[:len(midExposure)]
- print('number of tot. frames, recorded frames, exposures start, end :',numberOfFrames,len(recordedFrames), len(expStart), len(expEnd))
- if display:
- ledON = np.zeros(len(exposureArray))
- for i in range(11):
- ledON[((i*1.)<=arrayTimes) & ((i*1.+0.2)>arrayTimes)] = 1.
- ledON[29.<=arrayTimes] = 1.
- data = np.loadtxt('/home/mgraupe/2019.04.01_000-%s.csv' % (rec[-3:]),delimiter=',',skiprows=1,usecols=(0,1))
- print(len(data))
- plt.plot(arrayTimes,exposureArray/32.)
- plt.plot(arrayTimes,ledON)
- print('fist frame at %s sec' % arrayTimes[int(midExposure[0])],end='')
- if arrayTimes[int(midExposure[0])]<0.015 and arrayTimes[int(midExposure[0])]>=0.003:
- plt.plot(arrayTimes[midExposure.astype(int)], (data[3:(len(midExposure) + 3), 1] - 148.6) / 105.4, 'o-')
- print(3)
- elif arrayTimes[int(midExposure[0])]<0.003:
- plt.plot(arrayTimes[midExposure.astype(int)], (data[2:(len(midExposure) + 2), 1] - 148.6) / 105.4, 'o-')
- print(2)
- else:
- plt.plot(arrayTimes[midExposure.astype(int)], (data[:len(midExposure), 1] - 148.6) / 105.4, 'o-')
- print(0)
- #plt.plot(softFrameTimes[data[:-6,0].astype(np.int)+6],(data[:-6,1]-148.6)/105.4)
- #plt.plot(softFrameTimes,np.ones(len(softFrameTimes)),'|')
- plt.show()
- pdb.set_trace()
- return (expStartTime,expEndTime,recordedFrames)
- #################################################################################
- # detect spikes in ephys trace
- #################################################################################
- def applyImageNormalizationMask(frames,imageMetaInfo,normFrame,normImageMetaInfo,mouse, date, rec):
- print(imageMetaInfo, normImageMetaInfo)
- pixelRange = 10
- print('small, large frame : ', np.shape(frames), np.shape(normFrame))
- print('pixel-ratio, x-ratio, y-ratio', imageMetaInfo[4]/normImageMetaInfo[4],end='')
- fig = plt.figure()
- rect1 = patches.Rectangle(normImageMetaInfo[:2], normImageMetaInfo[2], normImageMetaInfo[3],linewidth=1,edgecolor='C0',facecolor='none')
- rect2 = patches.Rectangle(imageMetaInfo[:2],imageMetaInfo[2],imageMetaInfo[3],linewidth=1,edgecolor='C1',facecolor='none')
- framesF = np.array(frames,dtype=float)
- avgFrame = np.average(frames[:,:,:,0],axis=0)
- # rescale image stack to the resolution of the normalization image
- framesRescaled = scipy.ndimage.zoom(framesF, [1,imageMetaInfo[4]/normImageMetaInfo[4],imageMetaInfo[4]/normImageMetaInfo[4],1], order=3)
- # average across all time points of image stack
- #avgFrameZ = np.average(framesRescaled[:,:,:,0],axis=0)
- # rescale the average image to match pixel-size of normalization image, the re-scaling factor of the ratio of the pixel-sizes : stack/norm
- avgFrameZ = scipy.ndimage.zoom(avgFrame, imageMetaInfo[4]/normImageMetaInfo[4], order=3)
- # x,y location in pixel indices of the stack in the normalization image
- xLoc = int(np.round((imageMetaInfo[0] - normImageMetaInfo[0]) / normImageMetaInfo[4]))
- yLoc = int(np.round((imageMetaInfo[1] - normImageMetaInfo[1]) / normImageMetaInfo[4]))
- # dimensions of the rescaled image
- xDim = np.shape(avgFrameZ)[0]
- yDim = np.shape(avgFrameZ)[1]
- #####################
- ax0 = fig.add_subplot(2,3,4)
- ax0.set_title('avg. of image stack',size=7)
- ax0.imshow(np.transpose(avgFrame))
- ax1 = fig.add_subplot(2,3,5)
- ax1.set_title('avg. of image stack : rescaled to norm. image pixel size',size=7)
- ax1.imshow(np.transpose(avgFrameZ))
- print('image stack : ', np.shape(framesRescaled))
- scipy.io.savemat('%s_%s_%s_imageStackBeforeRescaling.mat' % (mouse, date, rec), mdict={'dataArray': framesF})
- scipy.io.savemat('%s_%s_%s_imageStack.mat' % (mouse, date, rec), mdict={'dataArray': framesRescaled})
- #img_stack_uint8 = mapToXbit(avgFrameZ,8)
- #pdb.set_trace()
- #tiff.imsave('avg_imageStack_scaled.tif', np.array(img_stack_uint8, dtype=np.uint8))
- ax2 = fig.add_subplot(2,3,1)
- ax2.set_title('normalization image with image stack rectangle',size=7)
- 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')
- ax2.imshow(np.transpose(normFrame[0,:,:,0]))
- ax2.add_patch(ret)
- scipy.io.savemat('%s_%s_%s_registrationImage.mat' % (mouse, date, rec), mdict={'dataArray': normFrame[0,:,:,0]})
- ax2 = fig.add_subplot(2,3,6)
- ax2.set_title('area of image stack from normalization image',size=7)
- #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')
- ax2.imshow(np.transpose(normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0]))
- #ax2.add_patch(ret)
- print('norm. image : ', np.shape(normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0]))
- scipy.io.savemat('%s_%s_%s_normalizationImage.mat' % (mouse, date, rec), mdict={'dataArray': normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0]})
- normFrameF = np.array(normFrame, dtype=float)
- test1 = scipy.ndimage.gaussian_filter1d(normFrameF[:,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),:], 2, axis=1)
- test2 = scipy.ndimage.gaussian_filter1d(test1, 2, axis=2)
- #test1 = scipy.ndimage.gaussian_filter1d(framesF, 2, axis=1)
- #test2 = scipy.ndimage.gaussian_filter1d(test1, 2, axis=2)
- #norm8bit = mapToXbit(test2,8)
- filter1D = scipy.ndimage.gaussian_filter1d(framesRescaled, 2, axis=1)
- filter2D = scipy.ndimage.gaussian_filter1d(filter1D, 2, axis=2)
- norm = filter2D / test2
- norm8bit = mapToXbit(norm,8)
- ax2 = fig.add_subplot(2,3,2)
- ax2.set_title('normalized average image',size=7)
- #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')
- ax2.imshow(np.transpose(np.average(norm[:,:,:,0],axis=0)))
- plt.show()
- #pdb.set_trace()
- errMatrix = np.zeros((pixelRange*2+1,pixelRange*2+1))
- # #row, col = np.indices(err)
- xRange = np.arange(pixelRange*2+1)
- yRange = np.copy(xRange)
- # #avgFrameZ = np.copy(normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0])
- for xy in itertools.product(xRange, yRange):
- #err = np.abs((normFrame[0,xLoc:(xLoc+xDim),yLoc:(yLoc+yDim),0] - avgFrameZ) ** 2).sum() / (xDim*yDim)
- xStart = xLoc + xy[0] - pixelRange
- yStart = yLoc + xy[1] - pixelRange
- normImg = normFrameF[0,xStart:(xStart+xDim),yStart:(yStart+yDim),0]
- normImgNorm = mapToXbit(normImg,8) #normImg - np.average(normImg)
- avgFrameZNorm = mapToXbit(np.average(framesRescaled[:,:,:,0],axis=0),8) #avgFrameZ - np.average(avgFrameZ)
- errMatrix[xy[0],xy[1]] = ((normImgNorm - avgFrameZNorm) ** 2).sum() / (xDim*yDim)
- minimumIndices = np.argwhere(errMatrix == np.min(errMatrix))
- print('MI :', minimumIndices)
- #pdb.set_trace()
- #xNorm = np.linspace(normImageMetaInfo[0],normImageMetaInfo[0]+
- #ax.add_patch(rect1)
- #ax.add_patch(rect2)
- #ax.set_ylim(normImageMetaInfo[1]-10,normImageMetaInfo[1]+normImageMetaInfo[3]+10)
- #ax.set_xlim(normImageMetaInfo[0]-10,normImageMetaInfo[0]+normImageMetaInfo[2]+10)
- #plt.patches.Rectangle(normImageMetaInfo[:2],normImageMetaInfo[2],normImageMetaInfo[3])
- #plt.patches.Rectangle(imageMetaInfo[:2],imageMetaInfo[2],imageMetaInfo[3])
- #plt.show()
- #pdb.set_trace()
- return norm8bit
- #################################################################################
- # detect spikes in ephys trace
- #################################################################################
- def detectPawTrackingOutlies(pawTraces,pawMetaData):
- jointNames = pawMetaData['data']['DLC-model-config file']['all_joints_names']
- threshold = 60
- def findOutliersBasedOnMaxSpeed(onePawData,jointName,i): # should be an 3 column array frame#, x, y
- frDisplOrig = np.sqrt((np.diff(onePawData[:, 1])) ** 2 + (np.diff(onePawData[:, 2])) ** 2) / np.diff(onePawData[:, 0])
- onePawDataTmp = np.copy(onePawData)
- onePawIndicies = np.arange(len(onePawData))
- # excursionsBoolOld = np.zeros(len(pawDataTmp)-1,dtype=bool)
- nIt = 0
- while True: # cycle as long as there are large displacements
- frDispl = np.sqrt((np.diff(onePawDataTmp[:,1])) ** 2 + (np.diff(onePawDataTmp[:,2])) ** 2) / np.diff(onePawDataTmp[:, 0]) # calculate displacement
- excursionsBoolTmp = frDispl > threshold # threshold displacement
- print(nIt, sum(excursionsBoolTmp))
- nIt += 1
- if sum(excursionsBoolTmp) == 0: # no excursions above threshold are found anymore -> exit loop
- break
- else:
- onePawDataTmp = onePawDataTmp[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
- onePawIndicies = onePawIndicies[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
- print('%s # of positions, # of detected mis-trackings, fraction : ' % (jointName), len(onePawData), len(onePawData) - len(onePawDataTmp), (len(onePawData) - len(onePawDataTmp)) / len(onePawData))
- if jointName=='tail_base_bottom':
- pdb.set_trace()
- return (len(onePawData),len(onePawDataTmp),onePawIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig)
- pawTrackingOutliers = []
- for i in range(len(jointNames)):
- (tot,correct,correctIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig) = findOutliersBasedOnMaxSpeed(np.column_stack((pawTraces[:,0],pawTraces[:,(i*3+1)],pawTraces[:,(i*3+2)])),jointNames[i],i)
- pawTrackingOutliers.append([i,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig])
- return pawTrackingOutliers
- #pdb.set_trace()
- #################################################################################
- def detectPawTrackingOutliersObstacle(pawTraces,pawMetaData):
- jointNames = pawMetaData['data']['DLC-model-config file']['all_joints_names']
- threshold = 70
- print(jointNames)
- def findOutliersBasedOnMaxSpeedObstacle(onePawData,jointName,i): # should be an 3 column array frame#, x, y
- # if jointName=='obstacle':
- # threshold=100
- # else:
- # threshold = 80
- frDisplOrig = np.sqrt((np.diff(onePawData[:, 1])) ** 2 + (np.diff(onePawData[:, 2])) ** 2) / np.diff(onePawData[:, 0])
- onePawDataTmp = np.copy(onePawData)
- onePawIndicies = np.arange(len(onePawData))
- # excursionsBoolOld = np.zeros(len(pawDataTmp)-1,dtype=bool)
- nIt = 0
- while True: # cycle as long as there are large displacements
- frDispl = np.sqrt((np.diff(onePawDataTmp[:,1])) ** 2 + (np.diff(onePawDataTmp[:,2])) ** 2) / np.diff(onePawDataTmp[:, 0]) # calculate displacement
- excursionsBoolTmp = frDispl > threshold # threshold displacement
- print(nIt, sum(excursionsBoolTmp))
- nIt += 1
- if sum(excursionsBoolTmp) == 0: # no excursions above threshold are found anymore -> exit loop
- break
- else:
- onePawDataTmp = onePawDataTmp[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
- onePawIndicies = onePawIndicies[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
- print('%s # of positions, # of detected mis-trackings, fraction : ' % (jointName), len(onePawData), len(onePawData) - len(onePawDataTmp), (len(onePawData) - len(onePawDataTmp)) / len(onePawData))
- # if jointName=='tail_base_bottom':
- # # pdb.set_trace()
- return (len(onePawData),len(onePawDataTmp),onePawIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig)
- pawTrackingOutliersDic = {}
- pawTrackingOutliersList=[]
- pawTrackingOutliersBot_paw=[]
- b=0
- bot_paw = ['front_left_bottom', 'front_right_bottom', 'hind_left_bottom', 'hind_right_bottom']
- for i in range(len(jointNames)):
- pawTrackingOutliersDic[jointNames[i]] = {}
- (tot,correct,correctIndicies,onePawData,onePawDataTmp,frDispl,frDisplOrig) = findOutliersBasedOnMaxSpeedObstacle(np.column_stack((pawTraces[:,0],pawTraces[:,(i*3+1)],pawTraces[:,(i*3+2)])),jointNames[i],i)
- pawTrackingOutliersList.append([i,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig]) #all pf these are parameters that we stock in each jointName
- parameters = [i,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig]
- parameters_string = ['i','tot','correct','correctIndicies','jointName','onePawData','onePawDataTmp','frDispl','frDisplOrig']
- if any([x in jointNames[i] for x in bot_paw]):
- pawTrackingOutliersBot_paw.append([b,tot,correct,correctIndicies,jointNames[i],onePawData,onePawDataTmp,frDispl,frDisplOrig])
- b+=1
- for j in range(len(parameters)):
- pawTrackingOutliersDic[jointNames[i]][parameters_string[j]] = parameters[j]
- return pawTrackingOutliersDic,pawTrackingOutliersList,pawTrackingOutliersBot_paw
- #################################################################################
- def detectPawTrackingOutliersObstacleVids(pawTraces, pawMetaData):
- jointNames = pawMetaData['data']['DLC-model-config file']['all_joints_names']
- threshold = 70
- print(jointNames)
- def findOutliersBasedOnMaxSpeedObstacle(onePawData, jointName, i): # should be an 3 column array frame#, x, y
- # if jointName=='obstacle':
- # threshold=100
- # else:
- # threshold = 80
- frDisplOrig = np.sqrt((np.diff(onePawData[:, 1])) ** 2 + (np.diff(onePawData[:, 2])) ** 2) / np.diff(
- onePawData[:, 0])
- onePawDataTmp = np.copy(onePawData)
- onePawIndicies = np.arange(len(onePawData))
- # excursionsBoolOld = np.zeros(len(pawDataTmp)-1,dtype=bool)
- nIt = 0
- while True: # cycle as long as there are large displacements
- frDispl = np.sqrt((np.diff(onePawDataTmp[:, 1])) ** 2 + (np.diff(onePawDataTmp[:, 2])) ** 2) / np.diff(
- onePawDataTmp[:, 0]) # calculate displacement
- excursionsBoolTmp = frDispl > threshold # threshold displacement
- print(nIt, sum(excursionsBoolTmp))
- nIt += 1
- if sum(excursionsBoolTmp) == 0: # no excursions above threshold are found anymore -> exit loop
- break
- else:
- onePawDataTmp = onePawDataTmp[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
- onePawIndicies = onePawIndicies[np.concatenate((np.array([True]), np.invert(excursionsBoolTmp)))]
- print('%s # of positions, # of detected mis-trackings, fraction : ' % (jointName), len(onePawData),
- len(onePawData) - len(onePawDataTmp), (len(onePawData) - len(onePawDataTmp)) / len(onePawData))
- # if jointName=='tail_base_bottom':
- # # pdb.set_trace()
- return (len(onePawData), len(onePawDataTmp), onePawIndicies, onePawData, onePawDataTmp, frDispl, frDisplOrig)
- pawTrackingOutliersDic = {}
- pawTrackingOutliersList = []
- pawTrackingOutliersBot_paw = []
- b = 0
- bot_paw = ['front_left_bottom', 'front_right_bottom', 'hind_left_bottom', 'hind_right_bottom']
- for i in range(len(jointNames)):
- pawTrackingOutliersDic[jointNames[i]] = {}
- tot=[]
- correct=[]
- correctIndicies=np.array([])
- onePawData=np.empty((3))
- onePawDataTmp=np.empty((3))
- frDispl=np.array([])
- frDisplOrig=np.array([])
- for v in np.unique(pawMetaData['obs_number']):
- # print('obstacle ids', np.unique(pawMetaData['obs_number']))
- vmask=pawMetaData['obs_number']==v
- try:
- (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)
- except:
- 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')
- pdb.set_trace()
- # print(len(onePawDataTmp_v), len(frDispl_v), len(onePawData_v), len(frDisplOrig_v))
- correctIndicies=np.concatenate((correctIndicies,correctIndicies_v))
- onePawData=np.vstack((onePawData,onePawData_v))
- onePawDataTmp=np.vstack((onePawDataTmp,onePawDataTmp_v))
- frDispl = np.concatenate((frDispl, frDispl_v))
- frDisplOrig=np.concatenate((frDisplOrig, frDisplOrig_v))
- onePawData=onePawData[1:]
- onePawDataTmp=onePawDataTmp[1:]
- # pdb.set_trace()
- pawTrackingOutliersList.append([i, tot, correct, correctIndicies, jointNames[i], onePawData, onePawDataTmp, frDispl,frDisplOrig]) # all pf these are parameters that we stock in each jointName
- parameters = [i, tot, correct, correctIndicies, jointNames[i], onePawData, onePawDataTmp, frDispl,frDisplOrig]
- parameters_string = ['i', 'tot', 'correct', 'correctIndicies', 'jointName', 'onePawData', 'onePawDataTmp','frDispl', 'frDisplOrig']
- if any([x in jointNames[i] for x in bot_paw]):
- pawTrackingOutliersBot_paw.append(
- [b, tot, correct, correctIndicies, jointNames[i], onePawData, onePawDataTmp, frDispl, frDisplOrig])
- b += 1
- for j in range(len(parameters)):
- pawTrackingOutliersDic[jointNames[i]][parameters_string[j]] = parameters[j]
- return pawTrackingOutliersDic, pawTrackingOutliersList, pawTrackingOutliersBot_paw
- #################################################################################
- #################################################################################
- # convert ca traces in easily usable numpy array
- #################################################################################
- def getCaWheelPawInterpolatedDictsPerDay(nSess,allCorrDataPerSession,allStepData,showFig = False):
- baselineTime = 5.
- # calcium traces ##############################################################
- trialStartUnixTimes = []
- fTraces = allCorrDataPerSession[nSess]['caImg']['Fluo'] #[3][0][0]
- timeStamps = allCorrDataPerSession[nSess]['caImg']['timeStamps'] # [3][0][3] # the array containing the time-stamp array
- recordings = np.unique(timeStamps[:, 1]) # determine how many recordings where performed
- caTracesDict = {}
- for n in range(len(recordings)):
- mask = (timeStamps[:, 1] == recordings[n])
- triggerStart = timeStamps[:, 5][mask]
- trialStartUnixTimes.append(timeStamps[:, 3][mask][0])
- if n > 0:
- if oldTriggerStart > triggerStart[0]:
- print('problem in trial order')
- sys.exit(1)
- # for i in range(len(fTraces)):
- # triggerstart - time of the acq start trigger for the current acquisition
- # timeStamps[:, 4][mask] - time of the first pixel in the frame passed since acqModeEpoch
- caTracesTime = (timeStamps[:, 4][mask] - triggerStart) # triggerStart is negative
- #pdb.set_trace()
- caTracesFluo = fTraces[:, mask]
- # pdb.set_trace()
- # caTraces.append(np.column_stack((caTracesTime,caTracesFluo)))
- caTracesDict[n] = np.row_stack((caTracesTime, caTracesFluo))
- #print(np.shape(np.row_stack((caTracesTime, caTracesFluo))))
- oldTriggerStart = triggerStart[0]
- # wheel speed ######################################################
- # also find calmest pre-motorization period
- minPreMotorMeanV = 1000.
- wheelTracks = allCorrDataPerSession[nSess]['wheel'] #[1]
- nRec = 0
- # print(len(wheelTracks))
- wheelSpeedDict = {}
- for n in range(len(wheelTracks)):
- wheelRecStartTime = wheelTracks[n]['timeStamp']#[3]
- if (trialStartUnixTimes[nRec] - wheelRecStartTime) < 1.:
- # if not wheelTracks[n][4]:
- # recStartTime = wheelTracks[0][3]
- if nRec > 0:
- if oldRecStartTime > wheelRecStartTime:
- print('problem in trial order')
- sys.exit(1)
- wheelTime = wheelTracks[n]['sTimes']#[2]
- wheelSpeed = wheelTracks[n]['linearSpeed']#[1] # linear wheel speed in cm/s
- angleSpeed = wheelTracks[n]['angluarSpeed']#[0]
- angleTime = wheelTracks[n]['angleTimes']#[5]
- wheelSpeedDict[nRec] = np.row_stack((wheelTime, wheelSpeed))
- #pdb.set_trace()
- preMMask = (wheelTime < baselineTime)
- preMmeanV = np.mean(np.abs(wheelSpeed[preMMask]))
- #print(nSess, nRec, preMmeanV)
- if preMmeanV < minPreMotorMeanV:
- slowestRec = nRec
- minPreMotorMeanV = np.copy(preMmeanV)
- nRec += 1
- oldRecStartTime = wheelRecStartTime
- print('trials with slowest baseline period :', slowestRec)
- # normalize ca-traces by baseline fluorescence : fluorescence during the first baselineTime seconds in the least active recording #############################################
- mask = (caTracesDict[slowestRec][0] < baselineTime)
- F0 = np.mean(caTracesDict[slowestRec][1:][:, mask], axis=1)
- #pdb.set_trace()
- for n in range(len(recordings)):
- normalizedCaTraces = (caTracesDict[n][1:] - F0[:, np.newaxis]) / F0[:, np.newaxis]
- caTracesDict[n][1:] = np.copy(normalizedCaTraces)
- # pdb.set_trace()
- # paw speed ######################################################
- pawTracks = allCorrDataPerSession[nSess]['paws']#[2]
- nRec = 0
- pawTracksDict = {}
- pawID = []
- for n in range(len(pawTracks)):
- # if not wheelTracks[n][4]:
- pawRecStartTime = pawTracks[n]['recStartTime']#[4]
- if (trialStartUnixTimes[nRec] - pawRecStartTime) < 1.:
- if nRec > 0:
- if oldRecStartTime > pawRecStartTime:
- print('problem in trial order')
- sys.exit(1)
- pawTracksDict[nRec] = {}
- for i in range(4):
- # pdb.set_trace()
- if nRec == 0:
- pawID.append(pawTracks[n]['jointNamesFramesInfo'][i][0])
- # pawTracksDict[nFig][i]['pawID'] = pawTracks[n][2][i][0]
- pawSpeedTime = pawTracks[n]['pawSpeed'][i][:,0] # times of cleared paw speed
- pawSpeed = pawTracks[n]['pawSpeed'][i][:,1] # 1 is combined speed in 2-d plane of the camera view from below,
- pawTracksDict[nRec][i] = np.row_stack((pawSpeedTime,pawSpeed)) # interp = interp1d(pawSpeedTime, pawSpeed) # newPawSpeedAtCaTimes = interp(caTracesTime[nFig]) # pawTracksDict[i]['pawSpeed'].extend(newPawSpeedAtCaTimes)
- oldRecStartTime = pawRecStartTime
- nRec += 1
- # interpolation #############################################################################
- # interp = interp1d(wheelTime, wheelSpeed)
- # interpMask = (caTracesTime[nFig] >= wheelTime[0]) & (caTracesTime[nFig] <= wheelTime[-1])
- # newWheelSpeedAtCaTimes = interp(caTracesTime[nFig][interpMask])
- # wheelSpeedAll.extend(newWheelSpeedAtCaTimes)
- wheelSpeedDictInterp = wheelSpeedDict.copy()
- pawTracksDictInterp = pawTracksDict.copy()
- caTracesDictInterp = caTracesDict.copy()
- for nrec in range(len(caTracesDict)):
- # determine interpolation range
- 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]))
- 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]))
- interpMask = (caTracesDict[nrec][0] >= startInterpTime) & (caTracesDict[nrec][0] <= endInterpTime)
- # restrict ca-traces to interpolation range
- #pdb.set_trace()
- #matrix = np.copy(caTracesDict[nrec])
- caTracesDictInterp[nrec] = caTracesDict[nrec][:,interpMask]
- # interpolate wheel speed
- interpWheel = interp1d(wheelSpeedDict[nrec][0], wheelSpeedDict[nrec][1])#,kind='cubic')
- newWheelSpeedAtCaTimes = interpWheel(caTracesDict[nrec][0][interpMask])
- wheelSpeedDictInterp[nrec] = np.row_stack((caTracesDict[nrec][0][interpMask], newWheelSpeedAtCaTimes))
- # interpolate paw speed
- for i in range(4):
- interpPaw = interp1d(pawTracksDict[nrec][i][0], pawTracksDict[nrec][i][1])#,kind='cubic')
- newPawSpeedAtCaTimes = interpPaw(caTracesDict[nrec][0][interpMask])
- pawTracksDictInterp[nrec][i] = np.row_stack((caTracesDict[nrec][0][interpMask], newPawSpeedAtCaTimes))
- #pawAll[i].extend(newPawSpeedAtCaTimes)
- if showFig:
- cc = ['C0','C1']
- fig, ax = plt.subplots(figsize=(10, 6))
- for i in range(2):
- ax.plot(pawTracksDictInterp[nrec][i][0],pawTracksDictInterp[nrec][i][1]/np.max(pawTracksDictInterp[nrec][i][1]),c=cc[i])
- idxSwings = allStepData[nSess][4][nrec][3][i][1]
- # print('Wow nSess',nDay,allStepData[nDay-1][0])
- recTimes = allStepData[nSess][4][nrec][4][i][2]
- # pdb.set_trace()
- idxSwings = np.asarray(idxSwings)
- for k in range(len(idxSwings)): # loop over all swings
- startSwingTime = recTimes[idxSwings[k, 0]]
- endSwingTime = recTimes[idxSwings[k, 1]]
- ax.fill_between((startSwingTime,endSwingTime), 0, 1, color='0.6', alpha=0.5, transform=ax.get_xaxis_transform())
- ax.fill_between((startSwingTime,endSwingTime), 0, 1, color='0.4', alpha=0.5, transform=ax.get_xaxis_transform())
- #ax.plot(caTracesDictInterp[nrec][0],caTracesDictInterp[nrec][1]/np.max(caTracesDictInterp[nrec][1]),c='black')
- plt.show()
- pdb.set_trace()
- return (wheelSpeedDictInterp,pawTracksDictInterp,caTracesDictInterp,wheelSpeedDict,pawTracksDict,caTracesDict,slowestRec)
- #################################################################################
- # calculate correlations between ca-imaging, wheel speed and paw speed
- #################################################################################
- def doRegressionAnalysis(mouse,allCorrDataPerSession,allStepData,borders=None,figShow=False):
- matplotlib.use('TkAgg') # WxAgg
- from sklearn.linear_model import LinearRegression
- from sklearn.svm import SVR
- #SVR(kernel='rbf', C=1e3, gamma=0.1)
- #from sklearn.ensemble import RandomForestRegressor
- regressionN = 6
- Rvalues = []
- #for nSess in range(len(allCorrDataPerSession)):
- for nDay in range(len(allCorrDataPerSession)):
- print(nDay,allCorrDataPerSession[nDay]['folder'])
- (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP, wheelSpeedDict, pawTracksDict, caTracesDict, slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay,allCorrDataPerSession,allStepData)
- # ATTENTION : all of the arrays also contain a time array
- # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTimeSame]
- # dims of pawTracksDictInterP : [nSessions][nPaw][2][valuesOverTimeSame]
- # dims of caTracesDictInterP : [nSessions][nRois+1][valuesOverTimeSame]
- # dims of wheelSpeedDict : [nSessions][2][valuesOverTime]
- # dims of pawTracksDict : [nSessions][nPaw][2][valuesOverTime]
- # dims of caTracesDict : [nSessions][nRois+1][valuesOverTime]
- #(wheelSpeedDict, pawTracksDict, caTracesDict,aa,bb,cc) = getCaWheelPawInterpolatedDictsPerDay(nSess, allCorrDataPerSession)
- nRecWheel = len(wheelSpeedDictInterP)
- nRecPaw = len(pawTracksDictInterP)
- nRecCa = len(caTracesDictInterP)
- print('Recording length :', nRecWheel, nRecPaw, nRecCa)
- if (nRecWheel != nRecPaw ) or (nRecWheel != nRecCa):
- print('problem in number of recordings listed in dictionaries')
- # loop over 5 different regressions, each using a different combination of test and train samples
- recs = range(nRecCa)
- RTempValues = []
- for reg in range(nRecCa): # loop over all recordings
- recsForTraining = recs.copy()
- recsForTraining.remove(reg)
- recsForTest = [reg]
- Rval1 = []
- for d in range(regressionN): # loop over wheel speed, the four paw speeds and the combined speed
- #print('recording iteration %s, variable %s' %(reg,d))
- #pdb.set_trace()
- # concatenate data
- if borders is not None:
- timeMaskTrain = (caTracesDictInterP[recsForTraining[0]][0]>=borders[0])&(caTracesDictInterP[recsForTraining[0]][0]<=borders[1])
- timeMaskTest = (caTracesDictInterP[recsForTest[0]][0] >= borders[0]) & (caTracesDictInterP[recsForTest[0]][0] <= borders[1])
- else:
- timeMaskTrain = (caTracesDictInterP[recsForTraining[0]][0]>=0)&(caTracesDictInterP[recsForTraining[0]][0]<=1000.)
- timeMaskTest = (caTracesDictInterP[recsForTest[0]][0]>=0)&(caTracesDictInterP[recsForTest[0]][0]<=1000.)
- X = np.copy(caTracesDictInterP[recsForTraining[0]][1:][:,timeMaskTrain])
- Xtest = np.copy(caTracesDictInterP[recsForTest[0]][1:][:,timeMaskTest])
- #pdb.set_trace()
- if d == 0:
- Y = np.copy(wheelSpeedDictInterP[recsForTraining[0]][1:][:,timeMaskTrain])
- Ytest = np.copy(wheelSpeedDictInterP[recsForTest[0]][1:][:,timeMaskTest])
- YtestTime = np.copy(wheelSpeedDictInterP[recsForTest[0]][0][timeMaskTest])
- elif (d>0) and (d<5):
- pawId = d-1
- Y = np.copy(pawTracksDictInterP[recsForTraining[0]][pawId][1:][:,timeMaskTrain])
- Ytest = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][1:][:,timeMaskTest])
- YtestTime = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][0][timeMaskTest])
- elif d==5: # case where all four paw speeds are added together
- pawSpeedTrain = []
- pawSpeedTest = []
- for i in range(4):
- pawSpeedTrain.append(np.copy(pawTracksDictInterP[recsForTraining[0]][i][1:][:, timeMaskTrain]))
- pawSpeedTest.append(np.copy(pawTracksDictInterP[recsForTest[0]][i][1:][:, timeMaskTest]))
- Y = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
- Ytest = pawSpeedTest[0] + pawSpeedTest[1] + pawSpeedTest[2] + pawSpeedTest[3]
- #pdb.set_trace()
- for t in recsForTraining[1:]:
- if borders is not None:
- timeMaskTrain = (caTracesDictInterP[t][0] >= borders[0]) & (caTracesDictInterP[t][0] <= borders[1])
- else:
- timeMaskTrain = (caTracesDictInterP[t][0] >= 0) & (caTracesDictInterP[t][0] <= 1000.)
- X = np.column_stack((X,caTracesDictInterP[t][1:][:,timeMaskTrain]))
- if d == 0:
- Y = np.column_stack((Y,wheelSpeedDictInterP[t][1:][:,timeMaskTrain]))
- elif (d>0) and (d<5):
- Y = np.column_stack((Y, pawTracksDictInterP[t][pawId][1:][:,timeMaskTrain]))
- elif d==5:
- pawSpeedTrain = []
- for i in range(4):
- pawSpeedTrain.append(np.copy(pawTracksDictInterP[t][i][1:][:, timeMaskTrain]))
- speedTemp = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
- Y = np.column_stack((Y, speedTemp))
- #pdb.set_trace()
- Y = Y[0]
- X = np.transpose(X)
- Ytest = Ytest[0]
- Xtest = np.transpose(Xtest)
- # linear regression ########################################
- linReg = LinearRegression()
- linReg.fit(X,Y)
- #svm_rbf = SVR(kernel='rbf', C=1e3, gamma=0.1)
- #svm_rbf.fit(X,Y)
- YTrainPred = linReg.predict(X)
- YTestPred = linReg.predict(Xtest)
- R2trainLR = linReg.score(X, Y)
- R2testLR = linReg.score(Xtest, Ytest) # 1. - np.sum((Ytest-YTestPred)**2)/np.sum((Ytest - np.mean(Ytest))**2)#linReg.score(Xtest, Ytest)
- #print(linReg.coef_)
- #print(linReg.intercept_)
- #yPred = linReg.predict(np.transpose(X))
- # random forest ##############################################
- #randForestReg = RandomForestRegressor(n_estimators=20)
- #randForestReg.fit(X, Y)
- #R2trainRF = randForestReg.score(X, Y)
- #R2testRF= randForestReg.score(Xtest, Ytest)
- #
- if figShow :
- print('R2 test :',R2testLR)
- fig = plt.figure()
- ax = fig.add_subplot(111)
- ax.plot(YtestTime,Ytest,lw=2)
- ax.plot(YtestTime,YTestPred,lw=2)
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- ax.spines['bottom'].set_position(('outward', 10))
- ax.spines['left'].set_position(('outward', 10))
- ax.yaxis.set_ticks_position('left')
- ax.xaxis.set_ticks_position('bottom')
- plt.show()
- Rval1.extend([R2trainLR,R2testLR])
- RTempValues.append(Rval1)
- #pdb.set_trace()
- Rs = np.zeros(regressionN*2)
- for reg in range(5):
- Rs += RTempValues[reg]
- Rs /=5.
- Rvalues.append(Rs)
- return Rvalues
- #################################################################################
- # perform linear regression between spiking activity and behavioral measures
- #################################################################################
- def crossValidatedRegression(regModels,X,y,t,fold,visualize=False):
- cols = ['C0','C1','C2','C3','C4','C5','C6','C7','C8','C9','C10']
- regOutput = {}
- ###########
- if visualize:
- plt.plot(t,y,color='black')
- # get coeffss : apply the regression models consecutively
- for j in range(len(regModels)):
- regOutput[j] = {}
- print('applying', regModels[j][0])
- regModel = regModels[j][1]
- regModel.fit(X, y)
- # print(regModel.coef_)
- regOutput[j]['name'] = regModels[j][0]
- regOutput[j]['coefficients'] = regModel.coef_
- regOutput[j]['fitScore'] = regModel.score(X, y)
- regOutput[j]['scores'] = np.zeros(fold*5)
- if visualize:
- plt.plot(t,regModel.predict(X),c=cols[j],label=regModels[j][0]+': %s ' % (np.round(regOutput[j]['fitScore'],3)))
- del regModel
- if visualize:
- plt.xlabel('time (s)')
- plt.ylabel('firing rate')
- plt.legend(frameon=False)
- plt.show()
- ########
- if fold>0:
- print('performing cross-validation')
- # to k-fold cross-validated regression to access the score
- for j in range(len(regModels)):
- #kf = KFold(n_splits=fold)
- i = 0
- #scores = cross_val_score(regModels[j][1], X, y, cv=fold)
- #print(regModels[j][0],scores)
- rkf = RepeatedKFold(n_splits=fold, n_repeats=5, random_state=42)
- #for train_index, test_index in kf.split(X):
- for train_index, test_index in rkf.split(X):
- #print(i,len(train_index),len(test_index))
- X_train, X_test = X[train_index], X[test_index]
- y_train, y_test = y[train_index], y[test_index]
- regModel = regModels[j][1]
- regModel.fit(X_train, y_train)
- regOutput[j]['scores'][i] = regModel.score(X_test, y_test)
- #if visualize:
- # plt.plot(t[test_index],regModel.predict(X_test),label=(regModels[j][0] if i==0 else None))
- del regModel
- i+=1
- for j in range(len(regModels)):
- print(regOutput[j]['name'],regOutput[j]['fitScore'],np.mean(regOutput[j]['scores']),np.std(regOutput[j]['scores']))#,regOutput[j]['scores'])
- return regOutput
- # # print(regModel.intercept_)
- # sc = regModel.score(Xregressors_scaled[tmask], YspikeCount[tmask])
- # print('score:', sc)
- # Ypred = regModel.predict(Xregressors_scaled[tmask])
- # del regModel
- # plt.plot(tbinCenters[tmask], Ypred, label=regs[j][0] + ' %s' % np.round(sc, 3))
- # plt.xlabel('time (s)')
- # plt.ylabel('firing rate') # pass
- #################################################################################
- # shuffles variable within chunks
- #################################################################################
- def shuffleVariable(y,ttime,dt,chunkSize):
- nChunk = int(chunkSize/dt)
- shuffleIterations = int(np.ceil(len(ttime)/nChunk))
- for i in range(shuffleIterations):
- np.random.shuffle(y[(i*nChunk):((i+1)*nChunk)])
- return y
- #################################################################################
- # perform linear regression between spiking activity and behavioral measures
- #################################################################################
- # cPawPos,pawSpeed,ephys,swingStanceDict,sTimes,linearSpeed)
- def performGLManalysis(date, rec, pawPos,pawSpeed,ephys,swingStanceD,sTimes,linearSpeed):
- matplotlib.use('TkAgg')
- # create spike-count vector
- spikeTimes = ephys[0]
- print('firing rate :', 1./np.mean(np.diff(spikeTimes)))
- dt = 0.01
- shiftRange = 0.2 # binary columns are shifted back- and forth in time by this delay in s
- nShift = int(shiftRange/dt)
- tbins = np.linspace(0., 60., int(60 / dt) + 1, endpoint=True)
- tbinCenters = (tbins[1:]+tbins[:-1])/2
- binnedspikes, _ = np.histogram(spikeTimes, tbins)
- spikecountwindow = 0.02
- nspikecountwindow = int(spikecountwindow / dt) # + 0.5)
- #YspikeCount = np.convolve(binnedspikes, np.ones(nspikecountwindow), 'same')
- binnedspikes=np.array(binnedspikes,dtype=float)
- YspikeCount = scipy.ndimage.gaussian_filter1d(binnedspikes, nspikecountwindow,axis=0) # convolve with Gaussian kernel
- print(len(YspikeCount),len(binnedspikes),nspikecountwindow)
- # create regressor matrix
- #Xregressors = np.zeros((len(tbinCenters),1+4*4))
- Xregressors = np.zeros((len(tbinCenters), 1))
- # interpolate wheel speed
- interpWheel = interp1d(sTimes, linearSpeed,fill_value='extrapolate')#,kind='cubic')
- Xregressors[:,0] = interpWheel(tbinCenters)
- # interpolate paw position and paw speed
- for i in range(4):
- interpPawPos = interp1d(pawPos[i][:,0],pawPos[i][:,1],fill_value='extrapolate')
- #Xregressors[:,1+i] = interpPawPos(tbinCenters)
- Xregressors = np.column_stack((Xregressors,interpPawPos(tbinCenters)))
- for i in range(4):
- interpPawSpeed = interp1d(pawSpeed[i][:,0],pawSpeed[i][:,2],fill_value='extrapolate') # pawSpeed[i][1] is total speed, 2 is x, 3 is y speed
- #Xregressors[:,5+i] = interpPawSpeed(tbinCenters)
- Xregressors = np.column_stack((Xregressors, interpPawSpeed(tbinCenters)))
- for i in range(4):
- idxSwings = swingStanceD['swingP'][i][1]
- recTimes = swingStanceD['forFit'][i][2]
- idxSwings = np.asarray(idxSwings)
- binnedSwingStartTimes, _ = np.histogram(recTimes[idxSwings[:,0]], tbins)
- binnedSwingEndTimes, _ = np.histogram(recTimes[idxSwings[:, 1]], tbins)
- startTshift = np.zeros((len(binnedSwingStartTimes),2*nShift+1))
- endTshift = np.zeros((len(binnedSwingEndTimes),2*nShift+1))
- n = 0
- for j in range(-nShift,nShift+1):
- startTshift[:,n] = np.roll(binnedSwingStartTimes,j)
- endTshift[:,n] = np.roll(binnedSwingEndTimes, j)
- #plt.plot(startTshift[:,n])
- n+=1
- #plt.show()
- Xregressors = np.column_stack((Xregressors, startTshift))
- Xregressors = np.column_stack((Xregressors, endTshift))
- #Xregressors[:,9+i] = binnedSwingStartTimes
- #Xregressors[:,13+i] = binnedSwingEndTimes
- #pdb.set_trace()
- print('shape or regressor matrix :', np.shape(Xregressors))
- ## preprocessing of the data
- #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.
- #Xregressors_scaled = scaler.transform(Xregressors)
- #pdb.set_trace()
- Xregressors_scaled = np.copy(Xregressors)
- Xregressors_scaled[:,:9] = scipy.stats.zscore(Xregressors[:,:9],axis=0)
- #Xregressors_scaled = np.copy(Xregressors)
- #pdb.set_trace()
- #Xregressors_scaled[:8] = Xregressors[:8] # preserve the sparse data
- timeLimits = [10,50]
- tmask = (tbinCenters>timeLimits[0]) & (tbinCenters<timeLimits[1])
- # generate list of regression models
- # alpha multiplies the penalty terms : for alpha=0 is equivalent to an ordinary least square
- # ridge regression : l2 regularization
- # 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.
- 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])))]
- #regs = [('Elastic Net Regression',linear_model.ElasticNet(alpha=0.01,l1_ratio=0.5,random_state=0))]\
- # ('Ridge regression', linear_model.Ridge(alpha=1.)),
- # ('Elastic Net Regression', linear_model.ElasticNet(alpha=0.01, random_state=0))]
- regResults = crossValidatedRegression(regs,Xregressors_scaled[tmask],YspikeCount[tmask],tbinCenters[tmask],fold=0,visualize=False)
- print(len(Xregressors_scaled[tmask]))
- #plt.plot(tbinCenters[tmask],YspikeCount[tmask])
- #coeffss = []
- #pdb.set_trace()
- #plt.legend(frameon=False)
- #plt.show()
- return regResults
- plt.clf()
- fig = plt.figure(figsize=(12,4))
- plt.subplots_adjust(left=0.05, right=0.96, top=0.94, bottom=0.1)
- cols = ['C0','C1','C2','C3','C4']
- pawID = ['FL','FR','HL','HR']
- ax = fig.add_subplot(1,5,1)
- for i in range(len(regs)):
- ax.plot(regResults[i]['coefficients'][:9],'o-',label=regs[i][0])
- ax.set_ylabel('beta-weight')
- 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)
- #plt.setp(ax.get_xticklabels(), rotation=45, ha="right",rotation_mode="anchor")
- plt.legend(frameon=False)
- tVector = np.linspace(-nShift,nShift,2*nShift+1,endpoint=True)*dt
- shifts = 2*nShift + 1
- for j in range(4):
- ax = fig.add_subplot(1,5,j+2)
- for i in range(len(regs)):
- ax.set_title(pawID[j])
- 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]))
- 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]))
- #ax.set_ylim(-0.7,1.6)
- ax.axvline(x=0,ls=':',c='0.4')
- ax.set_xlabel('time (s)')
- ax.set_ylabel('beta-weight')
- plt.legend(frameon=False)
- plt.show()
- pdb.set_trace()
- # perform linear regression : no regularization
- print('Linear regression')
- Lreg = linear_model.LinearRegression()
- Lreg.fit(Xregressors_scaled, YspikeCount)
- print(Lreg.coef_)
- print(Lreg.intercept_)
- print('score:',Lreg.score(Xregressors_scaled,YspikeCount))
- Ypred = Lreg.predict(Xregressors_scaled)
- plt.title('Linear reg.')
- plt.plot(YspikeCount)
- plt.plot(Ypred)
- plt.show()
- def constructDesignMatrix(pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
- Xregressors = np.zeros((len(tbinCenters), 1))
- # interpolate wheel speed
- interpWheel = interp1d(sTimes, linearSpeed, fill_value='extrapolate') # ,kind='cubic')
- newWheelSpeed = interpWheel(tbinCenters)
- if (shuffle) and ('wheelSpeed' in variable):
- newWheelSpeed = shuffleVariable(newWheelSpeed,tbinCenters,dt,chunkSize)
- Xregressors[:, 0] = newWheelSpeed
- # interpolate paw position and paw speed
- for i in range(4):
- interpPawPos = interp1d(pawPos[i][:, 0], pawPos[i][:, 1], fill_value='extrapolate')
- # Xregressors[:,1+i] = interpPawPos(tbinCenters)
- newPawPos = interpPawPos(tbinCenters)
- if (shuffle) and (i in pawID) and ('pawPosition' in variable):
- newPawPos = shuffleVariable(newPawPos,tbinCenters,dt,chunkSize)
- Xregressors = np.column_stack((Xregressors, newPawPos))
- for i in range(4):
- interpPawSpeed = interp1d(pawSpeed[i][:, 0], pawSpeed[i][:, 2], fill_value='extrapolate') # pawSpeed[i][1] is total speed, 2 is x, 3 is y speed
- # Xregressors[:,5+i] = interpPawSpeed(tbinCenters)
- newPawSpeed = interpPawSpeed(tbinCenters)
- if (shuffle) and (i in pawID) and ('pawSpeed' in variable):
- newPawSpeed = shuffleVariable(newPawSpeed, tbinCenters, dt, chunkSize)
- Xregressors = np.column_stack((Xregressors, newPawSpeed))
- for i in range(4):
- idxSwings = swingStanceD['swingP'][i][1]
- recTimes = swingStanceD['forFit'][i][2]
- idxSwings = np.asarray(idxSwings)
- binnedSwingStartTimes, _ = np.histogram(recTimes[idxSwings[:, 0]], tbins)
- binnedSwingEndTimes, _ = np.histogram(recTimes[idxSwings[:, 1]], tbins)
- if (shuffle) and (i in pawID) and ('swingStart' in variable):
- binnedSwingStartTimes = shuffleVariable(binnedSwingStartTimes, tbinCenters, dt, chunkSize)
- if (shuffle) and (i in pawID) and ('stanceStart' in variable):
- binnedSwingEndTimes = shuffleVariable(binnedSwingEndTimes, tbinCenters, dt, chunkSize)
- startTshift = np.zeros((len(binnedSwingStartTimes), 2 * nShift + 1))
- endTshift = np.zeros((len(binnedSwingEndTimes), 2 * nShift + 1))
- n = 0
- for j in range(-nShift, nShift + 1):
- startTshift[:, n] = np.roll(binnedSwingStartTimes, j)
- endTshift[:, n] = np.roll(binnedSwingEndTimes, j)
- n += 1
- Xregressors = np.column_stack((Xregressors, startTshift))
- Xregressors = np.column_stack((Xregressors, endTshift)) # Xregressors[:,9+i] = binnedSwingStartTimes # Xregressors[:,13+i] = binnedSwingEndTimes
- return Xregressors
- matplotlib.use('TkAgg')
- # create spike-count vector
- spikeTimes = ephys[0]
- print('firing rate :', 1./np.mean(np.diff(spikeTimes)))
- dt = 0.01
- shiftRange = 0.2 # binary columns are shifted back- and forth in time by this delay in s
- nShift = int(shiftRange/dt)
- tbins = np.linspace(0., 60., int(60 / dt) + 1, endpoint=True)
- tbinCenters = (tbins[1:]+tbins[:-1])/2
- binnedspikes, _ = np.histogram(spikeTimes, tbins)
- spikecountwindow = 0.03 # sigma of the Gaussian kernel
- nspikecountwindow = int(spikecountwindow / dt + 0.5)
- binnedspikes = np.array(binnedspikes,dtype=float)
- #YspikeCount = scipy.ndimage.gaussian_filter1d(binnedspikes, nspikecountwindow,axis=0) # convolve with Gaussian kernel
- YspikeCount = np.convolve(binnedspikes, np.ones(nspikecountwindow), 'same') # convolve spike-count with square kernel
- #print(len(YspikeCount),len(binnedspikes),nspikecountwindow)
- #plt.hist(YspikeCount,bins=30)
- #plt.show()
- #pdb.set_trace()
- # create regressor matrix
- #Xregressors = np.zeros((len(tbinCenters),1+4*4))
- timeLimits = [10,50]
- tmask = (tbinCenters>timeLimits[0]) & (tbinCenters<timeLimits[1])
- regs = [('Ridge Regression', linear_model.Ridge(alpha=1))]
- Xregressors = constructDesignMatrix()
- #pdb.set_trace()
- print('shape or regressor matrix :', np.shape(Xregressors))
- ## preprocessing of the data
- 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.
- Xregressors_scaled = scaler.transform(Xregressors)
- regResultsFullModel = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=10, visualize=False)
- coffs = regResultsFullModel[0]['coefficients']
- shuffleForWeights = False
- if shuffleForWeights:
- shuffCoeffsWheelSpeed = []
- for i in range(100):
- # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
- Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['wheelSpeed'],chunkSize=2.,shuffle=True)
- # pdb.set_trace()
- #print('shape or regressor matrix :', np.shape(Xregressors))
- ## preprocessing of the data
- 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.
- Xregressors_scaled = scaler.transform(Xregressors)
- regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
- shuffCoeffsWheelSpeed.append(regResultsShuffle[0]['coefficients'])
- shuffCoeffsPawPos = []
- for i in range(100):
- # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
- Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['pawPosition'],chunkSize=2.,shuffle=True)
- # pdb.set_trace()
- #print('shape or regressor matrix :', np.shape(Xregressors))
- ## preprocessing of the data
- 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.
- Xregressors_scaled = scaler.transform(Xregressors)
- regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
- shuffCoeffsPawPos.append(regResultsShuffle[0]['coefficients'])
- #
- shuffCoeffsPawSpeed = []
- for i in range(100):
- # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
- Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['pawSpeed'],chunkSize=2.,shuffle=True)
- # pdb.set_trace()
- #print('shape or regressor matrix :', np.shape(Xregressors))
- ## preprocessing of the data
- 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.
- Xregressors_scaled = scaler.transform(Xregressors)
- regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
- shuffCoeffsPawSpeed.append(regResultsShuffle[0]['coefficients'])
- shuffCoeffsSwingStance = []
- for i in range(100):
- # (pawID=0,variable='wheelSpeed',chunkSize=3,shuffle=False):
- Xregressors = constructDesignMatrix(pawID=[0,1,2,3],variable=['swingStart','stanceStart'],chunkSize=2.,shuffle=True)
- # pdb.set_trace()
- #print('shape or regressor matrix :', np.shape(Xregressors))
- ## preprocessing of the data
- 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.
- Xregressors_scaled = scaler.transform(Xregressors)
- regResultsShuffle = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
- shuffCoeffsSwingStance.append(regResultsShuffle[0]['coefficients'])
- coffs = np.asarray(coffs)
- shuffCoeffsWheelSpeed = np.asarray(shuffCoeffsWheelSpeed)
- shuffCoeffsPawPos = np.asarray(shuffCoeffsPawPos)
- shuffCoeffsPawSpeed = np.asarray(shuffCoeffsPawSpeed)
- shuffCoeffsSwingStance = np.asarray(shuffCoeffsSwingStance)
- shuffCoeffs = np.copy(shuffCoeffsWheelSpeed)
- shuffCoeffs[:,1:5] = shuffCoeffsPawPos[:,1:5]
- shuffCoeffs[:,5:10] = shuffCoeffsPawSpeed[:,5:10]
- shuffCoeffs[:, 10:] = shuffCoeffsSwingStance[:,10:]
- #pdb.set_trace()
- fig = plt.figure(figsize=(12,4))
- plt.subplots_adjust(left=0.05, right=0.96, top=0.94, bottom=0.1)
- cols = ['C0','C1','C2','C3','C4']
- #pawID = ['FL','FR','HL','HR']
- ax0 = fig.add_subplot(1,5,1)
- ax0.axhline(y=0,ls='--',c='0.5')
- ax0.plot(coffs[:9],'o-')
- ax0.plot(np.mean(shuffCoeffs,axis=0)[:9])
- ax0.fill_between(np.arange(9),np.percentile(shuffCoeffs,5,axis=0)[:9],np.percentile(shuffCoeffs,95,axis=0)[:9],alpha=0.5)
- #ax1 = fig.add_subplot(1,2,2)
- tVector = np.linspace(-nShift,nShift,2*nShift+1,endpoint=True)*dt
- shifts = 2*nShift + 1
- for j in range(4):
- ax1 = fig.add_subplot(1,5,j+2)
- for i in range(len(regs)):
- #ax.set_title(pawID[j])
- ax1.plot(tVector,coffs[(9+(2*j)*shifts):(9+(2*j+1)*shifts)])
- 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]))
- 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)
- plt.show()
- pdb.set_trace()
- #Xregressors_scaled = np.copy(Xregressors)
- #pdb.set_trace()
- #Xregressors_scaled[:8] = Xregressors[:8] # preserve the sparse data
- # generate list of regression models
- # alpha multiplies the penalty terms : for alpha=0 is equivalent to an ordinary least square
- # ridge regression : l2 regularization
- # 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.
- 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])))]
- #regs = [('Ridge Regression',linear_model.Ridge(alpha=1.))]#,('GLM with log link function',linear_model.PoissonRegressor(alpha=1e-6/len(YspikeCount[tmask])))]
- # ('Ridge regression', linear_model.Ridge(alpha=1.)),
- # ('Elastic Net Regression', linear_model.ElasticNet(alpha=0.01, random_state=0))]
- RRidge = []
- #import statsmodels.api as sm
- #gamma_model = sm.GLM(YspikeCount[tmask],Xregressors_scaled[tmask], family=sm.families.Poisson())
- #gamma_results = gamma_model.fit()
- #print(gamma_results.summary())
- #pdb.set_trace()
- checkAlphaVariable = False
- if checkAlphaVariable :
- for i in range(100):
- al = float(i) #1./(1.1220184543019633**i)
- #regs = [('GLM with log link function',linear_model.PoissonRegressor(alpha=al))]
- regs = [('Ridge Regression', linear_model.Ridge(alpha=al))]
- regResults = crossValidatedRegression(regs,Xregressors_scaled[tmask],YspikeCount[tmask],tbinCenters[tmask],fold=10,visualize=False)
- RRidge.append([i,al,regResults[0]['fitScore'],np.mean(regResults[0]['scores'])])
- RRidge = np.asarray(RRidge)
- plt.title('Ridge Regression with log-link function')
- plt.plot(RRidge[:,1], RRidge[:,2], 'o-',label='fit score')
- plt.plot(RRidge[:,1],RRidge[:,3],'o-',label='cross validated score')
- plt.xlabel('alpha (penalty weight)')
- plt.ylabel('R^2')
- #plt.xscale('log')
- #plt.legend(frameon=False)
- plt.show()
- pdb.set_trace()
- print(len(Xregressors_scaled[tmask]))
- #plt.plot(tbinCenters[tmask],YspikeCount[tmask])
- #coeffss = []
- #pdb.set_trace()
- #plt.legend(frameon=False)
- #plt.show()
- 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])))]
- #regResultsFullModel = {}
- #for i in range(len(regs)):
- regResultsFullModel = crossValidatedRegression(regs, Xregressors_scaled[tmask], YspikeCount[tmask], tbinCenters[tmask], fold=0, visualize=False)
- plt.clf()
- fig = plt.figure(figsize=(12,4))
- plt.subplots_adjust(left=0.05, right=0.96, top=0.94, bottom=0.1)
- cols = ['C0','C1','C2','C3','C4']
- pawID = ['FL','FR','HL','HR']
- ax = fig.add_subplot(1,5,1)
- for i in range(len(regs)):
- ax.plot(regResultsFullModel[i]['coefficients'][:9],'o-',label=regs[i][0])
- ax.set_ylabel('beta-weight')
- 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)
- #plt.setp(ax.get_xticklabels(), rotation=45, ha="right",rotation_mode="anchor")
- plt.legend(frameon=False)
- tVector = np.linspace(-nShift,nShift,2*nShift+1,endpoint=True)*dt
- shifts = 2*nShift + 1
- for j in range(4):
- ax = fig.add_subplot(1,5,j+2)
- for i in range(len(regs)):
- ax.set_title(pawID[j])
- 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]))
- 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]))
- ax.set_ylim(-0.7,1.6)
- ax.axvline(x=0,ls=':',c='0.4')
- ax.set_xlabel('time (s)')
- ax.set_ylabel('beta-weight')
- plt.legend(frameon=False)
- plt.show()
- pdb.set_trace()
- # perform linear regression : no regularization
- print('Linear regression')
- Lreg = linear_model.LinearRegression()
- Lreg.fit(Xregressors_scaled, YspikeCount)
- print(Lreg.coef_)
- print(Lreg.intercept_)
- print('score:',Lreg.score(Xregressors_scaled,YspikeCount))
- Ypred = Lreg.predict(Xregressors_scaled)
- plt.title('Linear reg.')
- plt.plot(YspikeCount)
- plt.plot(Ypred)
- plt.show()
- # perform elastic net regression : Linear regression with combined L1 and L2 priors as regularizer
- # 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.
- print('Elastic net regression')
- ENreg = linear_model.ElasticNet(alpha=0.1,random_state=0)
- ENreg.fit(Xregressors_scaled, YspikeCount)
- print(ENreg.coef_)
- print(ENreg.intercept_)
- print('score:',ENreg.score(Xregressors_scaled,YspikeCount))
- Ypred = ENreg.predict(Xregressors_scaled)
- plt.title('Elastic net reg.')
- plt.plot(YspikeCount)
- plt.plot(Ypred)
- plt.show()
- #pdb.set_trace()
- # perform Generalized Linear Model with a Poisson distribution. This regressor uses the ‘log’ link function.
- # from sklearn.ensemble import HistGradientBoostingRegressor
- print('GLM with log link function')
- #Preg = HistGradientBoostingRegressor(loss="poisson",l2_regularization=1, max_leaf_nodes=128) #
- Preg = linear_model.PoissonRegressor(alpha=0)
- Preg.fit(Xregressors_scaled, YspikeCount)
- print(Preg.coef_)
- print(Preg.intercept_)
- print('score:',Preg.score(Xregressors_scaled, YspikeCount))
- Ypred = Preg.predict(Xregressors_scaled)
- plt.title('GLM with log link function')
- plt.plot(YspikeCount)
- plt.plot(Ypred)
- plt.show()
- pdb.set_trace()
- #################################################################################
- # calculate correlations between ca-imaging, wheel speed and paw speed
- #################################################################################
- def doContinuousRegressionAnalysis(mouse,allCorrDataPerSession,allStepData,borders=None,figShow=False):
- matplotlib.use('TkAgg') # WxAgg
- from sklearn.linear_model import LinearRegression
- from sklearn.model_selection import train_test_split
- from sklearn.svm import SVR
- #SVR(kernel='rbf', C=1e3, gamma=0.1)
- #from sklearn.ensemble import RandomForestRegressor
- regressionN = 6
- Rvalues = []
- #for nSess in range(len(allCorrDataPerSession)):
- for nDay in range(len(allCorrDataPerSession)):
- print(nDay,allCorrDataPerSession[nDay]['folder'])
- (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP, wheelSpeedDict, pawTracksDict, caTracesDict, slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay,allCorrDataPerSession,allStepData)
- # ATTENTION : all of the arrays also contain a time array
- # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTimeSame]
- # dims of pawTracksDictInterP : [nSessions][nPaw][2][valuesOverTimeSame]
- # dims of caTracesDictInterP : [nSessions][nRois+1][valuesOverTimeSame]
- # dims of wheelSpeedDict : [nSessions][2][valuesOverTime]
- # dims of pawTracksDict : [nSessions][nPaw][2][valuesOverTime]
- # dims of caTracesDict : [nSessions][nRois+1][valuesOverTime]
- #(wheelSpeedDict, pawTracksDict, caTracesDict,aa,bb,cc) = getCaWheelPawInterpolatedDictsPerDay(nSess, allCorrDataPerSession)
- nRecWheel = len(wheelSpeedDictInterP)
- nRecPaw = len(pawTracksDictInterP)
- nRecCa = len(caTracesDictInterP)
- print('Recording length :', nRecWheel, nRecPaw, nRecCa)
- if (nRecWheel != nRecPaw ) or (nRecWheel != nRecCa):
- print('problem in number of recordings listed in dictionaries')
- # loop over 5 different regressions, each using a different combination of test and train samples
- recs = range(nRecCa)
- RTempValues = []
- #for reg in range(nRecCa): # loop over all recordings
- #recsForTraining = recs.copy()
- #recsForTraining.remove(reg)
- #recsForTest = [reg]
- varr = ['wheel speed','paw speed 0', 'paw speed 1','paw speed 2','paw speed 3','paw speed 0+1+2+3']
- for d in range(regressionN): # loop over wheel speed, the four paw speeds and the combined speed
- print('regressing %s' %varr[d])
- #pdb.set_trace()
- # concatenate data
- if borders is not None:
- timeMaskTrain = (caTracesDictInterP[0][0]>=borders[0])&(caTracesDictInterP[0][0]<=borders[1])
- timeMaskTest = (caTracesDictInterP[0][0] >= borders[0]) & (caTracesDictInterP[0][0] <= borders[1])
- else:
- timeMaskTrain = (caTracesDictInterP[0][0]>=0)&(caTracesDictInterP[0][0]<=1000.)
- timeMaskTest = (caTracesDictInterP[0][0]>=0)&(caTracesDictInterP[0][0]<=1000.)
- X = np.copy(caTracesDictInterP[0][1:][:,timeMaskTrain])
- #Xtest = np.copy(caTracesDictInterP[0][1:][:,timeMaskTest])
- #pdb.set_trace()
- if d == 0:
- Y = np.copy(wheelSpeedDictInterP[0][1:][:,timeMaskTrain])
- #Ytest = np.copy(wheelSpeedDictInterP[0][1:][:,timeMaskTest])
- #YtestTime = np.copy(wheelSpeedDictInterP[0][0][timeMaskTest])
- elif (d>0) and (d<5):
- pawId = d-1
- Y = np.copy(pawTracksDictInterP[0][pawId][1:][:,timeMaskTrain])
- #Ytest = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][1:][:,timeMaskTest])
- #YtestTime = np.copy(pawTracksDictInterP[recsForTest[0]][pawId][0][timeMaskTest])
- elif d==5: # case where all four paw speeds are added together
- pawSpeedTrain = []
- #pawSpeedTest = []
- for i in range(4):
- pawSpeedTrain.append(np.copy(pawTracksDictInterP[0][i][1:][:, timeMaskTrain]))
- #pawSpeedTest.append(np.copy(pawTracksDictInterP[recsForTest[0]][i][1:][:, timeMaskTest]))
- Y = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
- #Ytest = pawSpeedTest[0] + pawSpeedTest[1] + pawSpeedTest[2] + pawSpeedTest[3]
- #pdb.set_trace()
- for t in recs[1:]:
- if borders is not None:
- timeMaskTrain = (caTracesDictInterP[t][0] >= borders[0]) & (caTracesDictInterP[t][0] <= borders[1])
- else:
- timeMaskTrain = (caTracesDictInterP[t][0] >= 0) & (caTracesDictInterP[t][0] <= 1000.)
- X = np.column_stack((X,caTracesDictInterP[t][1:][:,timeMaskTrain]))
- if d == 0:
- Y = np.column_stack((Y,wheelSpeedDictInterP[t][1:][:,timeMaskTrain]))
- elif (d>0) and (d<5):
- Y = np.column_stack((Y, pawTracksDictInterP[t][pawId][1:][:,timeMaskTrain]))
- elif d==5:
- pawSpeedTrain = []
- for i in range(4):
- pawSpeedTrain.append(np.copy(pawTracksDictInterP[t][i][1:][:, timeMaskTrain]))
- speedTemp = pawSpeedTrain[0] + pawSpeedTrain[1] + pawSpeedTrain[2] + pawSpeedTrain[3]
- Y = np.column_stack((Y, speedTemp))
- #pdb.set_trace()
- Y = Y[0]
- X = np.transpose(X)
- #Ytest = Ytest[0]
- nRegressionIterations = 10
- Rval1 = np.zeros(2)
- for n in range(nRegressionIterations):
- #print(n,end='')
- X_train, X_test, y_train, y_test = train_test_split(X, Y, test_size = 0.2)
- #Xtest = np.transpose(Xtest)
- # linear regression ########################################
- linReg = LinearRegression()
- linReg.fit(X_train,y_train)
- #svm_rbf = SVR(kernel='rbf', C=1e3, gamma=0.1)
- #svm_rbf.fit(X,Y)
- #YTrainPred = linReg.predict(X)
- y_test_pred = linReg.predict(X_test)
- R2trainLR = linReg.score(X_train, y_train)
- R2testLR = linReg.score(X_test, y_test) # 1. - np.sum((Ytest-YTestPred)**2)/np.sum((Ytest - np.mean(Ytest))**2)#linReg.score(Xtest, Ytest)
- #print(linReg.coef_)
- #print(linReg.intercept_)
- #yPred = linReg.predict(np.transpose(X))
- # random forest ##############################################
- #randForestReg = RandomForestRegressor(n_estimators=20)
- #randForestReg.fit(X, Y)
- #R2trainRF = randForestReg.score(X, Y)
- #R2testRF= randForestReg.score(Xtest, Ytest)
- #
- if figShow :
- print('R2 test :',R2testLR)
- fig = plt.figure()
- ax = fig.add_subplot(111)
- ax.plot(y_test,lw=2)
- ax.plot(y_test_pred,lw=2)
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- ax.spines['bottom'].set_position(('outward', 10))
- ax.spines['left'].set_position(('outward', 10))
- ax.yaxis.set_ticks_position('left')
- ax.xaxis.set_ticks_position('bottom')
- plt.show()
- Rval1+= np.array([R2trainLR,R2testLR])
- RTempValues.extend(Rval1/nRegressionIterations)
- #pdb.set_trace()
- #Rs = np.zeros(regressionN*2)
- #for reg in range(nRegressionIterations):
- # Rs += RTempValues[reg]
- #Rs /=nRegressionIterations
- Rvalues.append([nDay,allCorrDataPerSession[nDay]['folder'],RTempValues])
- return Rvalues
- #################################################################################
- # calculate correlations between ca-imaging, wheel speed and paw speed
- #################################################################################
- def generateStepTriggeredCaTraces(mouse,allCorrDataPerSession,allStepData,trigger='swingOnset',calculateFast=False): # swingOnset or swingOffset
- matplotlib.use('TkAgg')
- # check for sanity
- if len(allCorrDataPerSession) != len(allStepData):
- print('both dictionaries are not of the same length')
- print('CaWheelPawDict:',len(allCorrDataPerSession),' StepStanceDict:,',len(allStepData))
- timeAxis = np.linspace(-0.4,0.6,int((0.4+0.6)/0.02)+1)
- timeAxisRescaled = np.linspace(-1.,2.,int((1.+2.)/0.02)+1)
- preStanceMask = timeAxis<-0.1
- preStanceRescaledMask = timeAxisRescaled<-0.2
- K = len(timeAxis)
- KRescaled = len(timeAxisRescaled)
- caTraces = []
- maxTimeDelay = 1.
- for nDay in range(len(allCorrDataPerSession)):
- print(allCorrDataPerSession[nDay]['folder'], allStepData[nDay][0], nDay)
- # consistency check
- if not (allCorrDataPerSession[nDay]['folder'] == allStepData[nDay][0]):
- print('All animal data and swing data not from the same day!')
- #
- caRecTime = allCorrDataPerSession[nDay]['caImg']['timeStamps'][0, 3]
- wheelRecTime = allStepData[nDay][1][0][3] # check recording start of first recording allStepData[nDay][1][3]
- pawRecTime = allStepData[nDay][2][0][4] # again, third indices picks first recording
- #pdb.set_trace()
- timeDiffCaWheel = np.abs(caRecTime-wheelRecTime)
- timeDiffCaPaw = np.abs(caRecTime-pawRecTime)
- timeDiffWheelPaw = np.abs(wheelRecTime-pawRecTime)
- if any(np.array([timeDiffCaWheel,timeDiffCaPaw,timeDiffWheelPaw]) > maxTimeDelay):
- print('PROBLEM in data consistency!')
- print('recordings are separated by %s - %s - %s s min' % (timeDiffCaWheel,timeDiffCaPaw,timeDiffWheelPaw))
- pdb.set_trace()
- else:
- print('Delay between recordings is :', timeDiffCaWheel,timeDiffCaPaw,timeDiffWheelPaw, 's')
- #print(allCorrDataPerSession[nDay][0],nDay) getCaWheelPawInterpolatedDictsPerDay
- (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP,wheelSpeedDict,pawTracksDict,caTracesDict,slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession,allStepData)
- #if len(allStepData[nDay-1][4])==6:
- # print('more recordings :',len(allStepData[nDay][4]))
- # addIdx = 1
- #else:
- # addIdx = 0
- #pdb.set_trace()
- N = len(caTracesDict[0][1:]) # number of ROIs
- 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))
- caSnippetsRescaled = [[[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)],[[] for i in range(N)]]
- recordingID = [[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)],[[] for i in range(N)]]
- recordingIDRescaled = [[[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)]]
- NRecs=len(allStepData[nDay][4])
- for nrec in range(NRecs): # loop over the five recordings of a day
- for i in range(4): # loop over the four paws
- #pdb.set_trace()
- idxSwings = allStepData[nDay][4][nrec][3][i][1]
- #print('Wow nSess',nDay,allStepData[nDay-1][0])
- recTimes = allStepData[nDay][4][nrec][4][i][2]
- #pdb.set_trace()
- idxSwings = np.asarray(idxSwings)
- if trigger == 'swingOnset':
- NstepCycles = len(idxSwings)
- elif trigger == 'swingOffset':
- NstepCyles = len(idxSwings)-1
- for k in range(NstepCyles): # loop over all swings
- startSwingTime = recTimes[idxSwings[k, 0]]
- endSwingTime = recTimes[idxSwings[k, 1]]
- if trigger == 'swingOnset':
- triggerTime = startSwingTime
- duration = (endSwingTime-startSwingTime)
- elif trigger == 'swingOffset':
- triggerTime = endSwingTime
- duration = recTimes[idxSwings[k+1, 0]] - endSwingTime
- if len(caTracesDict[nrec][1:])!=N:
- print('problem in number of ROIs')
- pdb.set_trace(0)
- for l in range(len(caTracesDict[nrec][1:])): # loop over all ROIs
- interpCa = interp1d(caTracesDict[nrec][0]-triggerTime, caTracesDict[nrec][l+1])#,kind='cubic')
- interpCaRescaled = interp1d((caTracesDict[nrec][0]-triggerTime)/(duration), caTracesDict[nrec][l+1])#,kind='cubic')
- ############
- try:
- newCaTraceAtSwing = interpCa(timeAxis)
- except ValueError:
- pass
- else:
- #caSnippets[i,l,:] += newCaTraceAtSwing
- caSnippets[i][l].append(newCaTraceAtSwing)
- recordingID[i][l].append(nrec)
- ############
- try:
- newCaTraceAtSwingRescaled = interpCaRescaled(timeAxisRescaled)
- except ValueError:
- #print('error')
- pass
- else:
- #caSnippets[i,l,:] += newCaTraceAtSwing
- caSnippetsRescaled[i][l].append(newCaTraceAtSwingRescaled)
- recordingIDRescaled[i][l].append(nrec)
- #pdb.set_trace()
- # 4 paws
- # N number of ROIS
- # 2 mean and std
- # K number of time points during average
- alpha = 0.05
- def get_CI(dat, alpha=.05):
- return np.array(list(map(lambda x: boot.ci(x, alpha=alpha), dat)))
- caSnippetsArray = np.zeros((4, N, 3 + 3*NRecs, K))
- caSnippetsRescaledArray = np.zeros((4, N, 3 + 3*NRecs, KRescaled))
- for i in range(4): # loop over four paws
- for l in range(N): # loop over all ROIs
- caTempArray = np.asarray(caSnippets[i][l])
- #pdb.set_trace()
- caSnippetsZscores = (caTempArray - np.mean(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]) #/np.std(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]
- if calculateFast:
- CI = np.zeros((len(timeAxis),2))
- else:
- print('calculating CIs of global average ...')
- CI = np.array(list(map(lambda x: boot.ci(x, alpha=.05), caSnippetsZscores.T)))
- print('done')
- caTemp = np.mean(caSnippetsZscores,axis=0)
- #caTempSTD = np.std(caSnippetsZscores,axis=0)
- caSnippetsArray[i,l,0,:] = caTemp
- caSnippetsArray[i,l,1,:] = CI[:,0]
- caSnippetsArray[i,l,2,:] = CI[:,1]
- caTempRec = []
- for nrec in range(NRecs):
- recMask = (np.asarray(recordingID[i][l]) == nrec)
- caTempRec.append(caSnippetsZscores[recMask])
- if calculateFast:
- ci_cell_trial = np.zeros((NRecs, len(timeAxis), 2))
- else:
- print('calculating CIs of recording average for paw %s and roi %s ...' % (i, l))
- ci_cell_trial = np.array(Parallel(n_jobs=numcores)(delayed(get_CI)(c.T, alpha) for c in caTempRec))
- print('done')
- for nrec in range(NRecs):
- caSnippetsArray[i, l, 3 + nrec*3, :] = np.mean(caTempRec[nrec], axis=0)
- caSnippetsArray[i, l, 4 + nrec*3, :] = ci_cell_trial[nrec][:,0]
- caSnippetsArray[i, l, 5 + nrec*3, :] = ci_cell_trial[nrec][:,1]
- #pdb.set_trace()
- caTempRescaledArray = np.asarray(caSnippetsRescaled[i][l])
- #pdb.set_trace()
- caSnippetsRescaledZscores = (caTempRescaledArray - np.mean(caTempRescaledArray[:,preStanceRescaledMask],axis=1)[:,np.newaxis])# /np.std(caTempRescaledArray[:,preStanceRescaledMask],axis=1)[:,np.newaxis]
- caTempRe = np.mean(caSnippetsRescaledZscores,axis=0)
- if calculateFast:
- CI = np.zeros((len(timeAxisRescaled),2))
- else:
- print('calculating CIs of global rescaled recording average')
- CI = np.array(list(map(lambda x: boot.ci(x, alpha=.05), caSnippetsRescaledZscores.T)))
- print('done')
- #caTempReSTD = np.std(caSnippetsRescaledZscores,axis=0)
- caSnippetsRescaledArray[i,l,0,:] = caTempRe
- caSnippetsRescaledArray[i,l,1,:] = CI[:,0]
- caSnippetsRescaledArray[i,l,2,:] = CI[:,1]
- caTempRecRescaled = []
- for nrec in range(NRecs):
- recMask = (np.asarray(recordingIDRescaled[i][l]) == nrec)
- caTempRecRescaled.append(caSnippetsRescaledZscores[recMask])
- #caTemp = np.mean(caSnippetsRescaledZscores[recMask], axis=0)
- #caTempSTD = np.std(caSnippetsRescaledZscores[recMask], axis=0)
- #CI = np.array(list(map(lambda x: boot.ci(x, alpha=.05), caSnippetsRescaledZscores[recMask].T)))
- if calculateFast :
- ci_cell_trial_rescaled = np.zeros((NRecs,len(timeAxisRescaled),2))
- else:
- print('calculating CIs of rescaled recording average for paw %s and roi %s ...' % (i, l))
- ci_cell_trial_rescaled = np.array(Parallel(n_jobs=numcores)(delayed(get_CI)(c.T, alpha) for c in caTempRecRescaled))
- print('done')
- for nrec in range(NRecs):
- caSnippetsRescaledArray[i, l, 3 + nrec*3, :] = np.mean(caTempRecRescaled[nrec], axis=0)
- caSnippetsRescaledArray[i, l, 4 + nrec*3, :] = ci_cell_trial_rescaled[nrec][:,0]
- caSnippetsRescaledArray[i, l, 5 + nrec*3, :] = ci_cell_trial_rescaled[nrec][:,1]
- caTraces.append([allCorrDataPerSession[nDay]['folder'],allStepData[nDay][0],nDay,timeAxis,caSnippetsArray,timeAxisRescaled,caSnippetsRescaledArray])
- return caTraces
- #################################################################################
- # calculate correlations between ca-imaging, wheel speed and paw speed
- #################################################################################
- def generateStepTriggeredCaTracesAllPaws(mouse,allCorrDataPerSession,allStepData):
- maxSeparation = 0.# min separation between swings in sec
- # check for sanity
- if len(allCorrDataPerSession) != len(allStepData):
- print('both dictionaries are not of the same length')
- print('CaWheelPawDict:',len(allCorrDataPerSession),' StepStanceDict:,',len(allStepData))
- timeAxis = np.linspace(-0.4,0.6,(0.6+0.4)/0.02+1)
- timeAxisRescaled = np.linspace(-1.,2.,(2+1)/0.02+1)
- preStanceMask = timeAxis<-0.1
- preStanceRescaledMask = timeAxisRescaled<-0.2
- K = len(timeAxis)
- KRescaled = len(timeAxisRescaled)
- caTraces = []
- swingT = []
- for nDay in range(len(allCorrDataPerSession)):
- (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP,wheelSpeedDict,pawTracksDict,caTracesDict) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession)
- if len(allStepData[nDay-1][4])==6:
- print('more recordings :',len(allStepData[nDay-1][4]))
- addIdx = 1
- else:
- addIdx = 0
- #pdb.set_trace()
- N = len(caTracesDict[0][1:])
- print(allCorrDataPerSession[nDay][0], allStepData[nDay - 1][0], nDay, N)
- caSnippets = [[] for i in range(N)] #np.zeros((4,N,K))
- caSnippetsRescaled = [[] for i in range(N)]
- caSnippetsArray = np.zeros((N,2,K))
- caSnippetsRescaledArray = np.zeros((N, 2, KRescaled))
- swingSnippets = [[] for i in range(5)]
- for nrec in range(5): # loop over the five recordings of a day
- swingTimes = np.zeros(3)
- for i in range(4): # loop over the four paws and lump all swing times together
- idxSwings = allStepData[nDay-1][4][nrec+addIdx][3][i][1]
- #print('Wow nSess',nDay,allStepData[nDay-1][0])
- recTimes = allStepData[nDay-1][4][nrec+addIdx][4][i][2]
- #pdb.set_trace()
- idxSwings = np.asarray(idxSwings)
- startSwingTimes = recTimes[idxSwings[:,0]]
- endSwingTimes = recTimes[idxSwings[:,1]]
- swingTimes = np.vstack((swingTimes,np.column_stack((startSwingTimes, endSwingTimes,np.repeat(i,len(startSwingTimes))))))
- swingTimes = swingTimes[1:] # remove first element which was zeros only
- # sort swing times according to swing start
- swingTimes = swingTimes[swingTimes[:,0].argsort()]
- # remove swings with fall within the minimum separation betweeen swings
- diffSwings = np.diff(swingTimes[:,0]) # calculate inter-swing intervals
- swingTimesSparse = swingTimes[np.concatenate((diffSwings>maxSeparation,np.array([True])))] # only use swings which fall above separation time
- swingSnippets[nrec] = swingTimes
- for k in range(len(swingTimesSparse)): # loop over all swings
- startSwingTime = swingTimesSparse[k, 0]
- endSwingTime = swingTimesSparse[k, 1]
- if len(caTracesDict[nrec][1:])!=N: print('problem in number of ROIs')
- for l in range(len(caTracesDict[nrec][1:])): # loop over all ROIs
- interpCa = interp1d(caTracesDict[nrec][0]-startSwingTime, caTracesDict[nrec][l+1])#,kind='cubic')
- interpCaRescaled = interp1d((caTracesDict[nrec][0]-startSwingTime)/(endSwingTime-startSwingTime), caTracesDict[nrec][l+1])#,kind='cubic')
- ############
- try:
- newCaTraceAtSwing = interpCa(timeAxis)
- except ValueError:
- pass
- else:
- #caSnippets[i,l,:] += newCaTraceAtSwing
- caSnippets[l].append(newCaTraceAtSwing)
- ############
- try:
- newCaTraceAtSwingRescaled = interpCaRescaled(timeAxisRescaled)
- except ValueError:
- #print('error')
- pass
- else:
- #caSnippets[i,l,:] += newCaTraceAtSwing
- caSnippetsRescaled[l].append(newCaTraceAtSwingRescaled)
- #pdb.set_trace()
- #for i in range(4):
- for l in range(N):
- caTempArray = np.asarray(caSnippets[l])
- caSnippetsZscores = (caTempArray - np.mean(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]) #/np.std(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]
- caTemp = np.mean(caSnippetsZscores,axis=0)
- caTempSTD = np.std(caSnippetsZscores,axis=0)
- caSnippetsArray[l,0,:] = caTemp
- caSnippetsArray[l,1,:] = caTempSTD
- #
- caTempRescaledArray = np.asarray(caSnippetsRescaled[l])
- #pdb.set_trace()
- caSnippetsRescaledZscores = (caTempRescaledArray - np.mean(caTempRescaledArray[:,preStanceRescaledMask],axis=1)[:,np.newaxis]) #/np.std(caTempArray[:,preStanceMask],axis=1)[:,np.newaxis]
- caTempRe = np.mean(caSnippetsRescaledZscores,axis=0)
- caTempReSTD = np.std(caSnippetsRescaledZscores,axis=0)
- caSnippetsRescaledArray[l,0,:] = caTempRe
- caSnippetsRescaledArray[l,1,:] = caTempReSTD
- swingT.append([allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay,swingSnippets])
- caTraces.append([allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay,caSnippetsArray,caSnippetsRescaledArray])
- return caTraces
- #################################################################################
- def calcualteAllDeltaTValues(t1,t2):
- l1 = len(t1)
- l2 = len(t2)
- fannedOut1 = np.tile(t1,(l2,1))
- fannedOut2 = np.tile(t2,(l1,1))
- transposedFannedOut2 = np.transpose(fannedOut2)
- #print l1, l2
- #print shape(fannedOut1), shape(transposedFannedOut2)
- #pdb.set_trace()
- differences = fannedOut1 - transposedFannedOut2 # subtract(fannedOut1,transposedFannedOut2)
- return differences.flatten()
- #################################################################################
- # calculate correlations between ca-imaging, wheel speed and paw speed
- #################################################################################
- def generateInterstepTimeHistogram(mouse,allCorrDataPerSession,allStepData):
- # check for sanity
- if len(allCorrDataPerSession) != len(allStepData):
- print('both dictionaries are not of the same length')
- print('CaWheelPawDict:',len(allCorrDataPerSession),' StepStanceDict:,',len(allStepData))
- #timeAxis = np.linspace(-0.4,0.6,(0.6+0.4)/0.02+1)
- #timeAxisRescaled = np.linspace(-1.,2.,(2+1)/0.02+1)
- #preStanceMask = timeAxis<-0.1
- #preStanceRescaledMask = timeAxisRescaled<-0.2
- #K = len(timeAxis)
- #KRescaled = len(timeAxisRescaled)
- pawSwingTimes = []
- for nDay in range(1,len(allCorrDataPerSession)):
- print(allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay)
- (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP,wheelSpeedDict,pawTracksDict,caTracesDict) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession)
- if len(allStepData[nDay-1][4])==6:
- print('more recordings :',len(allStepData[nDay-1][4]))
- addIdx = 1
- else:
- addIdx = 0
- #pdb.set_trace()
- #N = len(caTracesDict[0][1:])
- #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))
- #caSnippetsRescaled = [[[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)], [[] for i in range(N)]]
- #caSnippetsArray = np.zeros((4,N,2,K))
- #caSnippetsRescaledArray = np.zeros((4, N, 2, KRescaled))
- allData = []
- for nrec in range(5): # loop over the five recordings of a day
- startSwingTimes = [[] for i in range(4)]
- endSwingTimes = [[] for i in range(4)]
- for i in range(4): # loop over the four paws
- idxSwings = allStepData[nDay-1][4][nrec+addIdx][3][i][1]
- #print('Wow nSess',nDay,allStepData[nDay-1][0])
- recTimes = allStepData[nDay-1][4][nrec+addIdx][4][i][2]
- #pdb.set_trace()
- idxSwings = np.asarray(idxSwings)
- startSwingT = recTimes[idxSwings[:,0]]
- endSwingT = recTimes[idxSwings[:,1]]
- startSwingTimes[i].append(startSwingT)
- endSwingTimes[i].append(endSwingT)
- allData.append([startSwingTimes,endSwingTimes])
- interPawSwingTimes = [[] for i in range(4)]
- interStepTimes = []
- stepLengths = [[] for i in range(4)]
- #pdb.set_trace()
- for nrec in range(5):
- for i in range(4):
- interPawSwingTimes[i].extend(calcualteAllDeltaTValues(allData[nrec][0][i][0],allData[nrec][0][i][0]))
- stepLengths[i].extend(allData[nrec][1][i][0]-allData[nrec][0][i][0])
- interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][0][0],allData[nrec][0][1][0]))
- interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][0][0], allData[nrec][0][2][0]))
- interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][0][0], allData[nrec][0][3][0]))
- interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][1][0], allData[nrec][0][2][0]))
- interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][1][0], allData[nrec][0][3][0]))
- interStepTimes.extend(calcualteAllDeltaTValues(allData[nrec][0][2][0], allData[nrec][0][3][0]))
- #pdb.set_trace()
- pawSwingTimes.append([allCorrDataPerSession[nDay][0],allStepData[nDay-1][0],nDay,interPawSwingTimes,stepLengths,interStepTimes])
- return pawSwingTimes
- #################################################################################
- # remove empty columns and row - from the image registration routine
- #################################################################################
- def removeEmptyColumnAndRows(img):
- # 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
- # vmask = np.invert(np.sum(img, axis=1) == 0)
- htemp = (img == img[0,:]) # look instead for same values in row
- hmask = np.invert(htemp.all(axis=0))
- vtemp = (img == img[:,0])
- vmask = np.invert(vtemp.all(axis=1))
- idxH = np.arange(len(hmask))[np.hstack((False,np.diff(hmask)>0))]
- idxV = np.arange(len(vmask))[np.hstack((False, np.diff(vmask)>0))]
- #pdb.set_trace()
- if len(idxV) == 0:
- idxV = np.array([0,np.shape(img)[0]])
- if len(idxH) == 0:
- idxH = np.array([0,np.shape(img)[1]])
- croppedImg = img[idxV[0]:idxV[1],idxH[0]:idxH[1]]
- cutLengths = np.vstack((idxV,idxH))
- #pdb.set_trace()
- return cutLengths
- #################################################################################
- # remove empty columns and row - from the image registration routine
- #################################################################################
- def alignTwoImages(imgA,cutLengthsA,imgB,cutLengthsB,refDate,otherDate,movementValues,figSave=False,figDir=''):
- #matplotlib.use('TkAgg')
- column1 = np.maximum(cutLengthsA[:,0],cutLengthsB[:,0])
- column2 = np.minimum(cutLengthsA[:,1],cutLengthsB[:,1])
- cutLenghts = np.column_stack((column1,column2))
- imgA = imgA[cutLenghts[0,0]:cutLenghts[0,1],cutLenghts[1,0]:cutLenghts[1,1]]
- imgB = imgB[cutLenghts[0,0]:cutLenghts[0,1],cutLenghts[1,0]:cutLenghts[1,1]]
- # Find size of ref image
- sz = imgA.shape
- corr = signal.correlate(imgA - imgA.mean(), imgB - imgB.mean(), mode='same', method='fft')
- maxIdx = np.unravel_index(np.argmax(corr, axis=None), corr.shape)
- shifty = np.shape(imgA)[0]/2. - maxIdx[0]
- shiftx = np.shape(imgA)[1]/2. - maxIdx[1]
- print('max of cross-correlation : ', shiftx, shifty )
- #pdb.set_trace()
- # Define the motion model
- #warp_mode = cv2.MOTION_TRANSLATION #cv2.MOTION_EUCLIDEAN # cv2.MOTION_TRANSLATION # MOTION_EUCLIDEAN
- warp_mode = cv2.MOTION_AFFINE #EUCLIDEAN #HOMOGRAPHY
- warp_modes = [cv2.MOTION_AFFINE,cv2.MOTION_EUCLIDEAN,cv2.MOTION_TRANSLATION]
- # Define 2x3 or 3x3 matrices and initialize the matrix to identity
- if warp_mode == cv2.MOTION_HOMOGRAPHY:
- warp_matrix = np.eye(3, 3, dtype=np.float32)
- else:
- warp_matrix = np.eye(2, 3, dtype=np.float32)
- # Specify the number of iterations.
- number_of_iterations = 1000
- # Specify the threshold of the increment
- # in the correlation coefficient between two iterations
- termination_eps = 1e-10
- # Define termination criteria
- criteria = (cv2.TERM_CRITERIA_EPS | cv2.TERM_CRITERIA_COUNT, number_of_iterations, termination_eps)
- # Run the ECC algorithm. The results are stored in warp_matrix.
- #try:
- imA_u8 = (((imgA-np.min(imgA))/(np.max(imgA)-np.min(imgA)))*255).astype(np.uint8) #cv2.cvtColor(imgA,cv2.COLOR_BGR2GRAY)
- imB_u8 = (((imgB-np.min(imgB))/(np.max(imgB)-np.min(imgB)))*255).astype(np.uint8)
- #pdb.set_trace()
- warpResults = []
- corrMax = []
- for w in range(len(warp_modes)):
- print('testing ', warp_modes[w])
- # warp_matrix1[0, 2] = aS.xOffset
- # warp_matrix1[1, 2] = aS.yOffset
- if (movementValues[0] != 0) and (movementValues[1] != 0):
- warp_matrix[0, 2] = movementValues[0] # -20.
- warp_matrix[1, 2] = movementValues[1] # -40.
- else:
- warp_matrix[0, 2] = shiftx # -20.
- warp_matrix[1, 2] = shifty # -40.
- #try:
- # (cc, warp_matrixRet) = cv2.findTransformECC(imA_u8, imB_u8, warp_matrix, warp_modes[w], criteria, inputMask = None, gaussFiltSize=5)
- #except TypeError:
- try :
- (cc, warp_matrixRet) = cv2.findTransformECC(imA_u8, imB_u8, warp_matrix, warp_modes[w], criteria, inputMask = None)
- except:
- print('findTransformECC did not converge')
- cc = -1
- warp_matrixRet = warp_matrix
- #print(warp_matrixRet,warp_matrix)
- warpResults.append([w,warp_modes[w],np.copy(warp_matrixRet),np.copy(cc)])
- corrMax.append(cc)
- #if cc>0.8:
- # break
- print(warpResults)
- corrMax = np.asarray(corrMax)
- maxCorr = np.argmax(corrMax)
- cc_max = warpResults[maxCorr][3]
- warp_matrix_max = warpResults[maxCorr][2]
- #except:
- #print('findTransformECC output : ',cc,warp_matrixRet)
- #print('find image transformation did not converge')
- #warp_matrixRet = np.copy(warp_matrix_max)
- #cc = None
- #else:
- #pass
- #(cc2, warp_matrix2Ret) = cv2.findTransformECC(imBD, im820, warp_matrix2, warp_mode, criteria)
- if warp_mode == cv2.MOTION_HOMOGRAPHY:
- # Use warpPerspective for Homography
- imgB_aligned = cv2.warpPerspective(imgB, warp_matrix_max, (sz[1], sz[0]), flags=cv2.INTER_LINEAR + cv2.WARP_INVERSE_MAP)
- else:
- # Use warpAffine for Translation, Euclidean and Affine
- imgB_aligned = cv2.warpAffine(imgB, warp_matrix_max, (sz[1], sz[0]), flags=cv2.INTER_LINEAR + cv2.WARP_INVERSE_MAP);
- print('result of image alignment-> warp-matrix and correlation coefficient : ', warp_matrix_max, cc_max)
- if figSave :
- ##################################################################
- # Show final results
- # figure #################################
- fig_width = 10 # width in inches
- fig_height = 10 # height in inches
- fig_size = [fig_width, fig_height]
- params = {'axes.labelsize': 11, 'axes.titlesize': 11, 'font.size': 11, 'xtick.labelsize': 11, 'ytick.labelsize': 11, 'figure.figsize': fig_size, 'savefig.dpi': 600,
- 'axes.linewidth': 1.3, 'ytick.major.size': 4, # major tick size in points
- 'xtick.major.size': 4 # major tick size in points
- # 'edgecolor' : None
- # 'xtick.major.size' : 2,
- # 'ytick.major.size' : 2,
- }
- rcParams.update(params)
- # set sans-serif font to Arial
- rcParams['font.sans-serif'] = 'Arial'
- # create figure instance
- fig = plt.figure()
- # define sub-panel grid and possibly width and height ratios
- gs = gridspec.GridSpec(2, 2 # ,
- # width_ratios=[1.2,1]
- # height_ratios=[1,1]
- )
- # define vertical and horizontal spacing between panels
- gs.update(wspace=0.3, hspace=0.3)
- # possibly change outer margins of the figure
- plt.subplots_adjust(left=0.05, right=0.95, top=0.92, bottom=0.06)
- # sub-panel enumerations
- # plt.figtext(0.06, 0.92, 'A',clip_on=False,color='black', weight='bold',size=22)
- # first sub-plot #######################################################
- # gssub = gridspec.GridSpecFromSubplotSpec(1, 2, subplot_spec=gs[0],hspace=0.2)
- # ax0 = plt.subplot(gssub[0])
- # fig = plt.figure(figsize=(10,10))
- #plt.figtext(0.1, 0.95, '%s ' % (aS.animalID), clip_on=False, color='black', size=14)
- ax0 = plt.subplot(gs[0])
- ax0.set_title('reference image %s' % refDate)
- ax0.imshow(imgA)
- ax0 = plt.subplot(gs[1])
- ax0.set_title('to-be-aligned image %s' % otherDate )
- ax0.imshow(imgB)
- ax0 = plt.subplot(gs[2])
- ax0.set_title('overlay of both images')
- overlayBefore = cv2.addWeighted(imgA/np.max(imgA), 1, imgB/np.max(imgB), 1, 0)
- ax0.imshow(overlayBefore)
- ax0 = plt.subplot(gs[3])
- ax0.set_title('overlay after alignement c = %s \nof BD-AD images' % np.round(cc_max,4), fontsize=10)
- overlayAfter = cv2.addWeighted(imgA/np.max(imgA), 1, imgB_aligned/np.max(imgB_aligned), 1, 0)
- ax0.imshow(overlayAfter)
- #plt.show()
- plt.savefig(figDir + 'ImageAlignment_%s-%s.pdf' % (refDate,otherDate)) # plt.savefig(figOutDir+'ImageAlignment_%s.png' % aS.animalID) # plt.show()
- plt.close()
- return (warp_matrix_max,cc_max)
- #################################################################################
- # calculate correlations between ca-imaging, wheel speed and paw speed
- #################################################################################
- def alignROIsCheckOverlap(statRef,opsRef,statAlign,opsAlign,warp_matrix,refDate,otherDate,figSave=False,figDir=''):
- ncellsRef= len(statRef)
- ncellsAlign = len(statAlign)
- imMaskRef = np.zeros((opsRef['Ly'], opsRef['Lx']))
- imMaskAlign = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
- intersectionROIs = []
- intersectionROIsA = []
- for n in range(0,ncellsRef):
- imMaskRef[:] = 0
- #if iscellBD[n][0]==1:
- #pdb.set_trace()
- ypixRef = statRef[n]['ypix']
- xpixRef = statRef[n]['xpix']
- imMaskRef[ypixRef,xpixRef] = 1
- for m in range(0,ncellsAlign):
- imMaskAlign[:] = 0
- #if iscellAD[m][0]==1:
- ypixAl = statAlign[m]['ypix']
- xpixAl = statAlign[m]['xpix']
- # perform homographic transform : rotation + translation
- #pdb.set_trace()
- points = np.column_stack((xpixAl,ypixAl))
- newPoints = np.copy(points)
- #pdb.set_trace()
- warp_matrix_inverse = np.copy(warp_matrix)
- cv2.invertAffineTransform(warp_matrix,warp_matrix_inverse)
- #newPoints = cv2.transform(points,warp_matrix_inverse)
- xpixAlPrime = np.rint(xpixAl*warp_matrix_inverse[0,0] + ypixAl*warp_matrix_inverse[0,1] + warp_matrix_inverse[0,2])
- 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])
- xpixAlPrime = np.array(xpixAlPrime,dtype=int)
- ypixAlPrime = np.array(ypixAlPrime,dtype=int)
- #pdb.set_trace()
- # make sure pixels remain within
- xpixAlPrime2 = xpixAlPrime[(xpixAlPrime<opsAlign['Lx'])&(ypixAlPrime<opsAlign['Ly'])]
- ypixAlPrime2 = ypixAlPrime[(xpixAlPrime<opsAlign['Lx'])&(ypixAlPrime<opsAlign['Ly'])]
- imMaskAlign[ypixAlPrime2,xpixAlPrime2] = 1
- #imMaskAlign[xpixAlPrime2,ypixAlPrime2] = 1
- intersection = np.sum(np.logical_and(imMaskRef,imMaskAlign))
- eitherOr = np.sum(np.logical_or(imMaskRef,imMaskAlign))
- if intersection>0.2:
- #print(n,m,intersection,eitherOr,intersection/eitherOr)
- intersectionROIs.append([n,m,xpixRef,ypixRef,xpixAlPrime2,ypixAlPrime2,intersection,eitherOr,intersection/eitherOr])
- intersectionROIsA.append([n,m,intersection,eitherOr,intersection/eitherOr])
- # clean up intersection ROIs; each ROI should only overlap once
- def removeDoubleCellOccurrences(interROIs,column):
- uniquePerColumn = np.unique(interROIs[:,column],return_counts=True) # find unique occurrences
- multipleCells = uniquePerColumn[0][uniquePerColumn[1]>1] # which cells occur more than once in the first column
- indiciesToRemove = []
- for i in multipleCells:
- indicies = np.argwhere(interROIs[:,column]==i)
- maxIdx = np.argmax(interROIs[indicies[:,0]][:,4])
- delIndicies = np.delete(indicies,maxIdx)
- indiciesToRemove.extend(delIndicies)
- return indiciesToRemove
- intersectionROIsA = np.asarray(intersectionROIsA)
- removeIdicies0 = removeDoubleCellOccurrences(intersectionROIsA,0)
- removeIdicies1 = removeDoubleCellOccurrences(intersectionROIsA,1)
- removeIndicies = np.asarray(removeIdicies0 + removeIdicies1)
- removeIndicies = np.unique(removeIndicies)
- cleanedIntersectionROIs = []
- for i in range(len(intersectionROIs)):
- if i not in removeIndicies:
- cleanedIntersectionROIs.append(intersectionROIs[i])
- #pdb.set_trace()
- if len(removeIndicies)>0:
- intersectionROIsA = np.delete(intersectionROIsA,removeIndicies,axis=0)
- #pdb.set_trace()
- if figSave:
- imRef = opsRef['meanImg']
- imAlign = opsAlign['meanImg']
- ##################################################################
- # Show final results
- fig = plt.figure(figsize=(15, 15)) ########################
- plt.figtext(0.1, 0.95, '%s and %s' % (refDate,otherDate), clip_on=False, color='black', size=14)
- ax0 = fig.add_subplot(3, 2, 1) #############################
- ax0.set_title('reference img')
- ax0.imshow(imRef)
- ax0 = fig.add_subplot(3, 2, 2) #############################
- ax0.set_title('image to be aligned')
- ax0.imshow(imAlign)
- ax0 = fig.add_subplot(3, 2, 3) #############################
- ax0.set_title('ROIs in reference image')
- imRef = np.zeros((opsRef['Ly'], opsRef['Lx']))
- imRefB = np.zeros((opsRef['Ly'], opsRef['Lx']))
- for n in range(0, ncellsRef):
- ypixR = statRef[n]['ypix']
- xpixR = statRef[n]['xpix']
- imRef[ypixR, xpixR] = n + 1
- imRefB[ypixR, xpixR] = 1
- ax0.imshow(imRef, cmap='gist_ncar')
- ax0 = fig.add_subplot(3, 2, 4) #############################
- ax0.set_title('ROIs in aligned image')
- imAlign = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
- imAlignB = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
- for n in range(0, ncellsAlign):
- ypixA = statAlign[n]['ypix']
- xpixA = statAlign[n]['xpix']
- imAlign[ypixA, xpixA] = n + 1
- imAlignB[ypixA, xpixA] = 2
- ax0.imshow(imAlign, cmap='gist_ncar')
- ax0 = fig.add_subplot(3, 2, 5) #############################
- ax0.set_title('overlapping ROIs Ref-Aligned')
- imRef = np.zeros((opsRef['Ly'], opsRef['Lx']))
- imAlign = np.zeros((opsAlign['Ly'], opsAlign['Lx']))
- for n in range(0, len(cleanedIntersectionROIs)):
- ypixR = cleanedIntersectionROIs[n][3]
- xpixR = cleanedIntersectionROIs[n][2]
- ypixA = cleanedIntersectionROIs[n][5]
- xpixA = cleanedIntersectionROIs[n][4]
- imRef[ypixR, xpixR] = 1
- imAlign[ypixA, xpixA] = 2
- overlayBothROIs1 = cv2.addWeighted(imRef, 1, imAlign, 1, 0)
- #overlayBothROIs1B = cv2.addWeighted(imRefB, 1, imAlignB, 1, 0)
- ax0.imshow(overlayBothROIs1)
- ax0 = fig.add_subplot(3, 2, 6) #############################
- ax0.set_title('fraction of ROI overlap Ref-Aligned')
- interFractions1 = []
- for n in range(0, len(cleanedIntersectionROIs)):
- interFractions1.append(cleanedIntersectionROIs[n][8])
- ax0.hist(interFractions1, bins=15)
- plt.savefig(figDir + 'ROIalignment_%s-%s.pdf' % (refDate, otherDate))
- #plt.show()
- plt.close()
- return (cleanedIntersectionROIs,intersectionROIsA)
- #pickle.dump(intersectionROIs, open( dataOutDir + 'ROIintersections_%s.p' % aS.animalID, 'wb' ) )
- #################################################################################
- # find ROI recorded on ref day and on any other given day
- #################################################################################
- def findMatchingRois(mouse,allCorrDataPerSession,analysisLocation,refDate=0):
- # check for sanity
- nDays = len(allCorrDataPerSession)
- refDay = allCorrDataPerSession[refDate][0]
- print('fluo images will be aligned to recordings of :', refDay)
- refDayCaData = allCorrDataPerSession[refDate][3][0]
- refImg = refDayCaData[2]['meanImgE']
- refImgCutLengths = removeEmptyColumnAndRows(refImg)
- opsRef = refDayCaData[2]
- statRef = refDayCaData[4]
- # create list of recoridng day indicies
- recDaysList = [i for i in range(nDays)]
- movementValuesPreset = np.zeros((len(recDaysList), 2))
- # movementValuesPreset[0] = np.array([-20,-47])
- movementValuesPreset[1] = np.array([143, 153])
- # movementValuesPreset[3] = np.array([-1,15])
- # remove day used for referencing
- recDaysList.remove(refDate)
- if os.path.exists(analysisLocation+'/alignmentData.p'):
- allDataRead = pickle.load(open(analysisLocation+'/alignmentData.p'))
- else:
- allDataRead = None
- allData = []
- for nDay in recDaysList:
- print(allCorrDataPerSession[nDay][0],nDay)
- #imgE = allCorrDataPerSession[nDay][3][0][2]['meanImgE']
- img = allCorrDataPerSession[nDay][3][0][2]['meanImgE']
- cutLengths = removeEmptyColumnAndRows(img)
- if allDataRead is not None:
- warp_matrix = allDataRead[nDay][3]
- else:
- (warp_matrix,cc) = alignTwoImages(refImg,refImgCutLengths,img,cutLengths,allCorrDataPerSession[refDate][0],allCorrDataPerSession[nDay][0],movementValuesPreset[nDay],figShow=True,)
- opsAlign = allCorrDataPerSession[nDay][3][0][2]
- statAlign = allCorrDataPerSession[nDay][3][0][4]
- (cleanedIntersectionROIs,intersectionROIsA) = alignROIsCheckOverlap(statRef,opsRef,statAlign,opsAlign,warp_matrix,allCorrDataPerSession[refDate][0],allCorrDataPerSession[nDay][0],showFig=True)
- print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(statRef), len(statAlign), len(cleanedIntersectionROIs))
- allData.append([allCorrDataPerSession[nDay][0],nDay,cutLengths,warp_matrix,cc,cleanedIntersectionROIs,intersectionROIsA])
- intersectingCellsInRefRecording = np.arange(len(statRef))
- for nDay in recDaysList:
- intersectingCellsInRefRecording = np.intersect1d(intersectingCellsInRefRecording,allData[nDay][5][:,0])
- print(nDay,allCorrDataPerSession[nDay][0],intersectingCellsInRefRecording)
- pdb.set_trace()
- return 0
- #################################################################################
- # correlates mean fluo images recorded all possible recording day combinations
- #################################################################################
- def findOverlayMatchingRoisAllDayCombinations(allCorrDataPerSession, figLocation, allDataRead=None,saveFigure=True):
- nDays = len(allCorrDataPerSession)
- movementValuesPreset = np.zeros((nDays*nDays, 2))
- allOverlayData = {}
- #corrMatrix = np.zeros((nDays,nDays))
- nPair = 0
- for nDayA in range(nDays):
- for nDayB in range(nDays):
- if nDayA != nDayB :
- print(nDayA,nDayB, allCorrDataPerSession[nDayA]['folder'], allCorrDataPerSession[nDayB]['folder'])
- #imgA = allCorrDataPerSession[nDayA][3][0][2]['meanImg']
- imgA = allCorrDataPerSession[nDayA]['caImg']['ops']['meanImg']
- cutLengthsA = removeEmptyColumnAndRows(imgA)
- opsA = allCorrDataPerSession[nDayA]['caImg']['ops'] # allCorrDataPerSession[nDayA][3][0][2]
- statA = allCorrDataPerSession[nDayA]['caImg']['stat'] # allCorrDataPerSession[nDayA][3][0][4]
- imgB = allCorrDataPerSession[nDayB]['caImg']['ops']['meanImg'] # allCorrDataPerSession[nDayB][3][0][2]['meanImg']
- cutLengthsB = removeEmptyColumnAndRows(imgB)
- opsB = allCorrDataPerSession[nDayB]['caImg']['ops'] # allCorrDataPerSession[nDayB][3][0][2]
- statB = allCorrDataPerSession[nDayB]['caImg']['stat'] # allCorrDataPerSession[nDayB][3][0][4]
- if (allDataRead is not None) and (allDataRead[nPair][0]==allCorrDataPerSession[nDayA]['folder']) and (allDataRead[nPair][1]==allCorrDataPerSession[nDayB]['folder']):
- print('warp_matrix for current pair of recordings exists and will be used')
- warp_matrix = allDataRead[nPair][6]
- cc = allDataRead[nPair][7]
- else:
- (warp_matrix,cc) = alignTwoImages(imgA,cutLengthsA,imgB,cutLengthsB,allCorrDataPerSession[nDayA]['folder'],allCorrDataPerSession[nDayB]['folder'],movementValuesPreset[nPair],figSave=saveFigure,figDir=figLocation)
- #corrMatrix[nDayA,nDayB] = cc
- (cleanedIntersectionROIs, intersectionROIsA) = alignROIsCheckOverlap(statA, opsA, statB, opsB, warp_matrix, allCorrDataPerSession[nDayA]['folder'], allCorrDataPerSession[nDayB]['folder'],figSave=saveFigure, figDir=figLocation)
- print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(statA), len(statB), len(cleanedIntersectionROIs))
- #allOverlayData.append([allCorrDataPerSession[nDayA]['folder'], allCorrDataPerSession[nDayB]['folder'], nDayA, nDayB, cutLengthsA, cutLengthsB, warp_matrix, cc,cleanedIntersectionROIs, intersectionROIsA,statA,statB])
- allOverlayData[nPair] = {}
- allOverlayData[nPair]['folderA'] = allCorrDataPerSession[nDayA]['folder'] #0
- allOverlayData[nPair]['folderB'] = allCorrDataPerSession[nDayB]['folder'] #1
- allOverlayData[nPair]['nDayA'] = nDayA #2
- allOverlayData[nPair]['nDayB'] = nDayB #3
- allOverlayData[nPair]['cutLengthsA'] = cutLengthsA #4
- allOverlayData[nPair]['cutLengthsB'] = cutLengthsB #5
- allOverlayData[nPair]['warp_matrix'] = warp_matrix #6
- allOverlayData[nPair]['cc'] = cc #7
- allOverlayData[nPair]['cleanedIntersectionROIs'] = cleanedIntersectionROIs #8
- allOverlayData[nPair]['intersectionROIsA'] = intersectionROIsA #9
- allOverlayData[nPair]['statA'] = statA #10
- allOverlayData[nPair]['statB'] = statB #11
- nPair+=1
- #pdb.set_trace()
- return allOverlayData
- #################################################################################
- # correlates mean fluo images of one day recorded at 910 and 820 nm
- #################################################################################
- def findOverlayMatchingRoisDuringOneDay(allCorrDataPerSession910,allCorrDataPerSession820, figLocation, allDataRead=None,saveFigure=True):
- nDays910 = len(allCorrDataPerSession910)
- nDays820 = len(allCorrDataPerSession820)
- movementValuesPreset = np.zeros((nDays910, 2))
- allAlignData = {}
- #corrMatrix = np.zeros((nDays,nDays))
- maxTimeDelay = 45*60 # maximal 40 min difference btw. 910 and 820 recording
- nPair = 0
- for nDay910 in range(nDays910):
- for nDay820 in range(nDays820):
- if allCorrDataPerSession910[nDay910]['folder'][:-4] == allCorrDataPerSession820[nDay820]['folder'][:-4]:
- #pdb.set_trace()
- timeDiff = np.abs(allCorrDataPerSession910[nDay910]['caImg']['timeStamps'][0,3] - allCorrDataPerSession820[nDay820]['caImg']['timeStamps'][0,3])
- if timeDiff>maxTimeDelay:
- print('PROBLEM in data consistency!')
- print('910 and 820 recordings are separated by %s min' % str(timeDiff/60.))
- pdb.set_trace()
- else:
- print('Delay between 910 and 820 imaging sessions is :', timeDiff/60.,'min')
- print(nDay910,nDay820, allCorrDataPerSession910[nDay910]['folder'], allCorrDataPerSession820[nDay820]['folder'])
- img910 = allCorrDataPerSession910[nDay910]['caImg']['ops']['meanImg'] #allCorrDataPerSession910[nDay910][3][0][2]['meanImg']
- cutLengths910 = removeEmptyColumnAndRows(img910)
- ops910 = allCorrDataPerSession910[nDay910]['caImg']['ops'] # allCorrDataPerSession910[nDay910][3][0][2]
- stat910 = allCorrDataPerSession910[nDay910]['caImg']['stat']
- img820 = allCorrDataPerSession820[nDay820]['caImg']['ops']['meanImg'] #[3][0][2]['meanImg']
- cutLengths820 = removeEmptyColumnAndRows(img820)
- ops820 = allCorrDataPerSession820[nDay820]['caImg']['ops'] #[3][0][2]
- stat820 = allCorrDataPerSession820[nDay820]['caImg']['stat'] #[3][0][4]
- if (allDataRead is not None) and (allDataRead[nPair][0]==allCorrDataPerSession910[nDay910]['folder']) and (allDataRead[nPair][1]==allCorrDataPerSession820[nDay820]['folder']):
- print('warp_matrix for current pair of recordings exists and will be used')
- warp_matrix = allDataRead[nPair][6]
- cc = allDataRead[nPair][7]
- else:
- (warp_matrix,cc) = alignTwoImages(img910,cutLengths910,img820,cutLengths820,allCorrDataPerSession910[nDay910]['folder'],allCorrDataPerSession820[nDay820]['folder'],movementValuesPreset[nPair],figSave=saveFigure,figDir=figLocation)
- #corrMatrix[nDayA,nDayB] = cc
- (cleanedIntersectionROIs, intersectionROIsA) = alignROIsCheckOverlap(stat910, ops910, stat820, ops820, warp_matrix, allCorrDataPerSession910[nDay910]['folder'], allCorrDataPerSession820[nDay820]['folder'],figSave=saveFigure, figDir=figLocation)
- print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(stat910), len(stat820), len(cleanedIntersectionROIs))
- #allAlignData.append([allCorrDataPerSession910[nDay910][0], allCorrDataPerSession820[nDay820][0], nDay910, nDay820, cutLengths910, cutLengths820, warp_matrix, cc,cleanedIntersectionROIs, intersectionROIsA,stat910,stat820])
- allAlignData[nPair] = {}
- allAlignData[nPair]['folder910'] = allCorrDataPerSession910[nDay910]['folder'] #0
- allAlignData[nPair]['folder820'] = allCorrDataPerSession820[nDay820]['folder'] #1
- allAlignData[nPair]['nDay910'] = nDay910 #2
- allAlignData[nPair]['nDay820'] = nDay820 #3
- allAlignData[nPair]['cutLengths910'] = cutLengths910 #4
- allAlignData[nPair]['cutLengths820'] = cutLengths820 #5
- allAlignData[nPair]['warp_matrix'] = warp_matrix #6
- allAlignData[nPair]['cc'] = cc # 7
- allAlignData[nPair]['cleanedIntersectionROIs'] = cleanedIntersectionROIs #8
- allAlignData[nPair]['intersectionROIsA'] = intersectionROIsA #9
- allAlignData[nPair]['statA'] = stat910 #10
- allAlignData[nPair]['statB'] = stat820 #11
- nPair+=1
- return allAlignData
- #################################################################################
- # find ROIs recorded across successive recording days
- #################################################################################
- def findMatchingRoisSuccessivDays(mouse,allCorrDataPerSession,analysisLocation,expDate,figLocation, allDataRead=None):
- # check for sanity
- nDays = len(allCorrDataPerSession)
- # create list of recoridng day indicies
- #recDaysList = [i for i in range(nDays)]
- movementValuesPreset = np.zeros((nDays, 2))
- # movementValuesPreset[0] = np.array([-20,-47])
- # movementValuesPreset[1] = np.array([143, 153])
- # movementValuesPreset[3] = np.array([-1,15])
- allDataStore = []
- for nPair in range(nDays-1):
- nDayA = nPair
- nDayB = nPair + 1
- print(nPair, allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0])
- imgA = allCorrDataPerSession[nDayA][3][0][2]['meanImg']
- cutLengthsA = removeEmptyColumnAndRows(imgA)
- opsA = allCorrDataPerSession[nDayA][3][0][2]
- statA = allCorrDataPerSession[nDayA][3][0][4]
- imgB = allCorrDataPerSession[nDayB][3][0][2]['meanImg']
- cutLengthsB = removeEmptyColumnAndRows(imgB)
- opsB = allCorrDataPerSession[nDayB][3][0][2]
- statB = allCorrDataPerSession[nDayB][3][0][4]
- if (allDataRead is not None) and (allDataRead[nPair][0]==allCorrDataPerSession[nDayA][0]) and (allDataRead[nPair][1]==allCorrDataPerSession[nDayB][0]):
- print('warp_matrix for current pair of recordings exists and will be used')
- warp_matrix = allDataRead[nPair][6]
- cc = allDataRead[nPair][7]
- else:
- (warp_matrix,cc) = alignTwoImages(imgA,cutLengthsA,imgB,cutLengthsB,allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0],movementValuesPreset[nPair],figShow=True,figDir=figLocation)
- (cleanedIntersectionROIs,intersectionROIsA) = alignROIsCheckOverlap(statA,opsA,statB,opsB,warp_matrix,allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0],showFig=True,figDir=figLocation)
- print('Number of ROIs in Ref and aligned images, intersection ROIs :', len(statA), len(statB), len(cleanedIntersectionROIs))
- allDataStore.append([allCorrDataPerSession[nDayA][0],allCorrDataPerSession[nDayB][0],nDayA,nDayB,cutLengthsA,cutLengthsB,warp_matrix,cc,cleanedIntersectionROIs,intersectionROIsA])
- #intersectingCellsInRefRecording = np.arange(len(statA))
- #for nDay in recDaysList:
- # intersectingCellsInRefRecording = np.intersect1d(intersectingCellsInRefRecording,allData[nDay][5][:,0])
- # print(nDay,allCorrDataPerSession[nDay][0],intersectingCellsInRefRecording)
- #pdb.set_trace()
- return allDataStore
- #################################################################################
- # check which ROIs were recoreded across recordings days
- #################################################################################
- def roisRecordedAllDays(allCorrDataPerSession,allAlignData,alignData910And820,correlationThres):
- def getMatchingPairs(aD):
- ROIpairs = []
- for i in range(len(aD)):
- ROIpairs.append(aD[i][:2])
- ROIpairs = np.asarray(ROIpairs)
- return ROIpairs
- #pdb.set_trace()
- #allCombis = list(itertools.combinations((1,2,3,4,5,6,7,8,9),2))
- nDays = len(allCorrDataPerSession)
- correlationThreshold = correlationThres
- intersectData = {}
- nPair = 0
- for nDayA in range(nDays):
- nRef = 0
- corrDays = []
- # 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
- for nDayB in range(nDays):
- if nDayA != nDayB :
- #print(nDayA,nDayB, allCorrDataPerSession[nDayA][0], allCorrDataPerSession[nDayB][0], nRef, nPair)
- if not (nDayA == allAlignData[nPair]['nDayA'] and nDayB == allAlignData[nPair]['nDayB']):
- print(nDayA, nDayB, allAlignData[nPair][2], allAlignData[nPair][3])
- print('sanity check failed! The pairing doesn\'t correspond to the day-pair')
- if nRef==0:
- matchingRoisBefore = getMatchingPairs(allAlignData[nPair]['cleanedIntersectionROIs']) # [8]
- if allAlignData[nPair]['cc']>correlationThreshold:
- corrDays.append([nDayB,allCorrDataPerSession[nDayB]['folder'],nPair]) #[0]
- else:
- matchingRoisAfter = getMatchingPairs(allAlignData[nPair]['cleanedIntersectionROIs']) # [8]
- if allAlignData[nPair]['cc']>correlationThreshold:
- #print(nDayA, nDayB, allCorrDataPerSession[nDayA][0], allCorrDataPerSession[nDayB][0], nRef, nPair)
- idxRemaining = np.intersect1d(matchingRoisBefore[:,0], matchingRoisAfter[:,0]) #
- idxRemainingBefore = [key for key,val in enumerate(matchingRoisBefore[:,0]) if val in idxRemaining]
- idxRemainingAfter = [key for key,val in enumerate(matchingRoisAfter[:,0]) if val in idxRemaining]
- BeforeAlsoAfter = matchingRoisBefore[idxRemainingBefore]
- AfterAlsoBefore = matchingRoisAfter[idxRemainingAfter]
- print('ROIS remaining before and after : ', len(idxRemaining), nDayB, allCorrDataPerSession[nDayB]['folder'])
- matchingRoisBefore = np.copy(BeforeAlsoAfter)
- corrDays.append([nDayB,allCorrDataPerSession[nDayB]['folder'],nPair])
- idxRemainingGood = np.copy(idxRemaining)
- #intersectData.append([nPair,matchingRoisBefore,matchingRoisAfter,idxRemaining,BeforeAlsoAfter,AfterAlsoBefore])
- #print(nPair,matchingRoisBefore,matchingRoisAfter,idxRemaining,BeforeAlsoAfter,AfterAlsoBefore)
- #pdb.set_trace()
- nRef+=1
- nPair+=1
- remainingRoisExists = (True if 'idxRemainingGood' in locals() else False)
- print(nDayA, allCorrDataPerSession[nDayA]['folder'], len(corrDays), (len(idxRemainingGood) if remainingRoisExists else 0), (idxRemainingGood if remainingRoisExists else None), corrDays)
- #intersectData.append([nDayA, allCorrDataPerSession[nDayA]['folder'], len(corrDays), (len(idxRemainingGood) if remainingRoisExists else 0), (idxRemainingGood if remainingRoisExists else None), corrDays])
- intersectData[nDayA] = {}
- intersectData[nDayA]['nDayA'] = nDayA # 0
- intersectData[nDayA]['folder'] = allCorrDataPerSession[nDayA]['folder'] # 1
- intersectData[nDayA]['lenCorrDays'] = len(corrDays) # 2
- intersectData[nDayA]['lenIdxRemainingGood'] = (len(idxRemainingGood) if remainingRoisExists else 0) # 3
- intersectData[nDayA]['idxRemainingGood'] = (idxRemainingGood if remainingRoisExists else None) # 4
- intersectData[nDayA]['corrDays'] = corrDays #5
- # second loop to determine identity of the remaining ROIs in all recordings
- #idxDay = intersectData[idxRef][5][r][0]
- #idxRoi = intersectData[idxRef][5][r][3 + n][1]
- #dayID = intersectData[idxRef][5][r][1]
- #nPairB = 0
- for i in range(len(intersectData)):
- print(i, intersectData[i]['folder'], intersectData[i]['lenCorrDays'],intersectData[i]['lenIdxRemainingGood'])
- idxRemainingGood = intersectData[i]['idxRemainingGood']
- for nDay in range(intersectData[i]['lenCorrDays']):
- #print(allAlignData[intersectData[i][5][2]][2])
- matchingRois = getMatchingPairs(allAlignData[intersectData[i]['corrDays'][nDay][2]]['cleanedIntersectionROIs'])
- #print('match :',matchingRois,idxRemainingGood)
- idxRemaining = [key for key,val in enumerate(matchingRois[:,0]) if val in idxRemainingGood]
- RemainingRois = matchingRois[idxRemaining]
- intersectData[i]['corrDays'][nDay].append(RemainingRois)
- #roiData.append([intersectData[i][0],intersectData[i][1],])
- #pdb.set_trace()
- # another loop to append 820 identities of remaining ROIs
- for i in range(len(intersectData)):
- #print(i, intersectData[i][1], intersectData[i][2], intersectData[i][3])
- idxRemainingGood = intersectData[i]['idxRemainingGood']
- for n in range(len(alignData910And820)):
- if (intersectData[i]['folder'] == alignData910And820[n]['folder820']) and (idxRemainingGood is not None):
- print(i,n,intersectData[i]['folder'],alignData910And820[n]['folder820'])
- matchingRois = getMatchingPairs(alignData910And820[n]['cleanedIntersectionROIs'])
- idxRemaining = [key for key,val in enumerate(matchingRois[:,0]) if val in idxRemainingGood]
- RemainingRois = matchingRois[idxRemaining]
- #intersectData[i].append(RemainingRois)
- intersectData[i]['RemainingRois'] = RemainingRois
- return intersectData
- #################################################################################
- # calculate correlations between ca-imaging, wheel speed and paw speed
- #################################################################################
- def doCorrelationAnalysisLocomotionPeriod(allCorrDataPerSession,allStepData):
- #
- matplotlib.use('TkAgg')
- xPixToUm = 0.79
- yPixToUm = 0.8
- motorizationCorrelations = []
- varExplainedMotorization = []
- maxMotorization = [12, 53] # in sec
- for nDay in range(len(allCorrDataPerSession)):
- print(nDay,allCorrDataPerSession[nDay]['folder'])
- (wheelSpeedDictInterP, pawTracksDictInterP, caTracesDictInterP, wheelSpeedDict, pawTracksDict, caTracesDict, slowestTrial) = getCaWheelPawInterpolatedDictsPerDay(nDay, allCorrDataPerSession,allStepData)
- # ATTENTION : all of the arrays also contain a time array
- # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTimeSame]
- # dims of pawTracksDictInterP : [nSessions][nPaw][2][valuesOverTimeSame]
- # dims of caTracesDictInterP : [nSessions][nRois+1][valuesOverTimeSame]
- # dims of wheelSpeedDictInterP : [nSessions][2][valuesOverTime]
- # dims of pawTracksDict : [nSessions][nPaw][2][valuesOverTime]
- # dims of caTracesDict : [nSessions][nRois+1][valuesOverTime]
- # correlations between calcium traces ######################################################################
- stat = allCorrDataPerSession[nDay]['caImg']['stat']#[3][0][4]
- nTrials = len(caTracesDict)
- nRois = np.shape(caTracesDict[0])[0] - 1
- allCoords = []
- for i in range(nRois):
- allCoords.append([i, stat[i]['med'][0], stat[i]['med'][1]]) # first extract the coordinates of all ROIs
- allCoordsSorted = sorted(allCoords, key=lambda x: (x[1], x[2])) # list according to increasing x and y coordinates
- allCoordsSorted = np.asarray(allCoordsSorted)
- combis = list(itertools.combinations(np.array(allCoordsSorted[:, 0], dtype=int), 2)) # use the sorted list to create the combinations
- corrCaTraces = np.zeros((len(combis), nTrials, 9))
- # for moving average windows
- tWindow = 1. # removing changes on the order of 1 s and longer
- ttime = caTracesDict[0][0]
- dt = np.mean(np.diff(ttime))
- Nwindow = int(tWindow / dt + 0.5)
- for i in range(len(combis)):
- # get location information of both ROIs
- xy0 = stat[combis[i][0]]['med']
- xy1 = stat[combis[i][1]]['med']
- # calculate eucleadian distance and x, y distance between cells
- euclDist = np.sqrt(((xy0[1]-xy1[1])*xPixToUm)**2 + ((xy0[0]-xy1[0])*yPixToUm)**2)
- 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
- #pdb.set_trace()
- for t in range(nTrials):
- mask = (caTracesDict[t][0] >= maxMotorization[0]) & (caTracesDict[t][0] < maxMotorization[1])
- corrTemp = scipy.stats.pearsonr(caTracesDict[t][combis[i][0]+1][mask],caTracesDict[t][combis[i][1]+1][mask])
- caTemp0 = np.convolve(caTracesDict[t][combis[i][0]+1][mask], np.ones((Nwindow,))/Nwindow, mode='same')
- caTemp1 = np.convolve(caTracesDict[t][combis[i][1]+1][mask], np.ones((Nwindow,))/Nwindow, mode='same')
- shortCorrTemp = scipy.stats.pearsonr(caTracesDict[t][combis[i][0]+1][mask]-caTemp0,caTracesDict[t][combis[i][1]+1][mask]-caTemp1)
- corrCaTraces[i,t] = np.array([combis[i][0],combis[i][1],corrTemp[0],corrTemp[1],euclDist,xyDist[0],xyDist[1],shortCorrTemp[0],shortCorrTemp[1]])
- # activity measures of ca traces before vs during walking ######################################################################
- tBaseline = 5
- tMotorization = [12,53]
- activityCaTraces = np.zeros((nRois,nTrials,4))
- for n in range(nRois):
- for t in range(nTrials):
- baselineMask = caTracesDict[t][0]<tBaseline
- activityMask = (caTracesDict[t][0]>=tMotorization[0]) & (caTracesDict[t][0]<tMotorization[1])
- (baseLMean,baseLSTD) = (np.mean(caTracesDict[t][n+1][baselineMask]),np.std(caTracesDict[t][n+1][baselineMask]))
- (actMean, actSTD) = (np.mean(caTracesDict[t][n + 1][activityMask]), np.std(caTracesDict[t][n + 1][activityMask]))
- activityCaTraces[n,t] = np.array([baseLMean,baseLSTD,actMean, actSTD])
- ###################################################################
- # correlation between calcium and paw as well as paw speed
- corrCaWheel = np.zeros((nRois,nTrials, 3))
- corrCaPawTraces = np.zeros((nRois,nTrials, 9))
- for i in range(nRois):
- for t in range(nTrials):
- caMask = (caTracesDictInterP[t][0]>= maxMotorization[0]) & (caTracesDictInterP[t][0] < maxMotorization[1])
- wheelMask = (wheelSpeedDictInterP[t][0]>= maxMotorization[0]) & (wheelSpeedDictInterP[t][0] < maxMotorization[1])
- paw0Mask = (pawTracksDictInterP[t][0][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][0][0] < maxMotorization[1])
- paw1Mask = (pawTracksDictInterP[t][1][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][1][0] < maxMotorization[1])
- paw2Mask = (pawTracksDictInterP[t][2][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][2][0] < maxMotorization[1])
- paw3Mask = (pawTracksDictInterP[t][3][0]>=maxMotorization[0]) & (pawTracksDictInterP[t][3][0] < maxMotorization[1])
- corrWheelTemp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], wheelSpeedDictInterP[t][1][wheelMask])
- corrCaWheel[i][t] = np.array([i,corrWheelTemp[0],corrWheelTemp[1]])
- corrPaw0Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][0][1][paw0Mask])
- corrPaw1Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][1][1][paw1Mask])
- corrPaw2Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][2][1][paw2Mask])
- corrPaw3Temp = scipy.stats.pearsonr(caTracesDictInterP[t][i+1][caMask], pawTracksDictInterP[t][3][1][paw3Mask])
- corrCaPawTraces[i][t] = np.array([i,corrPaw0Temp[0],corrPaw0Temp[1],corrPaw1Temp[0],corrPaw1Temp[1],corrPaw2Temp[0],corrPaw2Temp[1],corrPaw3Temp[0],corrPaw3Temp[1]])
- ###################################################################
- # correlation btw. PCA components and wheel, paw speeds
- # concatenate calcium, wheel speed and paw speed for PCA
- pawAll = {}
- pawMask = {}
- for t in range(nTrials):
- caMask = (caTracesDictInterP[t][0] >= maxMotorization[0]) & (caTracesDictInterP[t][0] < maxMotorization[1])
- wheelMask = (wheelSpeedDictInterP[t][0] >= maxMotorization[0]) & (wheelSpeedDictInterP[t][0] < maxMotorization[1])
- for i in range(4):
- pawMask[i] = (pawTracksDictInterP[t][i][0] >= maxMotorization[0]) & (pawTracksDictInterP[t][i][0] < maxMotorization[1])
- if t == 0:
- caAll = caTracesDictInterP[t][:,caMask]
- wheelAll = wheelSpeedDictInterP[t][:,wheelMask]
- for i in range(4):
- pawAll[i] = pawTracksDictInterP[t][i][:,pawMask[i]]
- else:
- caAll = np.column_stack((caAll,caTracesDictInterP[t][:,caMask]))
- wheelAll = np.column_stack((wheelAll,wheelSpeedDictInterP[t][:,wheelMask]))
- for i in range(4):
- pawAll[i] = np.column_stack((pawAll[i],pawTracksDictInterP[t][i][:,pawMask[i]]))
- #pdb.set_trace()
- pcaComponents = 5
- print('doing PCA ...')
- X = np.transpose(caAll[1:])
- pca = PCA(n_components=pcaComponents)
- pca.fit(X)
- X_pca = pca.transform(X)
- #print(pca.components_)
- pcaCorrs = np.zeros((pcaComponents,11))
- varExplainedMotorization.append(pca.explained_variance_ratio_)
- print(pca.explained_variance_ratio_)
- for i in range(pcaComponents):
- corrWheelTemp = scipy.stats.pearsonr(X_pca[:,i], wheelAll[1])
- corrPaw0Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[0][1])
- corrPaw1Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[1][1])
- corrPaw2Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[2][1])
- corrPaw3Temp = scipy.stats.pearsonr(X_pca[:,i], pawAll[3][1])
- pcaCorrs[i] = ([i,corrWheelTemp[0],corrWheelTemp[1],corrPaw0Temp[0],corrPaw0Temp[1],corrPaw1Temp[0],corrPaw1Temp[1],corrPaw2Temp[0],corrPaw2Temp[1],corrPaw3Temp[0],corrPaw3Temp[1]])
- ###################################################################
- motorizationCorrelations.append([nTrials,corrCaTraces,corrCaWheel,corrCaPawTraces,pcaCorrs,activityCaTraces])
- return (motorizationCorrelations,varExplainedMotorization)
- #################################################################################
- # chose day of reference from FOV alignment data
- #################################################################################
- def getRefIdxPerMouse(mouse):
- refIdxMouseDict = {'210214_m12':7,
- '210214_m13':4,
- '210214_m14':5,
- '210214_m15':6,
- '210214_m17':6,
- '210214_m18':9,
- '210214_m19':5,
- '210214_m20':3,
- '210122_f83':2,
- '210122_f84':6,
- '210120_m85':1,
- '210120_m86':0,
- }
- idxRef = refIdxMouseDict[mouse]
- return idxRef
- #################################################################################
- # chose day of reference from FOV alignment data
- ##########################################################
dataAnalysis.py at commit 7217d46, under gpl · at the source
Overview
- Université Paris Cité, CNRS, Saints-Pères Paris Institute for the Neurosciences, Paris, France
- Laboratoire de Neurosciences Cognitives et Computationnelles, INSERM U960, Département d’Études Cognitives, École Normale Supérieure, PSL University, Paris, France
- Institut des Systèmes Intelligents et de Robotique, Sorbonne Université, CNRS, Paris, France
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
7217d46bf438e44d9d4f735d109cab8e4d9f6d2a, 25 April 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
55 files
- analyzeEphysPsortData.py
, Python, 90 lines - analyzePsortTrainCluster
ing.py , Python, 109 lines - extractBehaviorAndWhiske
rCameraTiming.py , Python, 262 lines - extractPawTrackingOutlie
rs.py , Python, 105 lines - extractRungLocation.py, Python, 98 lines
- extractSwingStancePhases
.py , Python, 185 lines, 2 matches - getPsortWalkingActivityA
ndPawTraces.py , Python, 125 lines, 1 match - getPsortWalkingActivityA
ndPawTracesCalculatePSTH , Python, 235 lines, 1 match.py - getRawBehaviorImagesSave
Video.py , Python, 82 lines, 1 match - getWalkingActivity.py, Python, 71 lines
- groupAnalysisScripts/
ComplexSpikesGroupAnalys , Python, 98 linesis.py - groupAnalysisScripts/
CrossCorrelationAcrossAn , Python, 244 linesalysis_MG.py - groupAnalysisScripts/
CrossCorrelationAnalysis , Python, 246 lines_MG.py - groupAnalysisScripts/
PSTHGroupAnalysis.py , Python, 129 lines - groupAnalysisScripts/
PSTHGroupAnalysis_MG.py , Python, 86 lines - groupAnalysisScripts/
PSTHGroupAnalysis_contro , Python, 77 linesls.py - groupAnalysisScripts/
PSTHGroupAnalysis_missSt , Python, 89 lineseps.py - groupAnalysisScripts/
PSTHGroupAnalysis_swingL , Python, 81 linesength_range.py - groupAnalysisScripts/
clusterMLIkernelprofiles , Python, 203 lines.py - groupAnalysisScripts/
clusterPCkernelprofiles. , Python, 203 lines, 1 matchpy - groupAnalysisScripts/
generateAllMouseDictiona , Python, 93 linesry.py - groupAnalysisScripts/
generatePSTHGroupAnalysi , Python, 72 linessFigures.py - manuscriptFigures/
fig_ephys-clustering.py , Python, 35 lines - manuscriptFigures/
fig_ephys-recordings_MLI , Python, 82 liness.py - manuscriptFigures/
fig_experiment-PawPositi , Python, 85 lines, 1 matchonExtraction.py - manuscriptFigures/
fig_experiment-PawPositi , Python, 53 linesonExtraction_assembly.py - manuscriptFigures/
fig_locomotorLearning.py , Python, 38 lines - manuscriptFigures/
fig_locomotorLearning_su , Python, 42 linespplements.py - manuscriptFigures/
fig_miss-step-analysis.p , Python, 93 linesy - manuscriptFigures/
fig_psth_group_analysis_ , Python, 75 linesPCcell_based.py - manuscriptFigures/
fig_psth_group_analysis_ , Python, 69 linescell_based.py - manuscriptFigures/
fig_psth_group_analysis_ , Python, 69 linescell_based_supp.py - manuscriptFigures/
fig_psth_group_analysis_ , Python, 67 linesrecording_based.py - manuscriptFigures/
fig_psth_group_analysis_ , Python, 74 linesrecording_based_PC.py - manuscriptFigures/
fig_real_time_experiment , Python, 185 lines.py - manuscriptFigures/
fig_swing-phase-analysis , Python, 71 lines.py - runAnalysisOnMultipleMic
e.py , Python, 56 lines - tools/
createGroupVisualization , Python, 3,323 lines, 2 matchess.py - tools/
createPublicationVisuali , Python, 3,212 lines, 4 matcheszations.py - tools/
dataAnalysis.py , Python, 3,595 lines, 7 matches - tools/
dataAnalysis_cellCluster , Python, 243 lines, 1 matching.py - tools/
dataAnalysis_psth.py , Python, 1,447 lines - tools/
extractRungDistance.py , Python, 78 lines - tools/
extractRungDistanceBehav , Python, 154 linesior.py - tools/
extractSaveData.py , Python, 3,075 lines - tools/
generatePsortFiles.py , Python, 46 lines - tools/
groupAnalysis.py , Python, 3,057 lines, 2 matches - tools/
groupAnalysis_psth.py , Python, 1,364 lines - tools/
h5pyTools.py , Python, 87 lines - tools/
openCVImageProcessingToo , Python, 2,616 linesls.py - tools/
packageZenodoData.py , Python, 343 lines - tools/
parameters.py , Python, 7 lines - tools/
prepareForPSort.py , Python, 31 lines - tools/
zenodo_package.py , Python, 571 lines - README.md, Text, 149 lines
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:
- it points to the authors' code: mgraupe/
LocoReach-analysis-2026
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
- zenodo:19708054, at Zenodo; found in “Data availability”
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:
- it points to a dataset: Zenodo 19708054
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://
BibTeX
@article{andrianarivelo2
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/
url = {https://
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/
VL - 17
IS - 1
SP - 8861
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"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":
"volume": "17",
"issue": "1",
"page": "8861",
"DOI": "10.1038/
"PMID": "42476976",
"PMCID": "PMC13500626",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://
"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 controlJournal: eLifeIn 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: iScienceIn 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: eLifeIn 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 neuroscienceIn 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: iScienceIn 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 neuroscienceIn 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. MedicineIn 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: NatureIn 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 neuroscienceIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 54 scripts, and 23 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:7da43987bd7d6e06…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
