OSCR

Single-neuron network topology governs neural computation and learning in primate cortex.

Code ↔ Paper

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

The 2 matches
  1. [1] § Methods › Quantify the influence of single neuron on the activity dynamics within network › Demixed principal component analysis (dPCA) ↔ python/dPCA/dPCA.py, lines 21–93 · score 0.69 · demixed principal component, population activity, dPCA, dimensionality, variance
  2. [2] § Methods › Quantify the evolution of population activity › Alignment index ↔ python/dPCA/dPCA.py, lines 21–93 · score 0.63 · population activity, variance explained, principal component, PCA, axis, dimensional

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 · 984 lines · 40 KB · MIT · 2 matches

  1. """ demixed Principal Component Analysis
  2. """
  3. # Author: Wieland Brendel <[email hidden]>
  4. #
  5. # License: BSD 3 clause
  6. from __future__ import print_function
  7. import numpy as np
  8. from collections import OrderedDict
  9. from itertools import combinations, chain
  10. from scipy.sparse.linalg import svds
  11. from scipy.linalg import pinv
  12. from sklearn.base import BaseEstimator
  13. from sklearn.utils.extmath import randomized_svd
  14. import numexpr as ne
  15. from .utils import shuffle2D, classification, denoise_mask
  16. class dPCA(BaseEstimator):
  17. """ demixed Principal component analysis (dPCA)
  18. dPCA is a linear dimensionality reduction technique that automatically discovers
  19. and highlights the essential features of complex population activities. The
  20. population activity is decomposed into a few demixed components that capture most
  21. of the variance in the data and that highlight the dynamic tuning of the population
  22. to various task parameters, such as stimuli, decisions, rewards, etc.
  23. Parameters
  24. ----------
  25. labels : int or string
  26. Labels of feature axis.
  27. If int the corresponding number of labels are selected from the alphabet 'abcde...'
  28. join : None or dict
  29. Parameter combinations to join
  30. If a data set has parametrized by time t and stimulus s, then dPCA will split
  31. the data into marginalizations corresponding to 't', 's' and 'ts'. At times,
  32. we want to join different marginalizations (like 's' and 'ts'), e.g. if
  33. we are only interested in the time-modulated stimulus components. In this case,
  34. we would pass {'ts' : ['s','ts']}.
  35. regularizer : None, float, 'auto'
  36. Regularization parameter. If None or 0, then no regularization is applied.
  37. For float, the regularization weight is regularizer*var(data). If 'auto', the
  38. optimal regularization parameter is found during fitting (might take some time).
  39. n_components : None, int or dict
  40. Number of components to keep.
  41. If n_components is int, then the same number of components are kept in every
  42. marginalization. Otherwise, the dict allows to set the number of components
  43. in each marginalization (e.g. {'t' : 10, 'ts' : 5}). Defaults to 10.
  44. copy : bool
  45. If False, data passed to fit are overwritten and running
  46. fit(X).transform(X) will not yield the expected results,
  47. use fit_transform(X) instead.
  48. n_iter : int (default: 0)
  49. Number of iterations for randomized SVD solver (sklearn).
  50. Attributes
  51. ----------
  52. explained_variance_ratio_ : dict with arrays, [n_components]
  53. Dictionary in which each key refers to one marginalization and the \
  54. value is a vector with the percentage of variance explained by each of \
  55. the marginal components.
  56. Notes
  57. -----
  58. Implements the dPCA model from:
  59. D Kobak*, W Brendel*, C Constantinidis, C Feierstein, A Kepecs, Z Mainen, \
  60. R Romo, X-L Qi, N Uchida, C Machens
  61. Demixed principal component analysis of population activity in higher \
  62. cortical areas reveals independent representation of task parameters,
  63. Examples
  64. --------
  65. >>> import numpy as np
  66. >>> from dPCA import dPCA
  67. >>> X = np.array([[-1, -1], [-2, -1], [-3, -2], [1, 1], [2, 1], [3, 2]])
  68. >>> dpca = dPCA(n_components=2)
  69. >>> dpca.fit(X)
  70. PCA(copy=True, n_components=2, whiten=False)
  71. >>> print(pca.explained_variance_ratio_)
  72. [ 0.99244... 0.00755...]
  73. """
  74. def __init__(self, labels=None, join=None, n_components=10, regularizer=None, copy=True, n_iter=0):
  75. # create labels from alphabet if not provided
  76. if isinstance(labels,str):
  77. self.labels = labels
  78. elif isinstance(labels,int):
  79. alphabet = 'abcdefghijklmnopqrstuvwxyz'
  80. self.labels = alphabet[:labels]
  81. else:
  82. raise TypeError('Wrong type for labels. Please either set labels to the number of variables or provide the axis labels as a single string of characters (like "ts" for time and stimulus)')
  83. self._join = join
  84. self.join = join
  85. self.regularizer = 0 if regularizer == None else regularizer
  86. self.opt_regularizer_flag = regularizer == 'auto'
  87. self.n_components = n_components
  88. self.copy = copy
  89. self.marginalizations = self._get_parameter_combinations()
  90. self.n_iter = n_iter
  91. # set debug mode, 0 = no reports, 1 = warnings, 2 = warnings & progress, >2 = everything
  92. self.debug = 2
  93. if regularizer == 'auto':
  94. print("""You chose to determine the regularization parameter automatically. This can
  95. take substantial time and grows linearly with the number of crossvalidation
  96. folds. The latter can be set by changing self.n_trials (default = 3). Similarly,
  97. use self.protect to set the list of axes that are not supposed to get to get shuffled
  98. (e.g. upon splitting the data into test- and training, time-points should always
  99. be drawn from the same trial, i.e. self.protect = ['t']). This can significantly
  100. speed up the code.""")
  101. self.n_trials = 3
  102. self.protect = None
  103. def fit(self, X, trialX=None):
  104. """Fit the model with X.
  105. Parameters
  106. ----------
  107. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  108. Training data, where n_samples in the number of samples
  109. and n_features_j is the number of the j-features (where the axis correspond
  110. to different parameters).
  111. Returns
  112. -------
  113. self : object
  114. Returns the instance itself.
  115. """
  116. self._fit(X,trialX=trialX)
  117. return self
  118. def fit_transform(self, X, trialX=None):
  119. """Fit the model with X and apply the dimensionality reduction on X.
  120. Parameters
  121. ----------
  122. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  123. Training data, where n_samples in the number of samples
  124. and n_features_j is the number of the j-features (where the axis correspond
  125. to different parameters).
  126. Returns
  127. -------
  128. X_new : dict with arrays with the same shape as X
  129. Dictionary in which each key refers to one marginalization and the value is the
  130. latent component.
  131. """
  132. self._fit(X,trialX=trialX)
  133. return self.transform(X)
  134. def _get_parameter_combinations(self,join=True):
  135. ''' Returns all parameter combinations, e.g. for labels = 'xyz'
  136. {'x' : (0,), 'y' : (1,), 'z' : (2,), 'xy' : (0,1), 'xz' : (0,2), 'yz' : (1,2), 'xyz' : (0,1,2)}
  137. If join == True, parameter combinations are condensed according to self._join, Otherwise all
  138. combinations are returned.
  139. '''
  140. # subsets = () (0,) (1,) (2,) (0,1) (0,2) (1,2) (0,1,2)"
  141. subsets = list(chain.from_iterable(combinations(list(range(len(self.labels))), r) for r in range(len(self.labels))))
  142. # delete empty set & add (0,1,2)
  143. del subsets[0]
  144. subsets.append(list(range(len(self.labels))))
  145. # create dictionary
  146. pcombs = OrderedDict()
  147. for subset in subsets:
  148. key = ''.join([self.labels[i] for i in subset])
  149. pcombs[key] = set(subset)
  150. # condense dict if not None
  151. if isinstance(self._join,dict) and join:
  152. for key, combs in self._join.items():
  153. tmp = [pcombs[comb] for comb in combs]
  154. for comb in combs:
  155. del pcombs[comb]
  156. pcombs[key] = tmp
  157. return pcombs
  158. def _marginalize(self,X,save_memory=False):
  159. """ Marginalize the data matrix
  160. Parameters
  161. ----------
  162. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  163. Training data, where n_samples in the number of samples
  164. and n_features_j is the number of the j-features (where the axis correspond
  165. to different parameters).
  166. save_memory : bool, set to True if memory really is an issue (though optimization is not perfect yet)
  167. Returns
  168. -------
  169. mXs : dictionary, with values corresponding to the marginalized data (and the key refers to the marginalization)
  170. """
  171. def mmean(X,axes,expand=False):
  172. ''' Takes mean along several axis (given as list). If expand the averaged dimensions will be filled with
  173. new axis to retain the dimension.
  174. '''
  175. Z = X.copy()
  176. for ax in np.sort(axes)[::-1]:
  177. Z = np.mean(Z,ax)
  178. if expand == True:
  179. Z = np.expand_dims(Z,ax)
  180. return Z
  181. def dense_marg(Y,mYs):
  182. ''' The original marginalizations as returned by "get_marginalizations" are sparse in the sense that
  183. marginalized axis are newaxis. This functions blows them up to the original size of the data set
  184. (need for optimization).
  185. '''
  186. tmp = np.zeros_like(Y)
  187. for key in list(mYs.keys()):
  188. mYs[key] = (tmp + mYs[key]).reshape((Y.shape[0],-1))
  189. return mYs
  190. Xres = X.copy() # residual of data
  191. # center data
  192. Xres -= np.mean(Xres.reshape((Xres.shape[0],-1)),-1).reshape((Xres.shape[0],) + (len(Xres.shape)-1)*(1,))
  193. # init dict with marginals
  194. Xmargs = OrderedDict()
  195. # get parameter combinations
  196. pcombs = self._get_parameter_combinations(join=False)
  197. # subtract the mean
  198. S = list(pcombs.values())[-1] # full set of indices
  199. if save_memory:
  200. for key, phi in pcombs.items():
  201. S_without_phi = list(S - phi)
  202. # compute marginalization and save
  203. Xmargs[key] = mmean(Xres,np.array(S_without_phi)+1,expand=True)
  204. # subtract the marginalization from the data
  205. Xres -= Xmargs[key]
  206. else:
  207. # efficient precomputation of means
  208. pre_mean = {}
  209. for key, phi in pcombs.items():
  210. if len(key) == 1:
  211. pre_mean[key] = mmean(Xres,np.array(list(phi))+1,expand=True)
  212. else:
  213. pre_mean[key] = mmean(pre_mean[key[:-1]],np.array([list(phi)[-1]])+1,expand=True)
  214. # compute marginalizations
  215. for key, phi in pcombs.items():
  216. key_without_phi = ''.join(filter(lambda ch: ch not in key, self.labels))
  217. # self.labels.translate(None, key)
  218. # build local dictionary for numexpr
  219. X = pre_mean[key_without_phi] if len(key_without_phi) > 0 else Xres
  220. if len(key) > 1:
  221. subsets = list(chain.from_iterable(combinations(key, r) for r in range(1,len(key))))
  222. subsets = [''.join(subset) for subset in subsets]
  223. local_dict = {subset : Xmargs[subset] for subset in subsets}
  224. local_dict['X'] = X
  225. Xmargs[key] = ne.evaluate('X - ' + ' - '.join(subsets),local_dict=local_dict)
  226. else:
  227. Xmargs[key] = X
  228. # condense dict if not None
  229. if isinstance(self._join,dict):
  230. for key, combs in self._join.items():
  231. Xshape = np.ones(len(self.labels)+1,dtype='int')
  232. for comb in combs:
  233. sh = np.array(Xmargs[comb].shape)
  234. Xshape[(sh-1).nonzero()] = sh[(sh-1).nonzero()]
  235. tmp = np.zeros(Xshape)
  236. for comb in combs:
  237. tmp += Xmargs[comb]
  238. del Xmargs[comb]
  239. Xmargs[key] = tmp
  240. Xmargs = dense_marg(X,Xmargs)
  241. return Xmargs
  242. def _optimize_regularization(self,X,trialX,center=True,lams='auto'):
  243. """ Optimization routine to find optimal regularization parameter.
  244. TO DO: Routine is pretty dumb right now (go through predetermined
  245. list and find minimum). There are several ways to speed it up.
  246. """
  247. # center data
  248. if center:
  249. X = X - np.mean(X.reshape((X.shape[0],-1)),1).reshape((X.shape[0],)\
  250. + len(self.labels)*(1,))
  251. # compute variance of data
  252. varX = np.sum(X**2)
  253. # test different inits and regularization parameters
  254. if lams == 'auto':
  255. N = 45
  256. lams = np.logspace(0,N,num=N, base=1.4, endpoint=False)*1e-7
  257. # compute crossvalidated score over n_trials repetitions
  258. scores = self.crossval_score(lams,X,trialX,mean=False)
  259. # take mean over total scores
  260. totalscore = np.mean(np.sum(np.dstack([scores[key] for key in list(scores.keys())]),-1),0)
  261. # Raise warning if optimal lambda lies at boundaries
  262. if np.argmin(totalscore) == 0 or np.argmin(totalscore) == len(totalscore) - 1:
  263. if self.debug > 0:
  264. print("Warning: Optimal regularization parameter lies at the \
  265. boundary of the search interval. Please provide \
  266. different search list (key: lams).")
  267. # set minimum as new lambda
  268. self.regularizer = lams[np.argmin(totalscore)]
  269. if self.debug > 1:
  270. print('Optimized regularization, optimal lambda = ', self.regularizer)
  271. print('Regularization will be fixed; to compute the optimal \
  272. parameter again on the next fit, please \
  273. set opt_regularizer_flag to True.')
  274. self.opt_regularizer_flag = False
  275. def crossval_score(self,lams,X,trialX,mean=True):
  276. """ Calculates crossvalidation scores for a given set of regularization
  277. parameters. To this end it takes one parameter off the list,
  278. computes the model on a training set and then validates the
  279. reconstruction performance on a validation set.
  280. Parameters
  281. ----------
  282. lams: 1D array of floats
  283. Array of regularization parameters to test.
  284. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  285. Training data, where n_samples in the number of samples
  286. and n_features_j is the number of the j-features (where the
  287. axis correspond to different parameters).
  288. trialX: array-like, shape (n_trials, n_samples, n_features_1, n_features_2, ...)
  289. Trial-by-trial data. Shape is similar to X but with an additional axis at the beginning
  290. with different trials. If different combinations of features have different number
  291. of trials, then set n_samples to the maximum number of trials and fill unoccupied data
  292. points with NaN.
  293. mean: bool (default: True)
  294. Set True if the crossvalidation score should be averaged over
  295. all marginalizations, otherwise False.
  296. Returns
  297. -------
  298. mXs : dictionary, with values corresponding to the marginalized
  299. data (and the key refers to the marginalization)
  300. """
  301. # placeholder for scores
  302. scores = np.zeros((self.n_trials,len(lams))) if mean else {key : np.zeros((self.n_trials,len(lams))) for key in list(self.marginalizations.keys())}
  303. # compute number of samples in each condition
  304. N_samples = self._get_n_samples(trialX,protect=self.protect)
  305. for trial in range(self.n_trials):
  306. print("Starting trial ", trial + 1, "/", self.n_trials)
  307. # perform split into training and test trials
  308. trainX, validX = self.train_test_split(X,trialX,N_samples=N_samples)
  309. # compute marginalization of test and validation data
  310. trainmXs, validmXs = self._marginalize(trainX), self._marginalize(validX)
  311. # compute crossvalidation score for every regularization parameter
  312. for k, lam in enumerate(lams):
  313. # fit dpca model
  314. self.regularizer = lam
  315. self._fit(trainX,mXs=trainmXs,optimize=False)
  316. # compute crossvalidation score
  317. if mean:
  318. scores[trial,k] = self._score(validX,validmXs)
  319. else:
  320. tmp = self._score(validX,validmXs,mean=False)
  321. for key in list(self.marginalizations.keys()):
  322. scores[key][trial,k] = tmp[key]
  323. return scores
  324. def _score(self,X,mXs,mean=True):
  325. """ Scoring for crossvalidation. Predicts one observable (e.g. one neuron) of X at a time, using all other dimensions:
  326. \sum_phi ||X[n] - F_\phi D_phi^{-n} X^{-n}||^2
  327. where phi refers to the marginalization and X^{-n} (D_phi^{-n}) are all rows of X (D) except the n-th row.
  328. """
  329. n_features = X.shape[0]
  330. X = X.reshape((n_features,-1))
  331. error = {key: 0 for key in list(mXs.keys())}
  332. PDY = {key : np.dot(self.P[key],np.dot(self.D[key].T,X)) for key in list(mXs.keys())}
  333. trPD = {key : np.sum(self.P[key]*self.D[key],1) for key in list(mXs.keys())}
  334. for key in list(mXs.keys()):
  335. error[key] = np.sum((mXs[key] - PDY[key] + trPD[key][:,None]*X)**2)
  336. return error if not mean else np.sum(list(error.values()))
  337. def _randomized_dpca(self,X,mXs,pinvX=None):
  338. """ Solves the dPCA minimization problem analytically by using a randomized SVD solver from sklearn.
  339. Returns
  340. -------
  341. P : dict mapping strings to array-like,
  342. Holds encoding matrices for each term in variance decompostions (used in inverse_transform
  343. to map from low-dimensional representation back to original data space).
  344. D : dict mapping strings to array-like,
  345. Holds decoding matrices for each term in variance decompostions (used to transform data
  346. to low-dimensional space).
  347. """
  348. n_features = X.shape[0]
  349. rX = X.reshape((n_features,-1))
  350. pinvX = pinv(rX) if pinvX is None else pinvX
  351. P, D = {}, {}
  352. for key in list(mXs.keys()):
  353. mX = mXs[key].reshape((n_features,-1)) # called X_phi in paper
  354. C = np.dot(mX,pinvX)
  355. if isinstance(self.n_components,dict):
  356. U,s,V = randomized_svd(np.dot(C,rX),n_components=self.n_components[key],n_iter=self.n_iter,random_state=np.random.randint(10e5))
  357. else:
  358. U,s,V = randomized_svd(np.dot(C,rX),n_components=self.n_components,n_iter=self.n_iter,random_state=np.random.randint(10e5))
  359. P[key] = U
  360. D[key] = np.dot(U.T,C).T
  361. return P, D
  362. def _add_regularization(self,Y,mYs,lam,SVD=None,pre_reg=False):
  363. """ Prepares the data matrix and its marginalizations for the randomized_dpca solver (see paper)."""
  364. n_features = Y.shape[0]
  365. if not pre_reg:
  366. regY = np.hstack([Y.reshape((n_features,-1)),lam*np.eye(n_features)])
  367. else:
  368. regY = Y
  369. regY[:,-n_features:] = lam*eye(n_features)
  370. if not pre_reg:
  371. regmYs = OrderedDict()
  372. for key in list(mYs.keys()):
  373. regmYs[key] = np.hstack([mYs[key],np.zeros((n_features,n_features))])
  374. else:
  375. regmYs = mYs
  376. if SVD is not None:
  377. U,s,V = SVD
  378. M = ((s**2 + lam**2)**-1)[:,None]*U.T
  379. pregY = np.dot(np.vstack([V.T*s[None,:],lam*U]),M)
  380. else:
  381. pregY = np.dot(regY.reshape((n_features,-1)).T,np.linalg.inv(np.dot(Y.reshape((n_features,-1)),Y.reshape((n_features,-1)).T) + lam**2*np.eye(n_features)))
  382. return regY, regmYs, pregY
  383. def _fit(self, X, trialX=None, mXs=None, center=True, SVD=None, optimize=True):
  384. """ Fit the model on X
  385. Parameters
  386. ----------
  387. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  388. Training data, where n_samples in the number of samples
  389. and n_features_j is the number of the j-features (where the axis correspond
  390. to different parameters).
  391. trialX: array-like, shape (n_trials, n_samples, n_features_1, n_features_2, ...)
  392. Trial-by-trial data. Shape is similar to X but with an additional axis at the beginning
  393. with different trials. If different combinations of features have different number
  394. of trials, then set n_samples to the maximum number of trials and fill unoccupied data
  395. points with NaN.
  396. mXs: dict with values in the shape of X
  397. Marginalized data, should be the result of dpca._marginalize
  398. center: bool
  399. Centers data if center = True
  400. SVD: list of arrays
  401. Singular-value decomposition of the data. Don't provide!
  402. optimize: bool
  403. Flag to turn automatic optimization of regularization parameter on or off. Needed
  404. internally.
  405. """
  406. def flat2d(A):
  407. ''' Flattens all but the first axis of an ndarray, returns view. '''
  408. return A.reshape((A.shape[0],-1))
  409. # X = check_array(X)
  410. n_features = X.shape[0]
  411. # center data
  412. if center:
  413. X = X - np.mean(flat2d(X),1).reshape((n_features,) + len(self.labels)*(1,))
  414. # marginalize data
  415. if mXs is None:
  416. mXs = self._marginalize(X)
  417. # compute optimal regularization
  418. if self.opt_regularizer_flag and optimize:
  419. if self.debug > 0:
  420. print("Start optimizing regularization.")
  421. if trialX is None:
  422. raise ValueError('To optimize the regularization parameter, the trial-by-trial data trialX needs to be provided.')
  423. self._optimize_regularization(X,trialX)
  424. # add regularization
  425. if self.regularizer > 0:
  426. regX, regmXs, pregX = self._add_regularization(X,mXs,self.regularizer*np.sum(X**2),SVD=SVD)
  427. else:
  428. regX, regmXs, pregX = X, mXs, pinv(X.reshape((n_features,-1)))
  429. # compute closed-form solution
  430. self.P, self.D = self._randomized_dpca(regX,regmXs,pinvX=pregX)
  431. def _zero_mean(self,X):
  432. """ Subtracts the mean from each observable """
  433. return X - np.mean(X.reshape((X.shape[0],-1)),1).reshape((X.shape[0],) + (len(X.shape)-1)*(1,))
  434. def _roll_back(self,X,axes,invert=False):
  435. ''' Rolls all axis in list crossval_protect to the end (or inverts if invert=True) '''
  436. rX = X
  437. axes = np.sort(axes)
  438. if invert:
  439. for ax in reversed(axes):
  440. rX = np.rollaxis(rX,-1,start=ax)
  441. else:
  442. for ax in axes:
  443. rX = np.rollaxis(rX,ax,start=len(X.shape))
  444. return rX
  445. def _get_n_samples(self,trialX,protect=None):
  446. """ Computes the number of samples for each parameter combinations (except along protect) """
  447. n_unprotect = len(trialX.shape) - len(protect) - 1 if protect is not None else len(trialX.shape) - 1
  448. n_protect = len(protect) if protect is not None else 0
  449. return trialX.shape[0] - np.sum(np.isnan(trialX[(np.s_[:],) + (np.s_[:],)*n_unprotect + (0,)*n_protect]),0)
  450. def _check_protected(self,X,protect):
  451. ''' Checks if protect == None or, alternatively, if all protected axis are at the end '''
  452. if protect is None:
  453. protected = True
  454. else:
  455. # convert label in index
  456. protect = [self.labels.index(ax) for ax in protect]
  457. if set(protect) == set(np.arange(len(self.labels)-len(protect),len(self.labels))):
  458. protected = True
  459. else:
  460. protected = False
  461. print('Not all protected axis are at the end! While the algorithm will still work, the performance of the shuffling algorithm will substantially decrease due to unavoidable copies.')
  462. return protected
  463. def train_test_split(self,X,trialX,N_samples=None,sample_ax=0):
  464. """ Splits data in training and validation trial. To this end, we select one data-point in each observable for every
  465. combination of parameters (except along protected axis) for the validation set and average the remaining trial-by-trial
  466. data to get the training set.
  467. Parameters
  468. ----------
  469. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  470. Training data, where n_samples in the number of samples
  471. and n_features_j is the number of the j-features (where the axis correspond
  472. to different parameters).
  473. trialX: array-like, shape (n_trials, n_samples, n_features_1, n_features_2, ...)
  474. Trial-by-trial data. Shape is similar to X but with an additional axis at the beginning
  475. with different trials. If different combinations of features have different number
  476. of trials, then set n_samples to the maximum number of trials and fill unoccupied data
  477. points with NaN.
  478. N_samples: array-like with the same shape as X (except for protected axis).
  479. Number of trials in each condition. If None, computed from trial data.
  480. Returns
  481. -------
  482. trainX: array-like, same shape as X
  483. Training data
  484. blindX: array-like, same shape as X
  485. Validation data
  486. """
  487. def flat2d(A):
  488. ''' Flattens all but the first axis of an ndarray, returns view. '''
  489. return A.reshape((A.shape[0],-1))
  490. protect = self.protect
  491. n_samples = trialX.shape[-1] # number of samples
  492. n_unprotect = len(X.shape) - len(protect) if protect is not None else len(X.shape)
  493. n_protect = len(protect) if protect is not None else 0
  494. if sample_ax != 0:
  495. raise NotImplemented('The sample axis needs to come first.')
  496. # test if all protected axes lie at the end
  497. protected = self._check_protected(trialX,protect)
  498. # reorder matrix to protect certain axis (for speedup)
  499. if not(protected):
  500. # turn crossval_protect into index listX
  501. axes = [self.labels.index(ax) + 2 for ax in protect]
  502. # reorder matrix
  503. trialX = self._roll_back(trialX,axes)
  504. X = np.squeeze(self._roll_back(X[None,...],axes))
  505. # compute number of samples in each condition
  506. if N_samples is None:
  507. N_samples = self._get_n_samples(trialX,protect=self.protect)
  508. # get random indices
  509. idx = (np.random.rand(*N_samples.shape)*N_samples).astype(int)
  510. # select values
  511. blindX = np.empty(trialX.shape[1:])
  512. # iterate over multi_index
  513. it = np.nditer(np.empty(N_samples.shape), flags=['multi_index'])
  514. while not it.finished:
  515. blindX[it.multi_index + (np.s_[:],)*n_protect] = trialX[(idx[it.multi_index],) + it.multi_index + (np.s_[:],)*n_protect]
  516. it.iternext()
  517. # compute trainX
  518. trainX = (X*(N_samples/(N_samples-1))[(np.s_[:],)*n_unprotect + (None,)*n_protect] - blindX/(N_samples-1)[(np.s_[:],)*n_unprotect + (None,)*n_protect])
  519. # inverse rolled axis in blindX
  520. if not(protected):
  521. blindX = self._roll_back(blindX[...,None],axes,invert=True)[...,0]
  522. trainX = self._roll_back(trainX[...,None],axes,invert=True)[...,0]
  523. # remean datasets (both equally)
  524. trainX -= np.mean(flat2d(trainX),1)[(np.s_[:],) + (None,)*(len(X.shape)-1)]
  525. blindX -= np.mean(flat2d(blindX),1)[(np.s_[:],) + (None,)*(len(X.shape)-1)]
  526. return trainX, blindX
  527. def shuffle_labels(self,trialX):
  528. """ Shuffles *inplace* labels between conditions in trial-by-trial data, respecting the number of trials per condition.
  529. Parameters
  530. ----------
  531. trialX: array-like, shape (n_trials, n_samples, n_features_1, n_features_2, ...)
  532. Trial-by-trial data. Shape is similar to X but with an additional axis at the beginning
  533. with different trials. If different combinations of features have different number
  534. of trials, then set n_samples to the maximum number of trials and fill unoccupied data
  535. points with NaN.
  536. """
  537. # import shuffling algorithm from cython source
  538. protect = self.protect
  539. # test if all protected axes lie at the end
  540. protected = self._check_protected(trialX,protect)
  541. # reorder matrix to protect certain axis (for speedup)
  542. if not(protected):
  543. # turn crossval_protect into index list
  544. axes = [self.labels.index(ax) + 2 for ax in protect]
  545. # reorder matrix
  546. trialX = self._roll_back(trialX,axes)
  547. # reshape all non-protect axis into one vector
  548. original_shape = trialX.shape
  549. trialX = trialX.reshape((-1,) + trialX.shape[-len(protect):])
  550. # reshape all protected axis into one
  551. original_shape_protected = trialX.shape
  552. trialX = trialX.reshape((trialX.shape[0],-1))
  553. # shuffle within non-protected axis
  554. shuffle2D(trialX)
  555. # inverse reshaping of protected axis
  556. trialX = trialX.reshape(original_shape_protected)
  557. # inverse reshaping & sample axis
  558. trialX = trialX.reshape(original_shape)
  559. #trialX = np.rollaxis(trialX,0,len(original_shape))
  560. # inverse rolled axis in trialX
  561. if protected:
  562. trialX = self._roll_back(trialX,axes,invert=True)
  563. return trialX
  564. def significance_analysis(self,X,trialX,n_shuffles=100,n_splits=100,n_consecutive=1,axis=None,full=False):
  565. '''
  566. Cross-validated significance analysis of dPCA model. Here the generalization from the training
  567. to test data is tested by a simple classification measure in which one tries to predict the
  568. label of a validation test point from the training data. The performance is tested for n_splits
  569. test and training separations. The classification performance is then compared against
  570. the performance on data with randomly shuffled labels. Only if the performance is higher
  571. then the maximum in the shuffled data we regard the component as significant.
  572. Parameters
  573. ----------
  574. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  575. Training data, where n_samples in the number of samples
  576. and n_features_j is the number of the j-features (where the axis correspond
  577. to different parameters).
  578. trialX: array-like, shape (n_trials, n_samples, n_features_1, n_features_2, ...)
  579. Trial-by-trial data. Shape is similar to X but with an additional axis at the beginning
  580. with different trials. If different combinations of features have different number
  581. of trials, then set n_samples to the maximum number of trials and fill unoccupied data
  582. points with NaN.
  583. n_shuffles: integer
  584. Number of label shuffles over which the maximum is taken (default = 100, which
  585. is equivalent to p > 0.01)
  586. n_splits: integer
  587. Number of train-test splits per shuffle, from which the average performance is
  588. deduced.
  589. n_consecutive: integer
  590. Sometimes individual data points are deemed significant purely by chance. To reduced
  591. such noise one can demand that at least n consecutive data points are rated as significant.
  592. axis: None or True (default = None)
  593. Determines whether the significance is calculated over the last axis. More precisely,
  594. one is often interested in determining the significance of a component over time. In this
  595. case, set axis to True and make sure the last axis is time.
  596. full: Boolean (default = False)
  597. Whether or not all scores are returned. If False, only the significance matrix is returned.
  598. Returns
  599. -------
  600. masks: Dictionary
  601. Dictionary with keys corresponding to the marginalizations and with values that are
  602. binary nparrays that capture the significance of the demixed components.
  603. true_score: Dictionary (only returned when full = True)
  604. Dictionary with the scores of the data.
  605. scores: Dictionary (only returned when full = True)
  606. Dictionary with the scores of the shuffled data.
  607. '''
  608. assert axis in [None, True]
  609. def compute_mean_score(X,trialX,n_splits):
  610. K = 1 if axis is None else X.shape[-1]
  611. if type(self.n_components) == int:
  612. scores = {key : np.empty((self.n_components, n_splits, K)) for key in keys}
  613. else:
  614. scores = {key : np.empty((self.n_components[key], n_splits, K)) for key in keys}
  615. for shuffle in range(n_splits):
  616. print('.', end=' ')
  617. # do train-validation split
  618. trainX, validX = self.train_test_split(X,trialX)
  619. # fit a dPCA model to training data & transform validation data
  620. trainZ = self.fit_transform(trainX)
  621. validZ = self.transform(validX)
  622. # reshape data to match Cython input
  623. for key in keys:
  624. ncomps = self.n_components if type(self.n_components) == int else self.n_components[key]
  625. # mean over all axis not in key
  626. axset = self.marginalizations[key]
  627. axset = axset if type(axset) == set else set.union(*axset)
  628. axes = set(range(len(X.shape)-1)) - axset
  629. for ax in list(axes)[::-1]:
  630. trainZ[key] = np.mean(trainZ[key],axis=ax+1)
  631. validZ[key] = np.mean(validZ[key],axis=ax+1)
  632. # reshape
  633. if len(X.shape)-2 in axset and axis is not None:
  634. trainZ[key] = trainZ[key].reshape((ncomps,-1,K))
  635. validZ[key] = validZ[key].reshape((ncomps,-1,K))
  636. else:
  637. trainZ[key] = trainZ[key].reshape((ncomps,-1,1))
  638. validZ[key] = validZ[key].reshape((ncomps,-1,1))
  639. # compute classification score
  640. for key in keys:
  641. ncomps = self.n_components if type(self.n_components) == int else self.n_components[key]
  642. for comp in range(ncomps):
  643. scores[key][comp, shuffle] = classification(trainZ[key][comp],validZ[key][comp])
  644. for key in keys:
  645. scores[key] = np.nanmean(scores[key], axis=1)
  646. return scores
  647. if self.opt_regularizer_flag:
  648. print("Regularization not optimized yet; start optimization now.")
  649. self._optimize_regularization(X,trialX)
  650. keys = list(self.marginalizations.keys())
  651. keys.remove(self.labels[-1])
  652. # shuffling is in-place, so we need to copy the data
  653. trialX = trialX.copy()
  654. # compute score of original data
  655. print("Compute score of data: ", end=' ')
  656. true_score = compute_mean_score(X,trialX,n_splits)
  657. print("Finished.")
  658. # data collection
  659. scores = {key : [] for key in keys}
  660. # iterate over shuffles
  661. for it in range(n_shuffles):
  662. print("\rCompute score of shuffled data: ", str(it), "/", str(n_shuffles), end=' ')
  663. # shuffle labels
  664. self.shuffle_labels(trialX)
  665. # mean trial-by-trial data
  666. X = np.nanmean(trialX,axis=0)
  667. score = compute_mean_score(X,trialX,n_splits)
  668. for key in keys:
  669. scores[key].append(score[key])
  670. # binary mask, if data score is above maximum shuffled score make true
  671. masks = {}
  672. for key in keys:
  673. maxscore = np.amax(np.dstack(scores[key]),-1)
  674. masks[key] = true_score[key] > maxscore
  675. if n_consecutive > 1:
  676. for key in keys:
  677. mask = masks[key]
  678. for k in range(mask.shape[0]):
  679. masks[key][k,:] = denoise_mask(masks[key][k].astype(np.int32),n_consecutive)
  680. if full:
  681. return masks, true_score, scores
  682. else:
  683. return masks
  684. def transform(self, X, marginalization=None):
  685. """Apply the dimensionality reduction on X.
  686. X is projected on the first principal components previous extracted
  687. from a training set.
  688. Parameters
  689. ----------
  690. X: array-like, shape (n_samples, n_features_1, n_features_2, ...)
  691. Training data, where n_samples in the number of samples
  692. and n_features_j is the number of the j-features (where the axis correspond
  693. to different parameters).
  694. marginalization : str or None
  695. Marginalization subspace upon which to project, if None return dict
  696. with projections on all marginalizations
  697. Returns
  698. -------
  699. X_new : dict with arrays of the same shape as X
  700. Dictionary in which each key refers to one marginalization and the value is the
  701. latent component. If specific marginalization is given, returns only array
  702. """
  703. X = self._zero_mean(X)
  704. total_variance = np.sum(X**2)
  705. Xmargs = self._marginalize(X)
  706. def marginal_variances(marginal):
  707. ''' Computes the relative variance explained of each component
  708. within a marginalization
  709. '''
  710. D, P, Xmarg = self.D[marginal], self.P[marginal], Xmargs[marginal]
  711. return [(np.sum(Xmarg**2)-np.sum((Xmarg-np.outer(P[:, k], np.dot(D[:,k], Xmarg)))**2)) / total_variance for k in range(D.shape[1])]
  712. if marginalization is not None:
  713. D, Xr = self.D[marginalization], X.reshape((X.shape[0],-1))
  714. X_transformed = np.dot(D.T, Xr).reshape((D.shape[1],) + X.shape[1:])
  715. self.explained_variance_ratio_ = {marginalization : marginal_variances(marginalization)}
  716. else:
  717. X_transformed = {}
  718. self.explained_variance_ratio_ = {}
  719. for key in list(self.marginalizations.keys()):
  720. X_transformed[key] = np.dot(self.D[key].T, X.reshape((X.shape[0],-1))).reshape((self.D[key].shape[1],) + X.shape[1:])
  721. self.explained_variance_ratio_[key] = marginal_variances(key)
  722. return X_transformed
  723. def inverse_transform(self, X, marginalization):
  724. """ Transform data back to its original space, i.e.,
  725. return an input X_original whose transform would be X
  726. Parameters
  727. ----------
  728. X : array-like, shape (n_samples, n_components)
  729. New data, where n_samples is the number of samples
  730. and n_components is the number of components.
  731. Returns
  732. -------
  733. X_original array-like, shape (n_samples, n_features)
  734. """
  735. X = self._zero_mean(X)
  736. X_transformed = np.dot(self.P[marginalization],X.reshape((X.shape[0],-1))).reshape((self.P[marginalization].shape[0],) + X.shape[1:])
  737. return X_transformed
  738. def reconstruct(self, X, marginalization):
  739. """ Transform data first into reduced space before projecting
  740. it back into data space. Equivalent to inverse_transform(transform(X)).
  741. Parameters
  742. ----------
  743. X : array-like, shape (n_samples, n_components)
  744. New data, where n_samples is the number of samples
  745. and n_components is the number of components.
  746. Returns
  747. -------
  748. X_original array-like, shape (n_samples, n_features)
  749. """
  750. return self.inverse_transform(self.transform(X,marginalization),marginalization)

dPCA.py at commit 1def5b1, under MIT · at the source

Overview

Authors: Zhuangyi Jiang1, Ziang Liu1, Li Shi1, Fang Fang1,2,3,4, Shiming Tang3,5, Yang Zhou1,2,3,4
  1. Peking-Tsinghua Center for Life Sciences, Peking University,Beijing, China
  2. School of Psychological and Cognitive Sciences, Peking University,Beijing, China
  3. PKU-IDG/McGovern Institute for Brain Research, Peking University,Beijing, China
  4. School of Life Sciences, Peking University,Beijing, China
  5. Beijing Key Laboratory of Behavior and Mental Health, Beijing, China
Institutions: Peking University (China)
Journal: Nature communications, volume 17, issue 1, article 5909
Dates: received 21 September 2025; accepted 15 April 2026; published online 30 April 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41467-026-72510-9 · PMID 42062292 · PMCID PMC13338231 · OpenAlex W7159689184
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: non-human primate (organism), cognitive (subfield)
Methods: Connectivity, Statistics, Smoothing, state filtering, decompositions, Graphs, fMRI & imaging, Single-unit activity, calcium imaging, Physiology & signal measures
Keywords: Cortex, Cognitive neuroscience
MeSH: Learning*, Nerve Net*, Neurons*, Parietal Lobe*, Animals, Macaca mulatta, Male, Models, Neurological (* major topic)
Topic: Neural dynamics and brain function (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 95 references in the paper

Abstract

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

Repositories

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

machenslab/dPCA

License: MIT
State: the link answers, verified on 30 September 2026
Evidence: files inventoried
Commit: 1def5b15854638811a25257dc0b68074ab5a1be0, 29 July 2026
Languages: MATLAB (15), Python (4), Jupyter (1)
Size: 28 files, 20 scripts
Software Heritage: not archived
Found in: the text, “Demixed principal component analysis (dPCA)”
Holds: README, license file, environment (python/pyproject.toml, python/requirements.txt, python/setup.py), continuous integration, 1 notebook
Not found: CITATION.cff, tests, documentation
Tools: NumPy (4 files), Statistics and Machine Learning Toolbox (3 files), Numba (2 files), scikit-learn (2 files), SciPy (2 files)
Availability: 1 check, the latest on 30 September 2026: the link answers
  • 30 September 2026: the link answers
22 files

codeocean:8493005

License: none: the authors keep all their rights
State: cannot be verified, verified on 30 September 2026
Evidence: found in the paper
Software Heritage: not checked
Found in: “Code availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 30 September 2026: cannot be verified
  • 30 September 2026: cannot be verified
At the source:

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: codeocean:8493005

Read it in the paper: doi.org/10.1038/s41467-026-72510-9.

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:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 20 scripts, each with its path and the digest of its content;
  • 2 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

No dataset and no data link were found in the paper.

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 says that the data are available on request

Read it in the paper: doi.org/10.1038/s41467-026-72510-9.

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, 30 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 6 authors, 2 keywords, 8 MeSH terms, 3 funders, 94 references.

Cite

This paper

Jiang, Z., Liu, Z., Shi, L., Fang, F., Tang, S., & Zhou, Y. (2026). Single-neuron network topology governs neural computation and learning in primate cortex. Nature communications, 17(1), 5909. https://doi.org/10.1038/s41467-026-72510-9

BibTeX

@article{jiang2026single,
author = {Jiang, Zhuangyi and Liu, Ziang and Shi, Li and Fang, Fang and Tang, Shiming and Zhou, Yang},
title = {{Single-neuron network topology governs neural computation and learning in primate cortex}},
journal = {Nature communications},
year = {2026},
month = apr,
volume = {17},
number = {1},
pages = {5909},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-72510-9},
url = {https://doi.org/10.1038/s41467-026-72510-9},
pmid = {42062292},
pmcid = {PMC13338231}
}

RIS

TY - JOUR
AU - Jiang, Zhuangyi
AU - Liu, Ziang
AU - Shi, Li
AU - Fang, Fang
AU - Tang, Shiming
AU - Zhou, Yang
TI - Single-neuron network topology governs neural computation and learning in primate cortex
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/04/30
VL - 17
IS - 1
SP - 5909
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-72510-9
UR - https://doi.org/10.1038/s41467-026-72510-9
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-72510-9",
"type": "article-journal",
"title": "Single-neuron network topology governs neural computation and learning in primate cortex",
"container-title": "Nature communications",
"author": [
{
"family": "Jiang",
"given": "Zhuangyi"
},
{
"family": "Liu",
"given": "Ziang"
},
{
"family": "Shi",
"given": "Li"
},
{
"family": "Fang",
"given": "Fang"
},
{
"family": "Tang",
"given": "Shiming"
},
{
"family": "Zhou",
"given": "Yang"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "5909",
"DOI": "10.1038/s41467-026-72510-9",
"PMID": "42062292",
"PMCID": "PMC13338231",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-72510-9",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
30
]
]
}
}

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:10.1038/s41467-026-75959-w [code]
Charting higher-order models of brain function beyond pairwise interactions.
Journal: Nature communications
In common: Numba, Statistics and Machine Learning Toolbox, scikit-learn, 2 other tools, 7 references
[2] doi:10.1038/s41467-026-74466-2 [code]
Neuromorphic hierarchical modular reservoirs.
Journal: Nature communications
In common: Numba, scikit-learn, SciPy, 1 other tool, 7 references
[3] doi:10.1038/s41467-026-71725-0 [code]
Interactions across hemispheres in prefrontal cortex reflect global cognitive processing.
Journal: Nature communications
In common: Statistics and Machine Learning Toolbox, scikit-learn, SciPy, 1 other tool, non-human primate, cognitive, 5 references
[4] doi:10.1007/s12311-026-02042-x
The Cerebellar Connectome.
Journal: Cerebellum (London, England)
In common: 9 references
[5] doi:10.7554/elife.107518 [code]
Continuous flash suppression of neural responses and population orientation coding in macaque V1.
Journal: eLife
In common: Statistics and Machine Learning Toolbox, non-human primate, 2 references, author Shi-Ming Tang
[6] doi:10.1038/s41467-026-75585-6 [code]
Brain network dynamics reflect psychiatric illness status and transdiagnostic symptom profiles across health and disease.
Journal: Nature communications
In common: scikit-learn, SciPy, NumPy, 6 references
[7] doi:10.1038/s41467-026-71458-0 [code]
Early differential impact of MeCP2 mutations on functional networks in Rett syndrome patient-derived human cortical organoids.
Journal: Nature communications
In common: Numba, Statistics and Machine Learning Toolbox, scikit-learn, 2 other tools, 4 references
[8] doi:10.1002/hbm.70485 [code]
Exploring the Role of the Rich Club in Network Control of Neurocognitive States.
Journal: Human brain mapping
In common: Statistics and Machine Learning Toolbox, scikit-learn, SciPy, 1 other tool, 5 references
[9] doi:10.1038/s42003-026-09831-4 [code]
Comprehensive large-scale analyses reveal association between brain structure and cognitive ability during adolescence.
Journal: Communications biology
In common: Statistics and Machine Learning Toolbox, scikit-learn, SciPy, 1 other tool, 5 references
[10] doi:10.1371/journal.pbio.3003915 [code]
Noise-invariant representations of sound emerge along the canonical cortical hierarchy.
Journal: PLoS biology
In common: Numba, Statistics and Machine Learning Toolbox, scikit-learn, 2 other tools, 3 references

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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