OSCR

'Backpropagation and the brain' realized in cortical error neuron microcircuits.

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] § 2. Theory › 2.8. Synaptic plasticity ↔ numpy_model/src/microcircuit.py, lines 1973–2036 · score 0.64 · pre synaptic rate, post synaptic, synapses, computational, microcircuit, plasticity
  2. [2] § 6. Methods › 6.5. Simulations › 6.5.1. Cart pole and DMS task. ↔ numpy_model/cart-pole plot.ipynb, lines 85–111 · score 0.59 · penalty log, cart pole, epoch, reset, dt

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 · 2,444 lines · 86 KB · no license · 1 match

  1. import numpy as np
  2. import copy
  3. import inspect
  4. from scipy import signal
  5. #import matplotlib.pyplot as plt
  6. import logging
  7. import os.path as path
  8. # define activation functions
  9. def linear(x, slope=1.0, offset=0.0):
  10. x = x-offset
  11. return np.array(x*slope)
  12. def d_linear(x, slope=1.0, offset=0.0):
  13. x = x-offset
  14. return np.ones_like(x)*slope
  15. def relu(x, slope=1.0, offset=0.0):
  16. # offset is a bias
  17. return np.maximum(x-offset, 0, np.array(x))*slope
  18. def d_relu(x, slope=1.0, offset=0.0):
  19. return np.heaviside(x-offset, 0)*slope
  20. def soft_relu(x, slope=1.0, offset=0.0):
  21. x = x-offset
  22. exp = np.exp(-10.0*x)
  23. return np.log((1 + exp)/exp)/10.0*slope
  24. def d_soft_relu(x, slope=1.0, offset=0.0):
  25. x = x-offset
  26. return 1 / (1 + np.exp(-10.0*x))*slope
  27. def logistic(x, slope=1.0, offset=0.0):
  28. return 1/(1 + np.exp(-(x-offset)*slope))
  29. def d_logistic(x, slope=1.0, offset=0.0):
  30. y = logistic(x, slope=slope, offset=offset)
  31. return y * (1.0 - y)
  32. def tanh(x, slope=1.0, offset=0.0):
  33. return np.tanh((x-offset)*slope)
  34. def d_tanh(x, slope=1.0, offset=0.0):
  35. y = tanh(x*slope, offset=offset)
  36. return 1 - y**2
  37. # def tanh_offset(x, slope=1.0):
  38. # # tanh offset by +1 to produce positive output only
  39. # return np.tanh(x*slope) + 1
  40. # def d_tanh_offset(x, slope=1.0):
  41. # y = tanh(x*slope)
  42. # return 1 - y**2
  43. def hard_sigmoid(x, slope=1.0, offset=0.0):
  44. # see defintion at torch.nn.Hardsigmoid
  45. x = x-offset
  46. return np.maximum(0, np.minimum(1, x/6 + 1/2), np.array(x))*slope
  47. def d_hard_sigmoid(x, slope=1.0, offset=0.0):
  48. x = x-offset
  49. return np.heaviside(x/6 + 1/2, 0) * np.heaviside(1 - (x/6 + 1/2), 0)*slope
  50. # define dict between activations and derivatives
  51. dict_d_activation = {
  52. "linear": d_linear,
  53. "relu": d_relu,
  54. "soft_relu": d_soft_relu,
  55. "logistic": d_logistic,
  56. "tanh": d_tanh,
  57. # "tanh_offset": d_tanh_offset,
  58. "hard_sigmoid": d_hard_sigmoid
  59. }
  60. # cosine similarity between tensors
  61. def cos_sim(A, B):
  62. if A.ndim == 1 and B.ndim == 1:
  63. return A.T @ B / np.linalg.norm(A) / np.linalg.norm(B)
  64. else:
  65. return np.trace(A.T @ B) / np.linalg.norm(A) / np.linalg.norm(B)
  66. def dist(A, B):
  67. return np.linalg.norm(A-B)
  68. def deg(cos, precision=1e-6):
  69. if np.abs(1.0 - cos) < precision:
  70. return 0.0
  71. else:
  72. # calculates angle in deg from cosine
  73. return np.arccos(cos) * 180 / np.pi
  74. def deepcopy_array(array):
  75. """ makes a deep copy of an array of np-arrays """
  76. out = [nparray.copy() for nparray in array]
  77. return out.copy()
  78. def MSE(output, target):
  79. return np.linalg.norm(output - target)**2
  80. def moving_average(x, w):
  81. return np.convolve(x, np.ones(w), 'valid') / w
  82. def label_to_onehot(label, classes):
  83. onehot = np.zeros(classes)
  84. onehot[label] = 1
  85. return onehot
  86. def unison_shuffled_copies(a, b, rng):
  87. """
  88. Shuffles two np arrays in the same way
  89. by mtrw at SO
  90. """
  91. assert len(a) == len(b)
  92. p = rng.permutation(len(a))
  93. return a[p], b[p]
  94. def convert_eta_to_matrix_for_skip(eta_fw, WPP, layers):
  95. """
  96. Converts a learning rate array of shape (# of WPP x # of WPP)
  97. to an array of shape (# of neurons # of neurons) == shape of WPP
  98. """
  99. # pad eta_fw with zeros to align shape with WPP
  100. eta_fw = np.array(eta_fw)
  101. eta_fw = np.pad(eta_fw, ((1,0),(0,1)), 'constant')
  102. # repeat to match number of neurons in each layer
  103. eta_fw = np.repeat(eta_fw, layers, axis=0)
  104. eta_fw = np.repeat(eta_fw, layers, axis=1)
  105. return eta_fw
  106. def init_nonhierarchical_WPP(WPP_hierarchical_range, layers, WPP_skip_connection_range):
  107. """
  108. Generates a WPP matrix (for class errormc_model) which connects all representation units.
  109. WPP_hierarchical_range: uniform range for init values (hierarchical feedback)
  110. WPP_skip_connection_range: uniform range for init values (skip connections)
  111. """
  112. WPP_init = np.zeros(shape=(sum(layers), sum(layers)))
  113. # we populate WPP_init the same way as BPP_init (below), then transpose
  114. for idx, _ in enumerate(layers[:-1]):
  115. lower_row = np.sum(layers[:idx], dtype=int)
  116. upper_row = np.sum(layers[:idx+1])
  117. lower_col = np.sum(layers[:idx+1])
  118. upper_col = np.sum(layers[:idx+2])
  119. # generate hierarchical part of WPP
  120. WPP_init[lower_row:upper_row, lower_col:upper_col] = np.random.uniform(WPP_hierarchical_range[0], \
  121. WPP_hierarchical_range[1], size=(layers[idx], layers[idx+1]))
  122. # generate skip connections of WPP
  123. WPP_init[:lower_col, upper_col:] = np.random.uniform(WPP_skip_connection_range[0], \
  124. WPP_skip_connection_range[1], size=(WPP_init[:lower_col, upper_col:].shape))
  125. return [WPP_init.T]
  126. def init_nonhierarchical_BII(BII_hierarchical_range, error_layers, BII_skip_connection_range):
  127. """
  128. Generates a BII matrix (for class errormc_model) which connects all error units.
  129. BII_hierarchical_range: uniform range for init values (hierarchical feedback)
  130. BII_skip_connection_range: uniform range for init values (skip connections)
  131. """
  132. # filter out "none" layers (no error units in first layer)
  133. BII_init = np.zeros(shape=(sum(filter(None, error_layers)), sum(filter(None, error_layers))))
  134. for idx, _ in enumerate(error_layers[:-1]):
  135. lower_row = np.sum(error_layers[:idx], dtype=int)
  136. upper_row = np.sum(error_layers[:idx+1])
  137. lower_col = np.sum(error_layers[:idx+1])
  138. upper_col = np.sum(error_layers[:idx+2])
  139. # generate hierarchical part of BII
  140. BII_init[lower_row:upper_row, lower_col:upper_col] = np.random.uniform(BII_hierarchical_range[0], \
  141. BII_hierarchical_range[1], size=(error_layers[idx], error_layers[idx+1]))
  142. # generate skip connections of BII
  143. BII_init[:lower_col, upper_col:] = np.random.uniform(BII_skip_connection_range[0], \
  144. BII_skip_connection_range[1], size=(BII_init[:lower_col, upper_col:].shape))
  145. return [BII_init]
  146. def init_viz_cort_conn(layers, error_layers,
  147. init_WPP_range, init_BII_range):
  148. # load connectivity matrix based on Markov et al., 2014
  149. cwd = path.dirname(path.abspath(__file__))
  150. conn = np.genfromtxt(cwd + '/connectivity.csv', delimiter=',').T
  151. if len(layers[1:]) > len(conn) or len(error_layers[1:]) > len(conn):
  152. raise ValueError(f"Too many layers, connectivity only implemented for {len(conn)} layers.")
  153. # rescale skip connections using hierarchical connections and connectivity matrix
  154. assert len(layers) == len(error_layers)
  155. fw_rescale_mask = conn.copy()
  156. # insert row/column for input
  157. fw_rescale_mask = np.pad(fw_rescale_mask,((1, 0), (1, 0)), mode='constant')
  158. fw_rescale_mask[1,0] = 1.0
  159. WPP_init = np.zeros(shape=(sum(layers), sum(layers)))
  160. for idx, _ in enumerate(layers):
  161. for idy, _ in enumerate(layers):
  162. lower_row = np.sum(layers[:idx], dtype=int)
  163. upper_row = np.sum(layers[:idx+1], dtype=int)
  164. lower_col = np.sum(layers[:idy], dtype=int)
  165. upper_col = np.sum(layers[:idy+1], dtype=int)
  166. weight_lower = init_WPP_range[0] * fw_rescale_mask[idx,idy]
  167. weight_upper = init_WPP_range[1] * fw_rescale_mask[idx,idy]
  168. # generate hierarchical part of WPP
  169. WPP_init[lower_row:upper_row, lower_col:upper_col] = np.random.uniform(weight_lower, \
  170. weight_upper, size=(layers[idx], layers[idy]))
  171. # ignore first entry (no backprojection from V1)
  172. error_layers = error_layers[1:]
  173. # no modifications for input
  174. bw_rescale_mask = conn.copy()
  175. BII_init = np.zeros(shape=(sum(error_layers), sum(error_layers)))
  176. for idx, _ in enumerate(error_layers):
  177. for idy, _ in enumerate(error_layers):
  178. lower_row = np.sum(error_layers[:idx], dtype=int)
  179. upper_row = np.sum(error_layers[:idx+1], dtype=int)
  180. lower_col = np.sum(error_layers[:idy], dtype=int)
  181. upper_col = np.sum(error_layers[:idy+1], dtype=int)
  182. weight_lower = init_BII_range[0] * bw_rescale_mask[idx,idy]
  183. weight_upper = init_BII_range[1] * bw_rescale_mask[idx,idy]
  184. # generate hierarchical part of WPP
  185. BII_init[lower_row:upper_row, lower_col:upper_col] = np.random.uniform(weight_lower, \
  186. weight_upper, size=(error_layers[idx], error_layers[idy]))
  187. WPP_init = np.tril(WPP_init)
  188. BII_init = np.triu(BII_init)
  189. return [WPP_init], [BII_init]
  190. # create matrix with list of matrices on diagonal
  191. # user7328723 @ SO
  192. def diag_mat(rem=[], result=np.empty((0, 0))):
  193. if not rem:
  194. return result
  195. m = rem.pop(0)
  196. result = np.block(
  197. [
  198. [result, np.zeros((result.shape[0], m.shape[1]))],
  199. [np.zeros((m.shape[0], result.shape[1])), m],
  200. ]
  201. )
  202. return diag_mat(rem, result)
  203. def convert_learning_rate_for_skip(eta, method):
  204. """
  205. Takes a list of learning rates, e.g. eta=[1,1,1],
  206. and converts it into a 2d array needed for skip connections
  207. by using three different rules:
  208. - fill_diag: uses eta to fill only the diagonal, all others = 0
  209. - fill_all: all learning rates *to* a given layer are equal
  210. - fill_scaled: uses viz cx conn to scale learning rate
  211. Returns an array of entries:
  212. [[eta10, 0, 0, ...],
  213. [eta20, eta21, 0, ...],
  214. [eta30, eta31, eta32, ...]]
  215. """
  216. assert method in ["fill_diag", "fill_all", "fill_scaled"], \
  217. f"Unknown conversion method for learning rate (not in 'fill_diag', 'fill_all', 'fill_scaled')"
  218. logging.info(f"Adapting eta_fw according to {method}")
  219. eta = np.array(eta)
  220. # converts the learning
  221. if method == 'fill_diag':
  222. eta = np.diag(eta)
  223. elif method == 'fill_all':
  224. eta = np.repeat(eta[:,np.newaxis], len(eta), axis=1)
  225. elif method == 'fill_scaled':
  226. cwd = path.dirname(path.abspath(__file__))
  227. conn = np.genfromtxt(cwd + '/connectivity.csv', delimiter=',').T
  228. # insert row/column for input
  229. conn = np.pad(conn,((1, 0), (1, 0)), mode='constant')
  230. conn[1,0] = 1.0
  231. conn = np.tril(conn[1:][:len(eta),:len(eta)])
  232. eta = np.diag(eta) @ conn
  233. return np.tril(eta)
  234. class base_model:
  235. """ This class implements a generic microcircuit model """
  236. def __init__(self, bw_connection_mode, dWPP_use_activation, dt, Tpres,
  237. model, activation, layers, uP_init, uI_init, WPP_init,
  238. WIP_init, BPP_init, BPI_init, gl, gden, gbas, gapi, gnI,
  239. gntgt, eta_fw, eta_bw, eta_PI, eta_IP, seed=123, WT_noise=0.0):
  240. self.rng = np.random.RandomState(seed)
  241. # connection_mode: skip or layered
  242. self.bw_connection_mode = bw_connection_mode
  243. # whether to use activation in updates of WPP
  244. self.dWPP_use_activation = dWPP_use_activation
  245. self.model = model # FA, BP or PBP
  246. self.layers = layers
  247. self.uP = deepcopy_array(uP_init)
  248. self.uI = deepcopy_array(uI_init)
  249. # we also set up a buffer of voltages,
  250. # which corresponds to the value at the last time step
  251. self.uP_old = deepcopy_array(self.uP)
  252. self.uI_old = deepcopy_array(self.uI)
  253. # if a list of activations has been passed, use it
  254. if isinstance(activation, list):
  255. self.activation = activation
  256. # else, set same activation for all layers
  257. else:
  258. self.activation = [activation for layer in layers[1:]]
  259. self.d_activation = [dict_d_activation[activation.__name__] for activation in self.activation]
  260. # whether the target provided should be a rate or voltage
  261. self.rate_target = False
  262. # define the compartment voltages
  263. self.vbas = [np.zeros_like(uP) for uP in self.uP]
  264. self.vden = [np.zeros_like(uI) for uI in self.uI]
  265. self.vapi = [np.zeros_like(uP) for uP in self.uP[:-1]]
  266. # and make copies
  267. self.vbas_old = deepcopy_array(self.vbas)
  268. self.vden_old = deepcopy_array(self.vden)
  269. self.vapi_old = deepcopy_array(self.vapi)
  270. self.WPP = deepcopy_array(WPP_init)
  271. self.WIP = deepcopy_array(WIP_init)
  272. self.BPP = deepcopy_array(BPP_init)
  273. self.BPI = deepcopy_array(BPI_init)
  274. self.dWPP = [np.zeros(shape=WPP.shape) for WPP in self.WPP]
  275. self.dWIP = [np.zeros(shape=WIP.shape) for WIP in self.WIP]
  276. self.dBPP = [np.zeros(shape=BPP.shape) for BPP in self.BPP]
  277. self.dBPI = [np.zeros(shape=BPI.shape) for BPI in self.BPI]
  278. # set perfect transpose weights for BP
  279. if self.model == "BP":
  280. self.set_weights(BPP = [WPP.T for WPP in self.WPP[1:]])
  281. # noise level in setting transpose weights
  282. self.WT_noise = WT_noise
  283. # noise matrix (calculated once)
  284. self.BPP_noise = [self.rng.uniform(-self.WT_noise, self.WT_noise, size=BPP.shape) for BPP in self.BPP]
  285. for i, _ in enumerate(self.BPP):
  286. # add noise
  287. self.BPP[i] += self.BPP_noise[i]
  288. self.gl = gl
  289. self.gden = gden
  290. self.gbas = gbas
  291. self.gapi = gapi
  292. self.gnI = gnI
  293. self.gntgt = gntgt
  294. self.Time = 0 # initialize a model timer
  295. self.dt = dt
  296. self.Tpres = Tpres
  297. self.taueffP, self.taueffP_notgt, self.taueffI = self.calc_taueff()
  298. # learning rates
  299. self.eta_fw = eta_fw
  300. self.eta_bw = eta_bw
  301. self.eta_IP = eta_IP
  302. self.eta_PI = eta_PI
  303. # calculate lookahead
  304. self.uP_breve = [self.prospective_voltage(
  305. self.uP[i],
  306. self.uP_old[i],
  307. self.taueffP[i]) for i in range(len(self.uP))]
  308. self.uI_breve = [self.prospective_voltage(
  309. self.uI[i],
  310. self.uI_old[i],
  311. self.taueffI[i]) for i in range(len(self.uI))]
  312. # calculate rate of lookahead: phi(ubreve)
  313. self.rP_breve = [self.activation[i](self.uP_breve[i])
  314. for i in range(len(self.uP_breve))]
  315. try:
  316. self.rI_breve = [self.activation[i+1](self.uI_breve[i])
  317. for i in range(len(self.uI_breve))]
  318. except (IndexError, ValueError):
  319. logging.info("rI_breve not defined")
  320. self.r0 = np.zeros(self.layers[0])
  321. self.dWPP_r_low_pass = False
  322. self.dWPP_post_low_pass = False
  323. def init_record(self, rec_per_steps=1, rec_MSE=False,
  324. rec_error=False, rec_input=False, rec_target=False,
  325. rec_WPP=False, rec_WIP=False, rec_BPP=False,
  326. rec_BII=False,
  327. rec_BPI=False, rec_dWPP=False, rec_dWIP=False,
  328. rec_dBPP=False, rec_dBPI=False, rec_dBII=False, rec_uP=False,
  329. rec_uP_breve=False, rec_rP_breve=False,
  330. rec_rP_breve_HI=False, rec_uI=False,
  331. rec_uI_breve=False, rec_rI_breve=False, rec_vbas=False,
  332. rec_vapi=False, rec_vapi_noise=False, rec_noise=False,
  333. rec_epsilon=False, rec_epsilon_LO=False,
  334. rec_lat_mismatch=False):
  335. # records the values of the variables given in var_array
  336. # e.g. WPP, BPP, uP_breve
  337. # rec_per_steps sets after how many steps data is recorded
  338. if rec_MSE:
  339. self.MSE_time_series = []
  340. self.MSE_val_time_series = []
  341. self.MSE_test_time_series = []
  342. if rec_error:
  343. self.error_time_series = []
  344. if rec_input:
  345. self.input_time_series = []
  346. if rec_target:
  347. self.target_time_series = []
  348. if rec_WPP:
  349. self.WPP_time_series = []
  350. if rec_WIP:
  351. self.WIP_time_series = []
  352. if rec_BPP:
  353. self.BPP_time_series = []
  354. if rec_BII:
  355. self.BII_time_series = []
  356. if rec_BPI:
  357. self.BPI_time_series = []
  358. if rec_dWPP:
  359. self.dWPP_time_series = []
  360. if rec_dWIP:
  361. self.dWIP_time_series = []
  362. if rec_dBPP:
  363. self.dBPP_time_series = []
  364. if rec_dBPI:
  365. self.dBPI_time_series = []
  366. if rec_dBII:
  367. self.dBII_time_series = []
  368. if rec_uP:
  369. self.uP_time_series = []
  370. if rec_uP_breve:
  371. self.uP_breve_time_series = []
  372. if rec_rP_breve:
  373. self.rP_breve_time_series = []
  374. if rec_rP_breve_HI:
  375. self.rP_breve_HI_time_series = []
  376. if rec_uI:
  377. self.uI_time_series = []
  378. if rec_uI_breve:
  379. self.uI_breve_time_series = []
  380. if rec_rI_breve:
  381. self.rI_breve_time_series = []
  382. if rec_vbas:
  383. self.vbas_time_series = []
  384. if rec_vapi:
  385. self.vapi_time_series = []
  386. if rec_vapi_noise:
  387. self.vapi_noise_time_series = []
  388. if rec_noise:
  389. self.noise_time_series = []
  390. if rec_epsilon:
  391. self.epsilon_time_series = []
  392. if rec_epsilon_LO:
  393. self.epsilon_LO_time_series = []
  394. if rec_lat_mismatch:
  395. self.lat_mismatch_time_series = []
  396. self.rec_per_steps = rec_per_steps
  397. self.rec_counter = 0
  398. def record_step(self, target=None, MSE_only=False, testing=False, validation=False):
  399. if hasattr(self, 'MSE_time_series') and target is not None:
  400. if hasattr(self, "rec_rate_MSE"):
  401. # if MSE should be rate, and target is rate already
  402. if self.rec_rate_MSE and self.rate_target:
  403. mse = MSE(self.rP_breve[-1], target)
  404. # if MSE should be rate, but target is voltage:
  405. # convert target to rate before recording
  406. elif self.rec_rate_MSE and not self.rate_target:
  407. mse = MSE(self.rP_breve[-1], self.activation[-1](target[-1]))
  408. # if MSE should be voltage, and rate is voltage
  409. elif not self.rec_rate_MSE and not self.rate_target:
  410. mse = MSE(self.uP_breve[-1], target)
  411. else:
  412. raise ValueError("MSE cannot be based on voltage if target is rate")
  413. elif self.rate_target:
  414. mse = MSE(self.rP_breve[-1], target)
  415. else:
  416. mse = MSE(self.uP_breve[-1], target)
  417. if testing:
  418. self.MSE_test_time_series.append(mse.copy())
  419. elif validation:
  420. self.MSE_val_time_series.append(mse.copy())
  421. else:
  422. self.MSE_time_series.append(mse.copy())
  423. if not MSE_only:
  424. if hasattr(self, 'error_time_series') and target is not None:
  425. if self.rate_target:
  426. self.error_time_series.append(
  427. self.rP_breve[-1] - target
  428. )
  429. else:
  430. self.error_time_series.append(
  431. self.uP_breve[-1] - target
  432. )
  433. if hasattr(self, 'input_time_series'):
  434. self.input_time_series.append(copy.deepcopy(self.r0))
  435. if hasattr(self, 'target_time_series') and target is not None:
  436. self.target_time_series.append(target.copy())
  437. if hasattr(self, 'WPP_time_series'):
  438. self.WPP_time_series.append(copy.deepcopy(self.WPP))
  439. if hasattr(self, 'WIP_time_series'):
  440. self.WIP_time_series.append(copy.deepcopy(self.WIP))
  441. if hasattr(self, 'BPP_time_series'):
  442. self.BPP_time_series.append(copy.deepcopy(self.BPP))
  443. if hasattr(self, 'BII_time_series'):
  444. self.BII_time_series.append(copy.deepcopy(self.BII))
  445. if hasattr(self, 'BPI_time_series'):
  446. self.BPI_time_series.append(copy.deepcopy(self.BPI))
  447. if hasattr(self, 'dWPP_time_series') and hasattr(self, 'dWPP'):
  448. self.dWPP_time_series.append(copy.deepcopy(self.dWPP))
  449. if hasattr(self, 'dWIP_time_series')and hasattr(self, 'dWIP'):
  450. self.dWIP_time_series.append(copy.deepcopy(self.dWIP))
  451. if hasattr(self, 'dBPP_time_series')and hasattr(self, 'dBPP'):
  452. self.dBPP_time_series.append(copy.deepcopy(self.dBPP))
  453. if hasattr(self, 'dBPI_time_series')and hasattr(self, 'dBPI'):
  454. self.dBPI_time_series.append(copy.deepcopy(self.dBPI))
  455. if hasattr(self, 'dBII_time_series')and hasattr(self, 'dBII'):
  456. self.dBII_time_series.append(copy.deepcopy(self.dBII))
  457. if hasattr(self, 'uP_time_series'):
  458. self.uP_time_series.append(copy.deepcopy(self.uP))
  459. if hasattr(self, 'uP_breve_time_series'):
  460. self.uP_breve_time_series.append(copy.deepcopy(self.uP_breve))
  461. if hasattr(self, 'rP_breve_time_series'):
  462. self.rP_breve_time_series.append(copy.deepcopy(self.rP_breve))
  463. if hasattr(self, 'rP_breve_HI_time_series'):
  464. self.rP_breve_HI_time_series.append(copy.deepcopy(self.rP_breve_HI))
  465. if hasattr(self, 'uI_time_series'):
  466. self.uI_time_series.append(copy.deepcopy(self.uI))
  467. if hasattr(self, 'uI_breve_time_series'):
  468. self.uI_breve_time_series.append(copy.deepcopy(self.uI_breve))
  469. if hasattr(self, 'rI_breve_time_series'):
  470. self.rI_breve_time_series.append(copy.deepcopy(self.rI_breve))
  471. if hasattr(self, 'vbas_time_series'):
  472. self.vbas_time_series.append(copy.deepcopy(self.vbas))
  473. if hasattr(self, 'vapi_time_series'):
  474. self.vapi_time_series.append(copy.deepcopy(self.vapi))
  475. if hasattr(self, 'vapi_noise_time_series'):
  476. self.vapi_noise_time_series.append(copy.deepcopy(self.vapi_noise))
  477. if hasattr(self, 'noise_time_series'):
  478. self.noise_time_series.append(copy.deepcopy(self.noise))
  479. if hasattr(self, 'epsilon_time_series'):
  480. self.epsilon_time_series.append(copy.deepcopy(self.epsilon))
  481. if hasattr(self, 'epsilon_LO_time_series'):
  482. self.epsilon_LO_time_series.append(copy.deepcopy(self.epsilon_LO))
  483. if hasattr(self, 'lat_mismatch_time_series'):
  484. vbashat = [self.gbas / (self.gl + self.gbas + self.gapi) * vbas for vbas in self.vbas]
  485. vbashat[-1] = (self.gl + self.gbas + self.gapi) / (self.gl + self.gbas) * vbashat[-1]
  486. lat_mismatch = [uI_breve - vbashat for uI_breve, vbashat in zip(self.uI_breve, vbashat[1:])]
  487. self.lat_mismatch_time_series.append(copy.deepcopy(lat_mismatch))
  488. def calc_taueff(self):
  489. # calculate tau_eff for pyramidals and interneuron
  490. # taueffP is one value per layer
  491. taueffP = []
  492. for i in self.uP:
  493. taueffP.append(1 / (self.gl + self.gbas + self.gapi))
  494. taueffP[-1] = 1 / (self.gl + self.gbas + self.gntgt)
  495. # tau_eff for output layer in absence of target
  496. taueffP_notgt = [1 / (self.gl + self.gbas)]
  497. taueffI = []
  498. for i in self.uI:
  499. taueffI.append(1 / (self.gl + self.gden + self.gnI))
  500. return taueffP, taueffP_notgt, taueffI
  501. def get_conductances(self):
  502. return self.gl, self.gden, self.gbas, self.gapi, self.gnI, self.gntgt
  503. def get_weights(self):
  504. return self.WPP, self.WIP, self.BPP, self.BPI
  505. def set_weights(self, model=None, WPP=None, WIP=None, BPP=None,
  506. BPI=None):
  507. # if another model is given, copy its weights
  508. if hasattr(model, '__dict__'):
  509. WPP, WIP, BPP, BPI = model.get_weights()
  510. logging.info(f"Copying weights from model {model}")
  511. if WPP is not None: self.WPP = deepcopy_array(WPP)
  512. if WIP is not None: self.WIP = deepcopy_array(WIP)
  513. if BPP is not None: self.BPP = deepcopy_array(BPP)
  514. if BPI is not None: self.BPI = deepcopy_array(BPI)
  515. def set_self_predicting_state(self):
  516. # set WIP and BPI to values corresponding to self-predicting state
  517. for i in range(len(self.BPP)):
  518. self.BPI[i] = - self.BPP[i].copy()
  519. for i in range(len(self.WIP)-1):
  520. self.WIP[i] = self.gbas * (self.gl + self.gden) / (
  521. self.gden * (self.gl + self.gbas +
  522. self.gapi)) * self.WPP[i+1].copy()
  523. if len(self.layers) > 2:
  524. self.WIP[-1] = self.gbas * (self.gl + self.gden) / (self.gden *
  525. (self.gl + self.gbas)) * self.WPP[-1].copy()
  526. def get_voltages(self):
  527. return self.uP, self.uI
  528. def get_old_voltages(self):
  529. return self.uP_old, self.uI_old
  530. def get_breve_voltages(self):
  531. return self.uP_breve, self.uI_breve
  532. def set_voltages(self, model=None, uP=None, uP_old=None, uP_breve=None, uI=None, uI_old=None, uI_breve=None):
  533. # if another model is given, copy its voltages
  534. if hasattr(model, '__dict__'):
  535. uP, uI = model.get_voltages()
  536. uP_old, uI_old = model.get_old_voltages()
  537. uP_breve, uI_breve = model.get_breve_voltages()
  538. logging.info(f"Copying voltages from model {model}")
  539. if uP is not None:
  540. for i in range(len(self.layers)-1):
  541. self.uP[i] = copy.deepcopy(uP[i])
  542. if uP_old is not None:
  543. for i in range(len(self.layers)-1):
  544. self.uP_old[i] = copy.deepcopy(uP_old[i])
  545. if uP_breve is not None:
  546. for i in range(len(self.layers)-1):
  547. self.uP_breve[i] = copy.deepcopy(uP_breve[i])
  548. if uI is not None:
  549. self.uI = copy.deepcopy(uI)
  550. if uI_old is not None:
  551. self.uI_old = copy.deepcopy(uI_old)
  552. if uI_breve is not None:
  553. self.uI_breve = copy.deepcopy(uI_breve)
  554. def calc_vapi(self, rPvec, BPP_mat, rIvec, BPI_mat):
  555. """
  556. returns apical voltages in pyramidals of a given layer
  557. input: rPvec: vector of rates from pyramidal voltages in
  558. output layer
  559. WPP_mat: matrix connecting pyramidal to pyramidal
  560. rIvec: vector of rates from interneuron voltages in output layer
  561. BPI_mat: matrix connecting interneurons to pyramidals
  562. """
  563. return BPP_mat @ rPvec + BPI_mat @ rIvec
  564. def calc_vbas(self, rPvec, WPP_mat):
  565. """
  566. returns basal voltages in pyramidals of a given layer
  567. input: rPvec: vector of rates from pyramidal voltages in
  568. layer below
  569. WPP_mat: matrix connecting pyramidal to pyramidal
  570. """
  571. return WPP_mat @ rPvec
  572. def calc_vden(self, rPvec, WIP_mat):
  573. """
  574. returns dendritic voltages in inteneurons
  575. input: rPvec: vector of rates from pyramidal voltages in
  576. layer below
  577. WIP_mat: matrix connecting pyramidal to pyramidal
  578. """
  579. return WIP_mat @ rPvec
  580. def prospective_voltage(self, uvec, uvec_old, tau, dt=None):
  581. """
  582. returns an approximation of the lookahead of voltage
  583. vector u at current time
  584. """
  585. if dt == None:
  586. dt = self.dt
  587. return uvec_old + tau * (uvec - uvec_old) / dt
  588. def evolve_system(self, r0=None, u_tgt=None, learn_weights=True,
  589. learn_lat_weights=True, learn_bw_weights=False,
  590. record=True, testing=False, validation=False, compare_dWPP=False):
  591. """
  592. evolves the system by one time step:
  593. updates synaptic weights and voltages given input rate r0
  594. """
  595. # increase timer by dt and round float to nearest dt
  596. self.Time = np.round(self.Time + self.dt,
  597. decimals=int(np.round(-np.log10(self.dt))))
  598. if testing or validation:
  599. self.duP, self.duI = self.evolve_voltages(r0, u_tgt=None) # includes recalc of rP_breve
  600. else:
  601. self.duP, self.duI = self.evolve_voltages(r0, u_tgt) # includes recalc of rP_breve
  602. if learn_weights or learn_bw_weights or learn_lat_weights or compare_dWPP:
  603. self.dWPP, self.dWIP, self.dBPP, self.dBPI = self.evolve_synapses(r0)
  604. # apply evolution
  605. for i in range(len(self.duP)):
  606. self.uP[i] += self.duP[i]
  607. for i in range(len(self.duI)):
  608. self.uI[i] += self.duI[i]
  609. if learn_weights:
  610. for i in range(len(self.dWPP)):
  611. self.WPP[i] += self.dWPP[i]
  612. if learn_lat_weights:
  613. for i in range(len(self.dWIP)):
  614. self.WIP[i] += self.dWIP[i]
  615. for i in range(len(self.dBPI)):
  616. self.BPI[i] += self.dBPI[i]
  617. if learn_bw_weights:
  618. for i in range(len(self.dBPP)):
  619. self.BPP[i] += self.dBPP[i]
  620. # logging.warning("SETTING SPS AFTER dW")
  621. # self.set_self_predicting_state()
  622. # record step
  623. if hasattr(self, 'rec_per_steps') and record:
  624. self.rec_counter += 1
  625. if self.rec_counter % self.rec_per_steps == 0:
  626. self.rec_counter = 0
  627. self.record_step(target=u_tgt)
  628. # during testing and validation, we record MSE for all steps
  629. if (testing or validation) and record:
  630. self.record_step(target=u_tgt, MSE_only=True, testing=testing, validation=validation)
  631. def evolve_voltages(self, r0=None, u_tgt=None):
  632. """
  633. Evolves the pyramidal and interneuron voltages by one dt
  634. using r0 as input rates
  635. """
  636. self.duP = [np.zeros(shape=uP.shape) for uP in self.uP]
  637. self.duI = [np.zeros(shape=uI.shape) for uI in self.uI]
  638. # same for dendritic voltages and rates
  639. self.rP_breve_old = deepcopy_array(self.rP_breve)
  640. self.rI_breve_old = deepcopy_array(self.rI_breve)
  641. if self.r0 is not None:
  642. self.r0_old = self.r0.copy()
  643. self.vbas_old = deepcopy_array(self.vbas)
  644. self.vden_old = deepcopy_array(self.vden)
  645. self.vapi_old = deepcopy_array(self.vapi)
  646. # calculate lookahead
  647. self.uP_breve = [self.prospective_voltage(self.uP[i],
  648. self.uP_old[i],
  649. self.taueffP[i]) for i in range(len(self.uP))]
  650. self.uI_breve = [self.prospective_voltage(self.uI[i],
  651. self.uI_old[i],
  652. self.taueffI[i]) for i in range(len(self.uI))]
  653. # calculate rate of lookahead: phi(ubreve)
  654. self.rP_breve = [self.activation[i](self.uP_breve[i])
  655. for i in range(len(self.uP_breve))]
  656. self.rI_breve = [self.activation[i+1](self.uI_breve[i])
  657. for i in range(len(self.uI_breve))]
  658. self.r0 = r0
  659. # before modifying uP and uI, we need to save copies
  660. # for future calculation of u_breve
  661. self.uP_old = deepcopy_array(self.uP)
  662. self.uI_old = deepcopy_array(self.uI)
  663. self.vbas, self.vapi, self.vden = self.calc_dendritic_updates(r0, u_tgt)
  664. self.duP, self.duI = self.calc_somatic_updates(u_tgt)
  665. return self.duP, self.duI
  666. def calc_dendritic_updates(self, r0=None, u_tgt=None):
  667. # calculate dendritic voltages from lookahead
  668. if r0 is not None:
  669. self.vbas[0] = self.WPP[0] @ self.r0
  670. for i in range(1, len(self.layers)-1):
  671. self.vbas[i] = self.calc_vbas(self.rP_breve[i-1],
  672. self.WPP[i])
  673. for i in range(len(self.WIP)):
  674. if self.bw_connection_mode == 'skip':
  675. self.vden[0] = self.calc_vden(self.rP_breve[-2],
  676. self.WIP[-1])
  677. elif self.bw_connection_mode == 'layered':
  678. self.vden[i] = self.calc_vden(self.rP_breve[i],
  679. self.WIP[i])
  680. for i in range(len(self.layers)-2):
  681. if self.bw_connection_mode == 'skip':
  682. self.vapi[i] = self.calc_vapi(self.rP_breve[-1],
  683. self.BPP[i],
  684. self.rI_breve[-1],
  685. self.BPI[i])
  686. elif self.bw_connection_mode == 'layered':
  687. self.vapi[i] = self.calc_vapi(self.rP_breve[i+1],
  688. self.BPP[i],
  689. self.rI_breve[i],
  690. self.BPI[i])
  691. return self.vbas, self.vapi, self.vden
  692. def calc_somatic_updates(self, u_tgt=None):
  693. """
  694. calculates somatic updates from dendritic potentials
  695. """
  696. # update somatic potentials
  697. for i in range(len(self.uI)):
  698. ueffI = self.taueffI[i] * (self.gden *
  699. self.vden[i] + self.gnI * self.uP_breve[i+1])
  700. delta_uI = (ueffI - self.uI[i]) / self.taueffI[i]
  701. self.duI[i] = self.dt * delta_uI
  702. for i in range(0, len(self.layers)-2):
  703. ueffP = self.taueffP[i] * (self.gbas *
  704. self.vbas[i] + self.gapi * self.vapi[i])
  705. delta_uP = (ueffP - self.uP[i]) / self.taueffP[i]
  706. self.duP[i] = self.dt * delta_uP
  707. if u_tgt is not None:
  708. ueffP = self.taueffP[-1] * (self.gbas *
  709. self.vbas[-1] + self.gntgt * u_tgt[-1])
  710. delta_uP = (ueffP - self.uP[-1]) / self.taueffP[-1]
  711. else:
  712. ueffP = self.taueffP_notgt[-1] * (self.gbas *
  713. self.vbas[-1])
  714. delta_uP = (ueffP - self.uP[-1]) / self.taueffP[-1]
  715. self.duP[-1] = self.dt * delta_uP
  716. return self.duP, self.duI
  717. def evolve_synapses(self, r0, learn_WIP=True, learn_BPI=True):
  718. """
  719. evolves all synapses by a dt
  720. plasticity of WPP
  721. """
  722. self.dWPP = [np.zeros(shape=WPP.shape) for WPP in self.WPP]
  723. self.dWIP = [np.zeros(shape=WIP.shape) for WIP in self.WIP]
  724. self.dBPP = [np.zeros(shape=BPP.shape) for BPP in self.BPP]
  725. self.dBPI = [np.zeros(shape=BPI.shape) for BPI in self.BPI]
  726. if self.dWPP_use_activation:
  727. if self.dWPP_r_low_pass:
  728. if r0 is not None:
  729. self.r0_LO_old += self.dt / self.tauLO * (self.r0_old - self.r0_LO_old)
  730. # logging.info("updating WPP0")
  731. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  732. self.rP_breve[0] - self.activation[0](self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[0]),
  733. self.r0_LO_old)
  734. for i in range(1, len(self.WPP)-1):
  735. self.r_LO_old[i-1] += self.dt / self.tauLO * (self.rP_breve_old[i-1] - self.r_LO_old[i-1])
  736. # hidden layers
  737. # logging.info(f"updating WPP{i}")
  738. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(
  739. self.rP_breve[i] - self.activation[i](self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[i]),
  740. self.r_LO_old[i-1])
  741. # output layer
  742. self.r_LO_old[-2] += self.dt / self.tauLO * (self.rP_breve_old[-2] - self.r_LO_old[-2])
  743. # logging.info("updating WPP-1")
  744. if len(self.layers) > 2:
  745. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(
  746. self.rP_breve[-1] - self.activation[-1](self.gbas / (self.gl + self.gbas) * self.vbas_old[-1]),
  747. self.r_LO_old[-2])
  748. elif self.dWPP_post_low_pass:
  749. if r0 is not None:
  750. self.dWPP_post_LO_old[0] += self.dt / self.tauLO * \
  751. (self.rP_breve[0] - self.activation[0](self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[0]) \
  752. - self.dWPP_post_LO_old[0])
  753. # logging.info("updating WPP0")
  754. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(self.dWPP_post_LO_old[0], self.r0_old)
  755. for i in range(1, len(self.WPP)-1):
  756. self.dWPP_post_LO_old[i] += self.dt / self.tauLO * \
  757. (self.rP_breve[i] - self.activation[i](self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[i]) \
  758. - self.dWPP_post_LO_old[i])
  759. # hidden layers
  760. # logging.info(f"updating WPP{i}")
  761. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(self.dWPP_post_LO_old[i], self.rP_breve_old[i-1])
  762. # output layer
  763. self.dWPP_post_LO_old[-1] += self.dt / self.tauLO * \
  764. (self.rP_breve[-1] - self.activation[-1](self.gbas / (self.gl + self.gbas) * self.vbas_old[-1]) \
  765. - self.dWPP_post_LO_old[-1])
  766. # logging.info("updating WPP-1")
  767. if len(self.layers) > 2:
  768. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(self.dWPP_post_LO_old[-1], self.rP_breve_old[-2])
  769. else:
  770. if len(self.layers) == 2:
  771. # logging.info("updating WPP0")
  772. # input layer
  773. if r0 is not None:
  774. # if the model no hidden layers
  775. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  776. self.rP_breve[0] - self.activation[0](self.gbas / (self.gl + self.gbas) * self.vbas_old[0]),
  777. self.r0_old)
  778. else:
  779. # input layer
  780. if r0 is not None:
  781. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  782. self.rP_breve[0] - self.activation[0](self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[0]),
  783. self.r0_old)
  784. # hidden layers
  785. for i in range(1, len(self.WPP)-1):
  786. # logging.info(f"updating WPP{i}")
  787. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(
  788. self.rP_breve[i] - self.activation[i](self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[i]),
  789. self.rP_breve_old[i-1])
  790. # output layer
  791. # logging.info("updating WPP-1")
  792. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(
  793. self.rP_breve[-1] - self.activation[-1](self.gbas / (self.gl + self.gbas) * self.vbas_old[-1]),
  794. self.rP_breve_old[-2])
  795. else:
  796. if self.dWPP_r_low_pass:
  797. if r0 is not None:
  798. self.r0_LO_old += self.dt / self.tauLO * (self.r0_old - self.r0_LO_old)
  799. # logging.info("updating WPP0")
  800. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  801. self.uP_breve[0] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[0]),
  802. self.r0_LO_old)
  803. for i in range(1, len(self.WPP)-1):
  804. self.r_LO_old[i-1] += self.dt / self.tauLO * (self.rP_breve_old[i-1] - self.r_LO_old[i-1])
  805. # hidden layers
  806. # logging.info(f"updating WPP{i}")
  807. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(
  808. self.uP_breve[i] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[i]),
  809. self.r_LO_old[i-1])
  810. # output layer
  811. self.r_LO_old[-2] += self.dt / self.tauLO * (self.rP_breve_old[-2] - self.r_LO_old[-2])
  812. # logging.info("updating WPP-1")
  813. if len(self.layers) > 2:
  814. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(
  815. self.uP_breve[-1] - (self.gbas / (self.gl + self.gbas) * self.vbas_old[-1]),
  816. self.r_LO_old[-2])
  817. elif self.dWPP_post_low_pass:
  818. if r0 is not None:
  819. self.dWPP_post_LO_old[0] += self.dt / self.tauLO * \
  820. (self.uP_breve[0] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[0]) \
  821. - self.dWPP_post_LO_old[0])
  822. # logging.info("updating WPP0")
  823. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(self.dWPP_post_LO_old[0], self.r0_old)
  824. for i in range(1, len(self.WPP)-1):
  825. self.dWPP_post_LO_old[i] += self.dt / self.tauLO * \
  826. (self.uP_breve[i] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[i]) \
  827. - self.dWPP_post_LO_old[i])
  828. # hidden layers
  829. # logging.info(f"updating WPP{i}")
  830. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(self.dWPP_post_LO_old[i], self.r_old[i-1])
  831. # output layer
  832. self.dWPP_post_LO_old[-1] += self.dt / self.tauLO * \
  833. (self.uP_breve[-1] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[-1]) \
  834. - self.dWPP_post_LO_old[-1])
  835. # logging.info("updating WPP-1")
  836. if len(self.layers) > 2:
  837. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(self.dWPP_post_LO_old[-1], self.r_old[-2])
  838. else:
  839. if len(self.layers) == 2:
  840. # input layer
  841. if r0 is not None:
  842. # logging.info("updating WPP0")
  843. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  844. self.uP_breve[0] - (self.gbas / (self.gl + self.gbas) * self.vbas_old[0]),
  845. self.r0_old)
  846. else:
  847. if r0 is not None:
  848. # logging.info("updating WPP0")
  849. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  850. self.uP_breve[0] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[0]),
  851. self.r0_old)
  852. # hidden layers
  853. for i in range(1, len(self.WPP)-1):
  854. # logging.info(f"updating WPP{i}")
  855. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(
  856. self.uP_breve[i] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[i]),
  857. self.rP_breve_old[i-1])
  858. # output layer
  859. # logging.info("updating WPP-1")
  860. if len(self.layers) > 2:
  861. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(
  862. self.uP_breve[-1] - (self.gbas / (self.gl + self.gbas) * self.vbas_old[-1]),
  863. self.rP_breve_old[-2])
  864. """
  865. plasticity of WIP
  866. """
  867. if learn_WIP:
  868. if self.bw_connection_mode == 'skip':
  869. self.dWIP[-1] = self.dt * self.eta_IP[-1] * np.outer(
  870. self.rI_breve[-1] - self.activation[-1](self.gden / (self.gl + self.gden) * self.vden_old[-1]),
  871. self.rP_breve_old[-2])
  872. elif self.bw_connection_mode == 'layered':
  873. for i in range(len(self.WIP)):
  874. self.dWIP[i] = self.dt * self.eta_IP[i] * np.outer(
  875. self.rI_breve[i] - self.activation[i+1](self.gden / (self.gl + self.gden) * self.vden_old[i]),
  876. self.rP_breve_old[i])
  877. """
  878. plasticity of BPI
  879. """
  880. if learn_BPI:
  881. for i in range(0, len(self.BPI)):
  882. if self.eta_PI[i] != 0.0:
  883. if self.bw_connection_mode == 'skip':
  884. self.dBPI[i] = self.dt * self.eta_PI[i] * np.outer(-self.vapi_old[i], self.rI_breve_old[-1])
  885. elif self.bw_connection_mode == 'layered':
  886. self.dBPI[i] = self.dt * self.eta_PI[i] * np.outer(-self.vapi_old[i], self.rI_breve_old[i])
  887. """
  888. plasticity of BPP
  889. """
  890. if self.model == 'FA':
  891. # do nothing
  892. pass
  893. elif self.model == 'BP':
  894. # self.set_weights(BPP = [WPP.T for WPP in self.WPP[1:]])
  895. # add noise to transpose
  896. BPP_noised = [BPP_noise + WPP.T for BPP_noise, WPP in zip(self.BPP_noise, self.WPP[1:])]
  897. self.set_weights(BPP=BPP_noised)
  898. # lateral weight is still perfect copy
  899. self.set_weights(BPI = [-BPP for BPP in self.BPP])
  900. return self.dWPP, self.dWIP, self.dBPP, self.dBPI
  901. class noise_model(base_model):
  902. """ This class inherits all properties from the base model class and adds the function to add noise """
  903. def __init__(self, bw_connection_mode, dWPP_use_activation, dt, dtxi, tauHP, tauLO, Tpres, noise_scale, alpha,
  904. inter_low_pass, pyr_hi_pass, dWPP_low_pass, dWPP_r_low_pass, dWPP_post_low_pass, gate_regularizer,
  905. noise_type, noise_mode,
  906. model, activation, layers,
  907. uP_init, uI_init, WPP_init, WIP_init, BPP_init, BPI_init,
  908. gl, gden, gbas, gapi, gnI, gntgt,
  909. eta_fw, eta_bw, eta_PI, eta_IP, seed=123, **kwargs):
  910. # init base_model with same settings
  911. super().__init__(bw_connection_mode=bw_connection_mode, dWPP_use_activation=dWPP_use_activation, dt=dt, Tpres=Tpres,
  912. model=model, activation=activation, layers=layers,
  913. uP_init=uP_init, uI_init=uI_init,
  914. WPP_init=WPP_init, WIP_init=WIP_init, BPP_init=BPP_init, BPI_init=BPI_init,
  915. gl=gl, gden=gden, gbas=gbas, gapi=gapi, gnI=gnI, gntgt=gntgt,
  916. eta_fw=eta_fw, eta_bw=eta_bw, eta_PI=eta_PI, eta_IP=eta_IP, seed=seed)
  917. self.rng = np.random.RandomState(seed)
  918. self.d_activation = [dict_d_activation[activation.__name__] for activation in self.activation]
  919. # type of noise (OU or white)
  920. self.noise_type = noise_type
  921. # mode of noise injection (order vapi or uP or uP_adative)
  922. self.noise_mode = noise_mode
  923. # for uP_adaptive, we need epsilon: measures angle between BPP, WPP.T
  924. self.epsilon = [np.float64(1.0) for BPP in self.BPP]
  925. if noise_mode == 'uP_adaptive':
  926. self.noise_deg = kwargs.get('noise_deg')
  927. self.tau_eps = kwargs.get('taueps')
  928. if noise_type == 'OU':
  929. self.tauxi = kwargs.get('tauxi')
  930. # low-pass filtered version of epsilon
  931. self.epsilon_LO = deepcopy_array(self.epsilon)
  932. # whether to low-pass filter the interneuron dendritic input
  933. self.inter_low_pass = inter_low_pass
  934. # whether to high-pass filter rPbreve for updates of BPP
  935. self.pyr_hi_pass = pyr_hi_pass
  936. # whether to low-pass filter updates of WPP
  937. self.dWPP_low_pass = dWPP_low_pass
  938. self.dWPP_r_low_pass = dWPP_r_low_pass
  939. self.dWPP_post_low_pass = dWPP_post_low_pass
  940. # whether to gate application of the regularizer
  941. self.gate_regularizer = gate_regularizer
  942. # whether to use phi' B phi' as regularizer
  943. self.varphi_regularizer = kwargs.get('varphi_regularizer', False)
  944. if self.varphi_regularizer:
  945. self.d_rP = [np.diag(np.zeros(shape=uP.shape)) for uP in self.uP]
  946. # noise time scale
  947. self.dtxi = dtxi
  948. # decimals of dt
  949. self.dt_decimals = int(np.round(-np.log10(self.dt)))
  950. # synaptic time constant (sets the low-pass filter of interneuron)
  951. self.tauHP = tauHP
  952. self.tauLO = tauLO
  953. # gaussian noise properties
  954. self.noise_scale = noise_scale
  955. self.noise = [np.zeros(shape=uP.shape) for uP in self.uP]
  956. # self.noise_breve = [np.zeros(shape=uP.shape) for uP in self.uP]
  957. # we need a new variable: vapi after noise has been added
  958. # i.e. vapi = BPP rP + BPI rI (as usual), and vapi_noise = vapi + noise
  959. self.vapi_noise = deepcopy_array(self.vapi)
  960. # init a counter for time steps after which to resample noise
  961. self.noise_counter = 0
  962. self.noise_total_counts = np.round(self.dtxi / self.dt, decimals=self.dt_decimals)
  963. # init a high-pass filtered version of rP_breve
  964. self.rP_breve_HI = deepcopy_array(self.rP_breve)
  965. # init a low-pass filtered version of dWPP
  966. self.dWPP_LO = [np.zeros(shape=WPP.shape) for WPP in self.WPP]
  967. self.dWPP_post_LO_old = [np.zeros(shape=rP_breve.shape) for rP_breve in self.rP_breve]
  968. self.r0_LO_old = np.zeros(shape=layers[0])
  969. self.r_LO_old = [np.zeros(shape=rP_breve.shape) for rP_breve in self.rP_breve]
  970. # regularizer for backward weights
  971. self.alpha = alpha
  972. def evolve_system(self, r0=None, u_tgt=None, learn_weights=True, learn_lat_weights=True, learn_bw_weights=True, \
  973. record=True, testing=False, validation=False, compare_dWPP=False):
  974. """
  975. This overwrites the vanilla system evolution and implements
  976. additional noise
  977. """
  978. # # update which backwards weights to learn
  979. # if learn_bw_weights and self.Time % self.Tbw == 0:
  980. # # logging.info(f"Current time: {self.Time}s")
  981. # self.active_bw_syn = 0 if self.active_bw_syn == len(self.BPP) - 1 else self.active_bw_syn + 1
  982. # # logging.info(f"Learning backward weights to layer {self.active_bw_syn + 1}")
  983. # self.noise_counter = 0
  984. # calculate voltage evolution, including low pass on interneuron synapses
  985. # see calc_dendritic updates below
  986. if (testing or validation):
  987. self.duP, self.duI = self.evolve_voltages(r0, u_tgt=None, inject_noise=learn_bw_weights) # includes recalc of rP_breve
  988. else:
  989. self.duP, self.duI = self.evolve_voltages(r0, u_tgt, inject_noise=learn_bw_weights) # includes recalc of rP_breve
  990. if learn_weights or learn_lat_weights or compare_dWPP:
  991. self.dWPP, self.dWIP, _, self.dBPI = self.evolve_synapses(r0)
  992. if learn_bw_weights or compare_dWPP:
  993. self.dBPP = self.evolve_bw_synapses()
  994. # apply evolution
  995. for i in range(len(self.duP)):
  996. self.uP[i] += self.duP[i]
  997. for i in range(len(self.duI)):
  998. self.uI [i]+= self.duI[i]
  999. if learn_weights:
  1000. if self.dWPP_low_pass:
  1001. # calculate lo-passed update of WPP
  1002. self.dWPP_LO = self.calc_dWPP_LO()
  1003. for i in range(len(self.dWPP_LO)):
  1004. self.WPP[i] += self.dWPP_LO[i]
  1005. else:
  1006. for i in range(len(self.dWPP)):
  1007. self.WPP[i] += self.dWPP[i]
  1008. if learn_lat_weights:
  1009. for i in range(len(self.dWIP)):
  1010. self.WIP[i] += self.dWIP[i]
  1011. for i in range(len(self.dBPI)):
  1012. self.BPI[i] += self.dBPI[i]
  1013. if learn_bw_weights:
  1014. for i in range(len(self.dBPP)):
  1015. self.BPP[i] += self.dBPP[i]
  1016. # record step
  1017. if hasattr(self, 'rec_per_steps') and record:
  1018. self.rec_counter += 1
  1019. if self.rec_counter % self.rec_per_steps == 0:
  1020. self.rec_counter = 0
  1021. self.record_step(target=u_tgt)
  1022. # during testing and validation, we record MSE for all steps
  1023. if (testing or validation) and record:
  1024. self.record_step(target=u_tgt, MSE_only=True, testing=testing, validation=validation)
  1025. # increase timer
  1026. self.Time = np.round(self.Time + self.dt, decimals=self.dt_decimals)
  1027. # calculate hi-passed rP_breve for synapse BPP
  1028. self.rP_breve_HI = self.calc_rP_breve_HI()
  1029. def evolve_voltages(self, r0=None, u_tgt=None, inject_noise=False):
  1030. """
  1031. Overwrites voltage evolution:
  1032. Evolves the pyramidal and interneuron voltages by one dt
  1033. using r0 as input rates
  1034. >> Injects noise into vapi
  1035. """
  1036. self.duP = [np.zeros(shape=uP.shape) for uP in self.uP]
  1037. self.duI = [np.zeros(shape=uI.shape) for uI in self.uI]
  1038. # same for dendritic voltages and rates
  1039. self.rP_breve_old = deepcopy_array(self.rP_breve)
  1040. self.rI_breve_old = deepcopy_array(self.rI_breve)
  1041. if self.r0 is not None:
  1042. self.r0_old = self.r0.copy()
  1043. self.vbas_old = deepcopy_array(self.vbas)
  1044. self.vden_old = deepcopy_array(self.vden)
  1045. self.vapi_old = deepcopy_array(self.vapi)
  1046. # calculate lookahead
  1047. self.uP_breve = [self.prospective_voltage(self.uP[i], self.uP_old[i], self.taueffP[i]) for i in range(len(self.uP))]
  1048. self.uI_breve = [self.prospective_voltage(self.uI[i], self.uI_old[i], self.taueffI[i]) for i in range(len(self.uI))]
  1049. # calculate rate of lookahead: phi(ubreve)
  1050. self.rP_breve = [self.activation[i](self.uP_breve[i]) for i in range(len(self.uP_breve))]
  1051. self.rI_breve = [self.activation[i+1](self.uI_breve[i]) for i in range(len(self.uI_breve))]
  1052. self.r0 = r0
  1053. # before modifying uP and uI, we need to save copies
  1054. # for future calculation of u_breve
  1055. self.uP_old = deepcopy_array(self.uP)
  1056. self.uI_old = deepcopy_array(self.uI)
  1057. self.vbas, self.vapi, self.vapi_noise, self.vden = self.calc_dendritic_updates(r0, u_tgt)
  1058. # inject noise into newly calculated vapi before calculating update du
  1059. if inject_noise:
  1060. for i in range(len(self.BPP)):
  1061. self.generate_noise(layer=i, noise_type=self.noise_type, noise_scale=self.noise_scale)
  1062. self.duP, self.duI = self.calc_somatic_updates(u_tgt, inject_noise=inject_noise)
  1063. return self.duP, self.duI
  1064. def calc_somatic_updates(self, u_tgt=None, inject_noise=False):
  1065. """
  1066. this overwrites the somatic update rules
  1067. only difference to super: uses vapi_noise instead of vapi
  1068. """
  1069. # update somatic potentials
  1070. for i in range(len(self.uI)):
  1071. ueffI = self.taueffI[i] * (self.gden * self.vden[i] + self.gnI * self.uP_breve[i+1])
  1072. delta_uI = (ueffI - self.uI[i]) / self.taueffI[i]
  1073. self.duI[i] = self.dt * delta_uI
  1074. for i in range(0, len(self.layers)-2):
  1075. if inject_noise:
  1076. ueffP = self.taueffP[i] * (self.gbas * self.vbas[i] + self.gapi * self.vapi_noise[i])
  1077. else:
  1078. ueffP = self.taueffP[i] * (self.gbas * self.vbas[i] + self.gapi * self.vapi[i])
  1079. delta_uP = (ueffP - self.uP[i]) / self.taueffP[i]
  1080. self.duP[i] = self.dt * delta_uP
  1081. if u_tgt is not None:
  1082. ueffP = self.taueffP[-1] * (self.gbas * self.vbas[-1] + self.gntgt * u_tgt[-1])
  1083. delta_uP = (ueffP - self.uP[-1]) / self.taueffP[-1]
  1084. else:
  1085. ueffP = self.taueffP_notgt[-1] * (self.gbas * self.vbas[-1])
  1086. delta_uP = (ueffP - self.uP[-1]) / self.taueffP[-1]
  1087. self.duP[-1] = self.dt * delta_uP
  1088. return self.duP, self.duI
  1089. def generate_noise(self, layer, noise_type, noise_scale):
  1090. """
  1091. this function generates noise for a given layer
  1092. the noise is added to the apical potential in function calculate_dendritic_updates
  1093. there are two noise types:
  1094. - hold_white_noise: adds steps of width dtxi and height sampled from normal distribution
  1095. - OU: ornstein-uhlenbeck noise, i.e. low-pass filtered white noise updated at every dtxi
  1096. """
  1097. # if dtxi timesteps have passed, sample new noise
  1098. if np.all(self.noise[layer] == 0) or self.noise_counter % self.noise_total_counts == 0:
  1099. if self.noise_mode == 'uP_adaptive':
  1100. None # FIX THIS MODE
  1101. # # iterate over hidden layers
  1102. # for i in range(len(self.layers)-2):
  1103. # # first, we reset all noise values
  1104. # self.noise[i] = np.zeros(shape=self.uP[i].shape)
  1105. # self.vapi_noise[i] = self.vapi[i]
  1106. # # 'layer' is the layer which currently has active noise injection/bw learning
  1107. # if i == layer:
  1108. # # calculate Jacobian alignment factor epsilon
  1109. # self.epsilon[layer] = 1/2 * (1 - self.uP_breve[layer] @ self.BPP[layer] @ self.rP_breve[-1] \
  1110. # / np.linalg.norm(self.uP_breve[layer]) / np.linalg.norm(self.BPP[layer] @ self.rP_breve[-1]))
  1111. # # update low-pass filtered version of epsilon, with filter time-constant Tpres
  1112. # self.epsilon_LO[layer] += self.dt/self.tau_eps * (self.epsilon[layer] - self.epsilon_LO[layer])
  1113. # # generate noise, rescaled with epsilon
  1114. # # if epsilon is below threshold, do not inject noise
  1115. # if self.epsilon_LO[layer] > 1/2 * (1 - np.cos(self.noise_deg * np.pi/180)):
  1116. # white_noise = noise_scale[layer] * self.epsilon_LO[layer] * np.array([np.random.normal(0, np.abs(x)) for x in self.uP[layer]])
  1117. elif self.noise_mode == 'const':
  1118. # add noise with magnitude of rescaled uP
  1119. white_noise = noise_scale[layer] * self.rng.normal(0, 1, size=self.uP[layer].shape)
  1120. elif self.noise_mode == 'uP':
  1121. # add noise with magnitude of rescaled uP
  1122. white_noise = noise_scale[layer] * np.array([self.rng.normal(0, np.abs(x)) for x in self.uP[layer]])
  1123. elif self.noise_mode == 'uP_breve':
  1124. # add noise with magnitude of rescaled uP_breve
  1125. white_noise = noise_scale[layer] * np.array([self.rng.normal(0, np.abs(x)) for x in self.uP_breve[layer]])
  1126. elif self.noise_mode == 'vapi':
  1127. # add noise with magnitude of rescaled vapi
  1128. white_noise = noise_scale[layer] * np.array([self.rng.normal(0, np.abs(x)) for x in self.vapi[layer]])
  1129. # the noise will be added to vapi, depending on the mode
  1130. if self.noise_type == 'hold_white_noise':
  1131. self.noise[layer] = white_noise
  1132. elif self.noise_type == 'OU':
  1133. self.noise[layer] += 1 / (self.tauxi) * (- self.dt * self.noise[layer] + np.sqrt(self.dt * self.tauxi) * white_noise)
  1134. self.noise_counter = 0
  1135. self.noise_counter += 1
  1136. def calc_rP_breve_HI(self):
  1137. # updates the high-passed instantaneous rate rP_breve_HI which is used to update BPP
  1138. # High-pass has the form d v_out = d v_in - dt/tau * v_out
  1139. for i in range(len(self.rP_breve)):
  1140. self.rP_breve_HI[i] += (self.rP_breve[i] - self.rP_breve_old[i]) - self.dt / self.tauHP * self.rP_breve_HI[i]
  1141. # self.rP_breve_HI[-1] += (self.rP_breve[-1] - self.rP_breve_old[-1]) - self.dt / self.tauHP * self.rP_breve_HI[-1]
  1142. return self.rP_breve_HI
  1143. def calc_dWPP_LO(self):
  1144. # updates the low-passed update of WPP
  1145. # Low-pass has the form d v_out = dt/tau (v_in - v_out)
  1146. for i in range(len(self.dWPP_LO)):
  1147. self.dWPP_LO[i] += self.dt / self.tauLO * (self.dWPP[i] - self.dWPP_LO[i])
  1148. return self.dWPP_LO
  1149. def calc_dendritic_updates(self, r0=None, u_tgt=None):
  1150. """
  1151. this overwrites the dendritic updates by adding a low-pass on
  1152. the interneuron voltages
  1153. and adds vapi_noise somatic voltage after noise injection
  1154. """
  1155. # calculate dendritic voltages from lookahead
  1156. if r0 is not None:
  1157. self.vbas[0] = self.WPP[0] @ self.r0
  1158. for i in range(1, len(self.layers)-1):
  1159. self.vbas[i] = self.calc_vbas(self.rP_breve[i-1], self.WPP[i])
  1160. for i in range(len(self.WIP)):
  1161. if self.bw_connection_mode == 'skip':
  1162. # add slow response to dendritic compartment of interneurons
  1163. if self.inter_low_pass:
  1164. self.vden[0] += self.dt / self.tauLO * (self.calc_vden(self.rP_breve[-2], self.WIP[-1]) - self.vden[0])
  1165. else:
  1166. # else, instant response
  1167. self.vden[0] = self.calc_vden(self.rP_breve[-2], self.WIP[-1])
  1168. elif self.bw_connection_mode == 'layered':
  1169. if self.inter_low_pass:
  1170. self.vden[i] += self.dt / self.tauLO * (self.calc_vden(self.rP_breve[i], self.WIP[i]) - self.vden[i])
  1171. else:
  1172. self.vden[i] = self.calc_vden(self.rP_breve[i], self.WIP[i])
  1173. for i in range(0, len(self.layers)-2):
  1174. if self.bw_connection_mode == 'skip':
  1175. self.vapi[i] = self.calc_vapi(self.rP_breve[-1], self.BPP[i], self.rI_breve[-1], self.BPI[i])
  1176. elif self.bw_connection_mode == 'layered':
  1177. self.vapi[i] = self.calc_vapi(self.rP_breve[i+1], self.BPP[i], self.rI_breve[i], self.BPI[i])
  1178. self.vapi_noise[i] = self.vapi[i] + self.noise[i]
  1179. return self.vbas, self.vapi, self.vapi_noise, self.vden
  1180. def evolve_bw_synapses(self):
  1181. # evolve the synapses of BPPs
  1182. self.dBPP = [np.zeros(shape=BPP.shape) for BPP in self.BPP]
  1183. if self.varphi_regularizer:
  1184. for i in range(len(self.BPP)):
  1185. self.d_rP[i] = np.diag(self.d_activation[i](self.uP_breve[i]))
  1186. for i in range(len(self.BPP)):
  1187. if self.bw_connection_mode == 'skip':
  1188. if self.pyr_hi_pass:
  1189. r_pre = self.rP_breve_HI[-1]
  1190. else:
  1191. r_pre = self.rP_breve[-1]
  1192. elif self.bw_connection_mode == 'layered':
  1193. if self.pyr_hi_pass:
  1194. r_pre = self.rP_breve_HI[i+1]
  1195. else:
  1196. r_pre = self.rP_breve[i+1]
  1197. if self.model == "PAL":
  1198. self.dBPP[i] = self.dt * self.eta_bw[i] * np.outer(
  1199. self.noise[i], r_pre
  1200. )
  1201. # add regularizer with gate or without
  1202. if self.gate_regularizer:
  1203. self.dBPP[i] -= self.dt * self.alpha[i] * \
  1204. self.eta_bw[i] * self.BPP[i] * d_relu(r_pre)
  1205. else:
  1206. if self.varphi_regularizer:
  1207. self.dBPP[i] -= self.dt * self.alpha[i] * self.eta_bw[i] * self.d_rP[i] @ self.BPP[i] @ self.d_rP[i+1]
  1208. else:
  1209. # vanilla regularizer
  1210. self.dBPP[i] -= self.dt * self.alpha[i] * self.eta_bw[i] * self.BPP[i]
  1211. return self.dBPP
  1212. class errormc_model(base_model):
  1213. """ This class inherits all properties from the base model class and changes it to the error microcircuit """
  1214. def __init__(self, fw_connection_mode, bw_connection_mode, dWPP_use_activation,
  1215. varphi_transfer, dt, dtxi, tauHP, tauLO, Tpres,
  1216. noise_scale, alpha,
  1217. pyr_hi_pass, dWPP_low_pass,
  1218. noise_type, noise_mode,
  1219. model, activation, error_activation, layers,
  1220. uP_init, uI_init, WPP_init, WIP_init, BII_init, BPI_init,
  1221. gl, gden, gbas, gapi, gntgt,
  1222. eta_fw, eta_bw, eta_PI, eta_IP, seed=123, WT_noise=0.0, **kwargs):
  1223. # init base_model with same settings
  1224. super().__init__(bw_connection_mode=bw_connection_mode, dWPP_use_activation=dWPP_use_activation, dt=dt, Tpres=Tpres,
  1225. model=model, activation=activation, layers=layers,
  1226. uP_init=uP_init, uI_init=uI_init,
  1227. WPP_init=WPP_init, WIP_init=WIP_init, BPP_init=WPP_init, BPI_init=BPI_init,
  1228. gl=gl, gden=gden, gbas=gbas, gapi=gapi, gnI=0.0, gntgt=gntgt,
  1229. eta_fw=eta_fw, eta_bw=eta_bw, eta_PI=eta_PI, eta_IP=eta_IP, seed=seed, WT_noise=WT_noise)
  1230. self.rng = np.random.RandomState(seed)
  1231. # forward connection_mode: skip or layered
  1232. self.fw_connection_mode = fw_connection_mode
  1233. # need to provide learning rate in correct shape
  1234. if self.fw_connection_mode == 'skip':
  1235. # check that correct shape is passed (list of lists)
  1236. assert np.array(eta_fw).shape == (len(layers)-1, len(layers)-1), \
  1237. f"Forward connection is 'skip', but eta_fw is not correct shape. " \
  1238. f"Provide eta_fw in the shape of (len(layers)-1, len(layers)-1) ({(len(layers)-1, len(layers)-1)}). "\
  1239. f"Provided eta_fw: {eta_fw}"
  1240. self.eta_fw = np.array(self.eta_fw)
  1241. # self.eta_fw = convert_eta_to_matrix_for_skip(eta_fw, WPP_init, layers)
  1242. elif self.fw_connection_mode == 'layered':
  1243. # check that correct shape is passed (list)
  1244. assert ((np.array(eta_fw).shape[0] == len(layers)-1) and (np.array(eta_fw).ndim == 1)), \
  1245. f"Forward connection is 'layered', but eta_fw is not correct shape. " \
  1246. f"Provide eta_fw in the shape of (len(layers)-1)=={len(layers)-1}. " \
  1247. f"Provided eta_fw: {eta_fw}"
  1248. # check that same number of vectors has been passed
  1249. if len(uP_init) != len(uI_init):
  1250. raise ValueError(f"Error and representation init voltages do not have same number of entries")
  1251. # assert that skip connections are passed correctly
  1252. if len(BII_init) == 1 and len(layers) > 3 and bw_connection_mode == 'layered':
  1253. if model != 'BP':
  1254. raise ValueError(f"BII_init only has one array, but more are required for {len(layers)} layers \
  1255. using bw_connection_mode='layered'. Did you mean bw_connection_mode='skip'?")
  1256. else:
  1257. self.set_weights(BII=WPP_init[1:])
  1258. logging.info("Setting BII = WPP.T")
  1259. if self.activation[-1] is not linear:
  1260. logging.info("Output layer activation is not linear -- make sure that targets are rates")
  1261. # new connectivity
  1262. self.BII = deepcopy_array(BII_init)
  1263. self.dBII = [np.zeros(shape=BII.shape) for BII in self.BII]
  1264. self.d_activation = [dict_d_activation[activation.__name__] for activation in self.activation]
  1265. # time constants are different for this model, as errors are given to error units
  1266. self.taueffP, self.taueffI_notgt, self.taueffI = self.calc_taueff()
  1267. # the output neurons also have a vapi
  1268. self.vapi = [np.zeros_like(uP) for uP in self.uP]
  1269. self.error_layers = [0] + [len(v) for v in self.uI]
  1270. # if a list of error activations has been passed, use it
  1271. if isinstance(error_activation, list):
  1272. self.error_activation = error_activation
  1273. # else, set same activation for all layers
  1274. else:
  1275. self.error_activation = [error_activation for layer in layers[1:]]
  1276. # whether the target provided should be a rate or voltage
  1277. self.rate_target = True
  1278. # calculate rate of lookahead: phi(ubreve)
  1279. logging.info("Defining rI_breve")
  1280. self.rI_breve = [self.error_activation[i](self.uI_breve[i])
  1281. for i in range(len(self.uI_breve))]
  1282. # derivative of rep units
  1283. self.d_rP_breve = [np.zeros_like(rP_breve) for rP_breve in self.rP_breve]
  1284. # whether to transfer varphi from rep units to error units
  1285. self.varphi_transfer = varphi_transfer
  1286. if varphi_transfer and dWPP_use_activation:
  1287. logging.info("varphi_transfer AND dWPP_use_activation both enabled. Do you want this?")
  1288. # type of noise (OU or white)
  1289. self.noise_type = noise_type
  1290. # mode of noise injection (order vapi or uP or uP_adative)
  1291. self.noise_mode = noise_mode
  1292. # for uP_adaptive, we need epsilon: measures angle between BPP, WPP.T
  1293. if noise_type == 'OU':
  1294. self.tauxi = kwargs.get('tauxi')
  1295. # disabling of prospectivity for rate calcuation (uP_breve will still be used in dWPP!)
  1296. if "rep_lookahead" in kwargs:
  1297. self.rep_lookahead = kwargs.get('rep_lookahead')
  1298. if "error_lookahead" in kwargs:
  1299. self.error_lookahead = kwargs.get('error_lookahead')
  1300. # whether to high-pass filter rPbreve for updates of BPP
  1301. self.pyr_hi_pass = pyr_hi_pass
  1302. # whether to low-pass filter updates of WPP
  1303. self.dWPP_low_pass = dWPP_low_pass
  1304. # whether to use phi' B phi' as regularizer
  1305. self.varphi_regularizer = kwargs.get('varphi_regularizer', False)
  1306. if self.varphi_regularizer:
  1307. self.d_rP = [np.diag(np.zeros(shape=uP.shape)) for uP in self.uP]
  1308. # noise time scale
  1309. self.dtxi = dtxi
  1310. # decimals of dt
  1311. self.dt_decimals = int(np.round(-np.log10(self.dt)))
  1312. # synaptic time constant (sets the low-pass filter of interneuron)
  1313. self.tauHP = tauHP
  1314. self.tauLO = tauLO
  1315. # gaussian noise properties
  1316. self.noise_scale = noise_scale
  1317. self.noise = [np.zeros(shape=uP.shape) for uP in self.uP]
  1318. # self.noise_breve = [np.zeros(shape=uP.shape) for uP in self.uP]
  1319. # we need a new variable: vapi after noise has been added
  1320. # i.e. vapi = BPP rP + BPI rI (as usual), and vapi_noise = vapi + noise
  1321. self.vapi_noise = deepcopy_array(self.vapi)
  1322. # init a counter for time steps after which to resample noise
  1323. self.noise_counter = 0
  1324. self.noise_total_counts = np.round(self.dtxi / self.dt, decimals=self.dt_decimals)
  1325. # init a high-pass filtered version of rP_breve
  1326. self.rP_breve_HI = deepcopy_array(self.rP_breve)
  1327. # init a low-pass filtered version of dWPP
  1328. self.dWPP_LO = [np.zeros(shape=WPP.shape) for WPP in self.WPP]
  1329. self.dWPP_post_LO_old = [np.zeros(shape=rP_breve.shape) for rP_breve in self.rP_breve]
  1330. self.r0_LO_old = np.zeros(shape=layers[0])
  1331. self.r_LO_old = [np.zeros(shape=rP_breve.shape) for rP_breve in self.rP_breve]
  1332. # regularizer for backward weights
  1333. self.alpha = alpha
  1334. # determine whether lateral weights are identity
  1335. self.lateral_is_identity = True
  1336. if not self.error_layers[1:] == self.layers[1:]:
  1337. self.lateral_is_identity = False
  1338. for i in range(len(self.WIP)):
  1339. if not np.all(self.WIP[i] == np.eye(N=self.WIP[i].shape[0], M=self.WIP[i].shape[1])):
  1340. self.lateral_is_identity = False
  1341. if not np.all(self.BPI[i] == np.eye(N=self.BPI[i].shape[0], M=self.BPI[i].shape[1])):
  1342. self.lateral_is_identity = False
  1343. # noise level in setting transpose weights
  1344. self.WT_noise = WT_noise
  1345. # set transpose weights for BP
  1346. if self.model == "BP":
  1347. if self.bw_connection_mode == 'layered':
  1348. # determine BII = WPP.T, taking unequal dims into account
  1349. for i in range(0, len(self.BPI)-1):
  1350. WIP_eye = np.eye(N=self.WIP[i+1].shape[0], M=self.WIP[i+1].shape[1])
  1351. BPI_eye = np.eye(N=self.BPI[i].shape[0], M=self.BPI[i].shape[1])
  1352. # perfect transpose
  1353. self.BII[i] = (WIP_eye @ self.WPP[i+1] @ BPI_eye).T
  1354. # noise matrix (calculated once)
  1355. self.BII_noise = [self.rng.uniform(-self.WT_noise, self.WT_noise, size=BII.shape) for BII in self.BII]
  1356. # add noise
  1357. self.BII[i] += self.BII_noise[i]
  1358. elif self.bw_connection_mode == 'skip':
  1359. # determine BII = WPP.T, taking unequal dims into account
  1360. WIP_eye = [np.eye(N=WIP.shape[0], M=WIP.shape[1]) for WIP in self.WIP]
  1361. BPI_eye = [np.eye(N=BPI.shape[0], M=BPI.shape[1]) for BPI in self.BPI]
  1362. WIP_eye = diag_mat(WIP_eye)
  1363. BPI_eye = diag_mat(BPI_eye)
  1364. # perfect transpose
  1365. self.BII[0] = (WIP_eye @ self.WPP[0][self.layers[0]:,self.layers[0]:] @ BPI_eye).T
  1366. # noise matrix (calculated once)
  1367. self.BII_noise = [self.rng.uniform(-self.WT_noise, self.WT_noise, size=BII.shape) for BII in self.BII]
  1368. # add noise
  1369. self.BII[0] += self.BII_noise[0]
  1370. self.BII[0] = np.triu(self.BII[0])
  1371. def calc_taueff(self):
  1372. # calculate tau_eff for pyramidals and error units
  1373. # taueffP is one value per layer
  1374. taueffP = []
  1375. for i in self.uP:
  1376. taueffP.append(1 / (self.gl + self.gbas + self.gapi))
  1377. taueffI = []
  1378. for i in self.uI:
  1379. taueffI.append(1 / (self.gl + self.gden))
  1380. taueffI[-1] = 1 / (self.gl + self.gntgt)
  1381. # tau_eff for output layer error units in absence of target
  1382. taueffI_notgt = [1 / (self.gl + self.gntgt)]
  1383. return taueffP, taueffI_notgt, taueffI
  1384. def evolve_system(self, r0=None, u_tgt=None, learn_weights=True, learn_lat_weights=True, learn_bw_weights=False, \
  1385. record=True, testing=False, validation=False, compare_dWPP=False):
  1386. """
  1387. This overwrites the vanilla system evolution and implements
  1388. additional noise
  1389. """
  1390. # calculate voltage evolution, including low pass on interneuron synapses
  1391. # see calc_dendritic updates below
  1392. if testing or validation:
  1393. self.duP, self.duI = self.evolve_voltages(r0, u_tgt=None, inject_noise=learn_bw_weights) # includes recalc of rP_breve
  1394. else:
  1395. self.duP, self.duI = self.evolve_voltages(r0, u_tgt, inject_noise=learn_bw_weights) # includes recalc of rP_breve
  1396. if learn_weights or learn_bw_weights or learn_lat_weights or compare_dWPP:
  1397. self.dWPP, self.dWIP, _, self.dBPI = self.evolve_synapses(r0)
  1398. # apply evolution
  1399. for i in range(len(self.duP)):
  1400. self.uP[i] += self.duP[i]
  1401. for i in range(len(self.duI)):
  1402. self.uI [i]+= self.duI[i]
  1403. if learn_weights:
  1404. if self.dWPP_low_pass:
  1405. # calculate lo-passed update of WPP
  1406. self.dWPP_LO = self.calc_dWPP_LO()
  1407. for i in range(len(self.dWPP_LO)):
  1408. self.WPP[i] += self.dWPP_LO[i]
  1409. else:
  1410. for i in range(len(self.dWPP)):
  1411. self.WPP[i] += self.dWPP[i]
  1412. if learn_lat_weights:
  1413. for i in range(len(self.dWIP)):
  1414. self.WIP[i] += self.dWIP[i]
  1415. for i in range(len(self.dBPI)):
  1416. self.BPI[i] += self.dBPI[i]
  1417. # record step
  1418. if hasattr(self, 'rec_per_steps') and record:
  1419. self.rec_counter += 1
  1420. if self.rec_counter % self.rec_per_steps == 0:
  1421. self.rec_counter = 0
  1422. self.record_step(target=u_tgt)
  1423. # during testing and validation, we record MSE for all steps
  1424. if (testing or validation) and record:
  1425. self.record_step(target=u_tgt, MSE_only=True, testing=testing, validation=validation)
  1426. # increase timer
  1427. self.Time = np.round(self.Time + self.dt, decimals=self.dt_decimals)
  1428. # calculate hi-passed rP_breve for synapse BPP
  1429. # self.rP_breve_HI = self.calc_rP_breve_HI()
  1430. def evolve_voltages(self, r0=None, u_tgt=None, inject_noise=False):
  1431. """
  1432. Overwrites voltage evolution:
  1433. Evolves the pyramidal and interneuron voltages by one dt
  1434. using r0 as input rates
  1435. >> Injects noise into vapi
  1436. """
  1437. self.duP = [np.zeros(shape=uP.shape) for uP in self.uP]
  1438. self.duI = [np.zeros(shape=uI.shape) for uI in self.uI]
  1439. # same for dendritic voltages and rates
  1440. self.rP_breve_old = deepcopy_array(self.rP_breve)
  1441. self.rI_breve_old = deepcopy_array(self.rI_breve)
  1442. if self.r0 is not None:
  1443. self.r0_old = self.r0.copy()
  1444. self.vbas_old = deepcopy_array(self.vbas)
  1445. self.vden_old = deepcopy_array(self.vden)
  1446. self.vapi_old = deepcopy_array(self.vapi)
  1447. # calculate lookahead
  1448. self.uP_breve = [self.prospective_voltage(self.uP[i], self.uP_old[i], self.taueffP[i]) for i in range(len(self.uP))]
  1449. self.uI_breve = [self.prospective_voltage(self.uI[i], self.uI_old[i], self.taueffI[i]) for i in range(len(self.uI))]
  1450. if u_tgt is None:
  1451. self.uI_breve[-1] = self.prospective_voltage(self.uI[-1], self.uI_old[-1], self.taueffI_notgt[-1])
  1452. # calculate rate of lookahead: phi(ubreve)
  1453. if hasattr(self,"rep_lookahead"):
  1454. self.rP_breve = [self.activation[i](self.uP_breve[i]) if self.rep_lookahead[i]
  1455. else self.activation[i](self.uP[i])
  1456. for i in range(len(self.uP_breve))]
  1457. else:
  1458. self.rP_breve = [self.activation[i](self.uP_breve[i]) for i in range(len(self.uP_breve))]
  1459. self.d_rP_breve = [self.d_activation[i](self.uP_breve[i]) for i in range(len(self.uP_breve))]
  1460. if hasattr(self,"error_lookahead"):
  1461. self.rI_breve = [self.error_activation[i](self.uI_breve[i]) if self.error_lookahead[i]
  1462. else self.error_activation[i](self.uI[i])
  1463. for i in range(len(self.uI_breve))]
  1464. else:
  1465. self.rI_breve = [self.error_activation[i](self.uI_breve[i]) for i in range(len(self.uI_breve))]
  1466. self.r0 = r0
  1467. # before modifying uP and uI, we need to save copies
  1468. # for future calculation of u_breve
  1469. self.uP_old = deepcopy_array(self.uP)
  1470. self.uI_old = deepcopy_array(self.uI)
  1471. self.vbas, self.vapi, self.vapi_noise, self.vden = self.calc_dendritic_updates(r0, u_tgt)
  1472. # inject noise into newly calculated vapi before calculating update du
  1473. if inject_noise:
  1474. for i in range(len(self.BII)):
  1475. self.generate_noise(layer=i, noise_type=self.noise_type, noise_scale=self.noise_scale)
  1476. self.duP, self.duI = self.calc_somatic_updates(u_tgt, inject_noise=inject_noise)
  1477. return self.duP, self.duI
  1478. def calc_somatic_updates(self, u_tgt=None, inject_noise=False):
  1479. """
  1480. this overwrites the somatic update rules
  1481. difference to super: uses vapi_noise instead of vapi
  1482. and changes circuit to error mc
  1483. """
  1484. # update somatic potentials
  1485. for i in range(len(self.uI)-1):
  1486. ueffI = self.taueffI[i] * (self.gden * self.vden[i])
  1487. delta_uI = (ueffI - self.uI[i]) / self.taueffI[i]
  1488. self.duI[i] = self.dt * delta_uI
  1489. if u_tgt is not None:
  1490. ueffI = self.taueffI[-1] * (self.gntgt * self.WIP[-1] @ np.diag(self.d_rP_breve[-1]) @ (u_tgt[-1] - self.rP_breve[-1]))
  1491. delta_uI = (ueffI - self.uI[-1]) / self.taueffI[-1]
  1492. else:
  1493. ueffI = 0.0
  1494. delta_uI = (ueffI - self.uI[-1]) / self.taueffI_notgt[-1]
  1495. self.duI[-1] = self.dt * delta_uI
  1496. for i in range(0, len(self.layers)-1):
  1497. if inject_noise:
  1498. ueffP = self.taueffP[i] * (self.gbas * self.vbas[i] + self.gapi * self.vapi_noise[i])
  1499. else:
  1500. ueffP = self.taueffP[i] * (self.gbas * self.vbas[i] + self.gapi * self.vapi[i])
  1501. delta_uP = (ueffP - self.uP[i]) / self.taueffP[i]
  1502. self.duP[i] = self.dt * delta_uP
  1503. return self.duP, self.duI
  1504. def calc_vbas(self, WPP_rP_breve_below):
  1505. """
  1506. returns basal voltages in pyramidals of a given layer
  1507. WPP_rP_breve_below: activity from rep unit below times weight to current layer
  1508. """
  1509. return WPP_rP_breve_below
  1510. def calc_vapi(self, rIvec, BPI_mat):
  1511. """
  1512. returns apical voltages in pyramidals of a given layer for error mc
  1513. rIvec: vector of rates from interneuron voltages in output layer
  1514. BPI_mat: matrix connecting interneurons to pyramidals
  1515. """
  1516. return BPI_mat @ rIvec
  1517. def calc_vden(self, deriv_rPvec, WIP_mat, BII_rI_breve_above):
  1518. """
  1519. returns dendritic voltages in error units
  1520. deriv_rPvec: vector of rates from pyramidal voltages in
  1521. layer below
  1522. WIP_mat: matrix connecting pyramidal to pyramidal
  1523. BII_rI_breve_above: activity from error unit above times weight to current layer
  1524. """
  1525. if self.varphi_transfer:
  1526. vden = np.diag(WIP_mat @ deriv_rPvec) @ BII_rI_breve_above
  1527. else:
  1528. vden = BII_rI_breve_above
  1529. return vden
  1530. def generate_noise(self, layer, noise_type, noise_scale):
  1531. """
  1532. this function generates noise for a given layer
  1533. the noise is added to the apical potential in function calculate_dendritic_updates
  1534. there are two noise types:
  1535. - OU: ornstein-uhlenbeck noise, i.e. low-pass filtered white noise updated at every dtxi
  1536. """
  1537. # if dtxi timesteps have passed, sample new noise
  1538. if np.all(self.noise[layer] == 0) or self.noise_counter % self.noise_total_counts == 0:
  1539. if self.noise_type == 'OU':
  1540. self.noise[layer] += 1 / (self.tauxi) * (- self.dt * self.noise[layer] + np.sqrt(self.dt * self.tauxi) * white_noise)
  1541. self.noise_counter = 0
  1542. self.noise_counter += 1
  1543. def calc_dWPP_LO(self):
  1544. # updates the low-passed update of WPP
  1545. # Low-pass has the form d v_out = dt/tau (v_in - v_out)
  1546. for i in range(len(self.dWPP_LO)):
  1547. self.dWPP_LO[i] += self.dt / self.tauLO * (self.dWPP[i] - self.dWPP_LO[i])
  1548. return self.dWPP_LO
  1549. def calc_dendritic_updates(self, r0=None, u_tgt=None):
  1550. """
  1551. this overwrites the dendritic updates by using the error mc dynamics
  1552. and adds vapi_noise somatic voltage after noise injection
  1553. """
  1554. if self.fw_connection_mode == 'skip':
  1555. # if skip fw connection matrix has been given, first calculate total representation vector
  1556. # and multiply with total WPP matrix
  1557. r0 = np.zeros(self.layers[0]) if r0 is None else np.array(r0)
  1558. # concatenate input with hidden layer rates and multiply with forward weights
  1559. total_WPP_rP_breve_below = self.WPP[0] @ np.concatenate([r0] + self.rP_breve)
  1560. # the first entries are weight connecting to input, and do not correspond to neurons
  1561. total_WPP_rP_breve_below = total_WPP_rP_breve_below[self.layers[0]:]
  1562. pos_idx = 0
  1563. for i, curr_rP_breve in enumerate(self.rP_breve):
  1564. # select correct lines from total forward projected representation units
  1565. WPP_rP_breve_below = total_WPP_rP_breve_below[pos_idx:pos_idx+len(curr_rP_breve)]
  1566. self.vbas[i] = self.calc_vbas(WPP_rP_breve_below)
  1567. pos_idx += len(curr_rP_breve)
  1568. elif self.fw_connection_mode == 'layered':
  1569. if r0 is not None:
  1570. self.vbas[0] = self.WPP[0] @ self.r0
  1571. for i in range(1, len(self.layers)-1):
  1572. self.vbas[i] = self.calc_vbas(self.WPP[i] @ self.rP_breve[i-1])
  1573. # if skip bw connection matrix has been given, first calculate total error vector
  1574. # and multiply with total BII matrix
  1575. if self.bw_connection_mode == 'skip':
  1576. assert np.all(self.BII[0] == np.triu(self.BII[0])), "BII has none-triu entries"
  1577. total_BII_rI_breve_above = self.BII[0] @ np.concatenate(self.rI_breve)
  1578. pos_idx = 0
  1579. for i, curr_rI_breve in enumerate(self.rI_breve[:-1]):
  1580. if self.bw_connection_mode == 'skip':
  1581. # select correct lines from total backprojected error vector
  1582. BII_rI_breve_above = total_BII_rI_breve_above[pos_idx:pos_idx+len(curr_rI_breve)]
  1583. self.vden[i] = self.calc_vden(self.d_rP_breve[i],
  1584. self.WIP[i], BII_rI_breve_above)
  1585. pos_idx += len(curr_rI_breve)
  1586. elif self.bw_connection_mode == 'layered':
  1587. self.vden[i] = self.calc_vden(self.d_rP_breve[i],
  1588. self.WIP[i], self.BII[i] @ self.rI_breve[i+1])
  1589. for i in range(0, len(self.layers)-1):
  1590. self.vapi[i] = self.calc_vapi(self.rI_breve[i], self.BPI[i])
  1591. self.vapi_noise[i] = self.vapi[i] + self.noise[i]
  1592. return self.vbas, self.vapi, self.vapi_noise, self.vden
  1593. def evolve_synapses(self, r0):
  1594. """
  1595. evolves all synapses by a dt
  1596. plasticity of WPP
  1597. """
  1598. self.dWPP = [np.zeros(shape=WPP.shape) for WPP in self.WPP]
  1599. self.dWIP = [np.zeros(shape=WIP.shape) for WIP in self.WIP]
  1600. self.dBII = [np.zeros(shape=BII.shape) for BII in self.BII]
  1601. self.dBPI = [np.zeros(shape=BPI.shape) for BPI in self.BPI]
  1602. if self.fw_connection_mode == 'skip':
  1603. # if skip fw connection matrix has been given, first calculate total representation vector
  1604. # and multiply with total WPP matrix
  1605. r0 = np.zeros(self.layers[0]) if self.r0 is None else self.r0
  1606. r0_old = np.zeros(self.layers[0]) if self.r0_old is None else self.r0_old
  1607. # take all pre-synaptic rates (from previous dt)
  1608. rates_old_pre = [r0_old] + self.rP_breve_old
  1609. # take all postsynaptic rates (r0 will be skipped, just for ease of computation)
  1610. if self.dWPP_use_activation:
  1611. rates_post = [r0] + self.rP_breve
  1612. else:
  1613. rates_post = [r0] + self.uP_breve
  1614. # for all post-synaptic rates
  1615. post_pos_idx = len(r0)
  1616. for i, rate_post in enumerate(rates_post):
  1617. if i == 0:
  1618. # skip r0
  1619. continue
  1620. pre_pos_idx = 0
  1621. # for all pre-synaptic rates until the layer below the current one
  1622. for j, rate_old_pre in enumerate(rates_old_pre[:i]):
  1623. if self.dWPP_use_activation:
  1624. post_diff = rate_post - self.activation[i-1](self.gbas / \
  1625. (self.gl + self.gbas + self.gapi) * \
  1626. self.vbas_old[i-1])
  1627. else:
  1628. post_diff = rate_post - (self.gbas / \
  1629. (self.gl + self.gbas + self.gapi) * \
  1630. self.vbas_old[i-1])
  1631. tmp_dWPP = self.dt * self.eta_fw[i-1,j] * np.outer(post_diff, rate_old_pre)
  1632. # tmp_dWPP = np.outer(post_diff, rate_old_pre)
  1633. self.dWPP[0][post_pos_idx:post_pos_idx+len(rate_post),
  1634. pre_pos_idx:pre_pos_idx+len(rate_old_pre)] = tmp_dWPP
  1635. pre_pos_idx += len(rate_old_pre)
  1636. post_pos_idx += len(rate_post)
  1637. elif self.fw_connection_mode == 'layered':
  1638. if self.dWPP_use_activation:
  1639. # input layer
  1640. if r0 is not None:
  1641. # logging.info("updating WPP0")
  1642. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  1643. self.rP_breve[0] - \
  1644. self.activation[0](self.gbas / \
  1645. (self.gl + self.gbas + self.gapi) * \
  1646. self.vbas_old[0]),
  1647. self.r0_old)
  1648. # hidden layers
  1649. for i in range(1, len(self.WPP)-1):
  1650. # logging.info(f"updating WPP{i}")
  1651. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(
  1652. self.rP_breve[i] - \
  1653. self.activation[i](self.gbas / \
  1654. (self.gl + self.gbas + self.gapi) * \
  1655. self.vbas_old[i]),
  1656. self.rP_breve_old[i-1])
  1657. # output layer
  1658. # logging.info("updating WPP-1")
  1659. if len(self.layers) > 2:
  1660. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(
  1661. self.rP_breve[-1] - self.activation[-1](self.gbas / \
  1662. (self.gl + self.gbas + self.gapi) * \
  1663. self.vbas_old[-1]),
  1664. self.rP_breve_old[-2])
  1665. else:
  1666. # input layer
  1667. if r0 is not None:
  1668. # logging.info("updating WPP0")
  1669. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  1670. self.uP_breve[0] - (self.gbas / (self.gl + self.gbas + self.gapi) * self.vbas_old[0]),
  1671. self.r0_old)
  1672. # hidden layers
  1673. for i in range(1, len(self.WPP)-1):
  1674. # logging.info(f"updating WPP{i}")
  1675. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(
  1676. self.uP_breve[i] - (self.gbas / \
  1677. (self.gl + self.gbas + self.gapi) * \
  1678. self.vbas_old[i]),
  1679. self.rP_breve_old[i-1])
  1680. # output layer
  1681. if len(self.layers) > 2:
  1682. # logging.info("updating WPP-1")
  1683. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(
  1684. self.uP_breve[-1] - (self.gbas / \
  1685. (self.gl + self.gbas + self.gapi) \
  1686. * self.vbas_old[-1]),
  1687. self.rP_breve_old[-2])
  1688. """
  1689. plasticity of WIP: weight from rep unit to error unit
  1690. """
  1691. for i in range(len(self.WIP)):
  1692. if self.eta_IP[i] != 0.0:
  1693. self.dWIP[i] = - self.dt * self.eta_IP[i] * (self.WIP[i] - np.eye(N=self.WIP[i].shape[0], M=self.WIP[i].shape[1]))
  1694. """
  1695. plasticity of BPI: weight from error unit to rep unit
  1696. """
  1697. for i in range(0, len(self.BPI)):
  1698. if self.eta_PI[i] != 0.0:
  1699. self.dBPI[i] = - self.dt * self.eta_PI[i] * (self.BPI[i] - np.eye(N=self.BPI[i].shape[0], M=self.BPI[i].shape[1]))
  1700. """
  1701. plasticity of BII
  1702. """
  1703. if self.model == 'FA':
  1704. # do nothing
  1705. pass
  1706. # set transpose weights for BP
  1707. if self.model == "BP":
  1708. if self.bw_connection_mode == 'layered':
  1709. # determine BII = WPP.T, taking unequal dims into account
  1710. for i in range(0, len(self.BPI)-1):
  1711. WIP_eye = np.eye(N=self.WIP[i+1].shape[0], M=self.WIP[i+1].shape[1])
  1712. BPI_eye = np.eye(N=self.BPI[i].shape[0], M=self.BPI[i].shape[1])
  1713. # perfect transpose
  1714. self.BII[i] = (WIP_eye @ self.WPP[i+1] @ BPI_eye).T
  1715. # add noise
  1716. self.BII[i] += self.BII_noise[i]
  1717. elif self.bw_connection_mode == 'skip':
  1718. # determine BII = WPP.T, taking unequal dims into account
  1719. WIP_eye = [np.eye(N=WIP.shape[0], M=WIP.shape[1]) for WIP in self.WIP]
  1720. BPI_eye = [np.eye(N=BPI.shape[0], M=BPI.shape[1]) for BPI in self.BPI]
  1721. WIP_eye = diag_mat(WIP_eye)
  1722. BPI_eye = diag_mat(BPI_eye)
  1723. # perfect transpose
  1724. self.BII[0] = (WIP_eye @ self.WPP[0][self.layers[0]:,self.layers[0]:] @ BPI_eye).T
  1725. # add noise
  1726. self.BII[0] += self.BII_noise[0]
  1727. self.BII[0] = np.triu(self.BII[0])
  1728. return self.dWPP, self.dWIP, self.dBII, self.dBPI
  1729. def get_weights(self):
  1730. return self.WPP, self.WIP, self.BII, self.BPI
  1731. def set_weights(self, model=None, WPP=None, WIP=None, BPP=None,
  1732. BPI=None, BII=None):
  1733. # if another model is given, copy its weights
  1734. if hasattr(model, '__dict__'):
  1735. WPP, WIP, BII, BPI = model.get_weights()
  1736. logging.info(f"Copying weights from model {model}")
  1737. if WPP is not None: self.WPP = deepcopy_array(WPP)
  1738. if WIP is not None: self.WIP = deepcopy_array(WIP)
  1739. if BPP is not None: self.BPP = deepcopy_array(BPP)
  1740. if BPI is not None: self.BPI = deepcopy_array(BPI)
  1741. if BII is not None: self.BII = deepcopy_array(BII)
  1742. class ann_model(base_model):
  1743. """
  1744. This class inherits all properties from the base model class and changes it to an ann
  1745. Here, uP is equal to uP_breve, and only determined by basal input (gapi = 0)
  1746. uI represents the error and is directly backpropagated
  1747. """
  1748. def __init__(self, dt, Tpres,
  1749. model, activation, layers,
  1750. uP_init, uI_init, WPP_init, BPP_init,
  1751. gl, gbas,
  1752. eta_fw, seed=123, **kwargs):
  1753. # init base_model with same settings
  1754. super().__init__(bw_connection_mode='layered', dWPP_use_activation=False, dt=dt, Tpres=Tpres,
  1755. model=model, activation=activation, layers=layers,
  1756. uP_init=uP_init, uI_init=uI_init,
  1757. # use WPP_init as placeholder for other weights
  1758. WPP_init=WPP_init, WIP_init=WPP_init, BPP_init=BPP_init, BPI_init=BPP_init,
  1759. gl=gl, gden=0.0, gbas=gbas, gapi=0.0, gnI=0.0, gntgt=0.0,
  1760. eta_fw=eta_fw, eta_bw=None, eta_PI=None, eta_IP=None, seed=seed)
  1761. self.rng = np.random.RandomState(seed)
  1762. # whether the target provided should be a rate or voltage
  1763. self.rate_target = True
  1764. if self.activation[-1] is not linear:
  1765. logging.info("Output layer activation is not linear -- make sure that targets are rates")
  1766. self.d_activation = [dict_d_activation[activation.__name__] for activation in self.activation]
  1767. # derivative of rep units
  1768. self.d_rP_breve = [np.zeros_like(rP_breve) for rP_breve in self.rP_breve]
  1769. def prospective_voltage(self, uvec, uvec_old, tau, dt=None):
  1770. """
  1771. for this class, prospective voltage is the same as instantaneous voltage
  1772. """
  1773. return uvec
  1774. def evolve_voltages(self, r0=None, u_tgt=None):
  1775. """
  1776. Evolves the pyramidal and interneuron voltages by one dt
  1777. using r0 as input rates
  1778. """
  1779. self.duP = [np.zeros(shape=uP.shape) for uP in self.uP]
  1780. self.duI = [np.zeros(shape=uI.shape) for uI in self.uI]
  1781. # same for dendritic voltages and rates
  1782. self.rP_breve_old = deepcopy_array(self.rP_breve)
  1783. if self.r0 is not None:
  1784. self.r0_old = self.r0.copy()
  1785. self.vbas_old = deepcopy_array(self.vbas)
  1786. # calculate lookahead
  1787. self.uP_breve = [self.prospective_voltage(self.uP[i],
  1788. self.uP_old[i],
  1789. self.taueffP[i]) for i in range(len(self.uP))]
  1790. self.uI_breve = [self.prospective_voltage(self.uI[i],
  1791. self.uI_old[i],
  1792. self.taueffI[i]) for i in range(len(self.uI))]
  1793. # calculate rate of lookahead: phi(ubreve)
  1794. self.rP_breve = [self.activation[i](self.uP_breve[i]) for i in range(len(self.uP_breve))]
  1795. self.d_rP_breve = [self.d_activation[i](self.uP_breve[i]) for i in range(len(self.uP_breve))]
  1796. self.r0 = r0
  1797. # before modifying uP and uI, we need to save copies
  1798. # for future calculation of u_breve
  1799. self.uP_old = deepcopy_array(self.uP)
  1800. self.uI_old = deepcopy_array(self.uI)
  1801. self.duP, self.duI = self.calc_somatic_updates(u_tgt)
  1802. return self.duP, self.duI
  1803. def calc_somatic_updates(self, u_tgt=None):
  1804. """
  1805. calculates somatic updates from dendritic potentials
  1806. for ann: this is just bottom-up input for uP
  1807. whereas errors are represented by uI
  1808. """
  1809. # update errors
  1810. for i in range(len(self.uI)-1):
  1811. self.duI[i] = np.diag(self.d_rP_breve[i]) @ self.BPP[i] @ self.uI[i+1] - self.uI[i]
  1812. if u_tgt is not None:
  1813. self.d_rP_breve[-1] = self.d_activation[-1](self.uP_breve[-1])
  1814. out_error = np.diag(self.d_rP_breve[-1]) @ (u_tgt[-1] - self.rP_breve[-1])
  1815. else:
  1816. out_error = 0.0
  1817. self.duI[-1] = out_error - self.uI[-1]
  1818. # rep units
  1819. self.duP[0] = self.gbas / (self.gbas + self.gl) * self.WPP[0] @ self.r0 - self.uP[0]
  1820. for i in range(1, len(self.layers)-1):
  1821. self.duP[i] = self.gbas / (self.gbas + self.gl) * self.WPP[i] @ self.rP_breve[i-1] - self.uP[i]
  1822. return self.duP, self.duI
  1823. def evolve_synapses(self, r0):
  1824. """
  1825. evolves all synapses by a dt
  1826. plasticity of WPP
  1827. """
  1828. self.dWPP = [np.zeros(shape=WPP.shape) for WPP in self.WPP]
  1829. self.dWIP = [np.zeros(shape=WIP.shape) for WIP in self.WIP]
  1830. self.dBPP = [np.zeros(shape=BPP.shape) for BPP in self.BPP]
  1831. self.dBPI = [np.zeros(shape=BPI.shape) for BPI in self.BPI]
  1832. if len(self.layers) == 2:
  1833. # input layer
  1834. if r0 is not None:
  1835. # logging.info("updating WPP0")
  1836. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  1837. self.uI_breve[0], self.r0_old)
  1838. else:
  1839. if r0 is not None:
  1840. # logging.info("updating WPP0")
  1841. self.dWPP[0] = self.dt * self.eta_fw[0] * np.outer(
  1842. self.uI_breve[0], self.r0_old)
  1843. # hidden layers
  1844. for i in range(1, len(self.WPP)-1):
  1845. # logging.info(f"updating WPP{i}")
  1846. self.dWPP[i] = self.dt * self.eta_fw[i] * np.outer(
  1847. self.uI_breve[i], self.rP_breve_old[i-1])
  1848. # output layer
  1849. # logging.info("updating WPP-1")
  1850. if len(self.layers) > 2:
  1851. self.dWPP[-1] = self.dt * self.eta_fw[-1] * np.outer(
  1852. self.uI_breve[-1], self.rP_breve_old[-2])
  1853. """
  1854. plasticity of BPP
  1855. """
  1856. if self.model == 'FA':
  1857. # do nothing
  1858. pass
  1859. elif self.model == 'BP':
  1860. self.set_weights(BPP = [WPP.T for WPP in self.WPP[1:]])
  1861. return self.dWPP, self.dWIP, self.dBPP, self.dBPI
  1862. def set_self_predicting_state(self):
  1863. pass
  1864. class dPC_model(base_model):
  1865. """ This class inherits all properties from the base model class and changes it to the dendritic PC microcircuit """
  1866. def __init__(self, bw_connection_mode, dWPP_use_activation,
  1867. dt, Tpres,
  1868. model, activation, layers,
  1869. uP_init, uI_init, WPP_init, WIP_init, BPP_init, BPI_init,
  1870. gl, gden, gbas, gapi, gntgt,
  1871. eta_fw, eta_bw, eta_PI, eta_IP, seed=123, WT_noise=0.0, **kwargs):
  1872. # init base_model with same settings
  1873. super().__init__(bw_connection_mode=bw_connection_mode, dWPP_use_activation=dWPP_use_activation, dt=dt, Tpres=Tpres,
  1874. model=model, activation=activation, layers=layers,
  1875. uP_init=uP_init, uI_init=uI_init,
  1876. WPP_init=WPP_init, WIP_init=WIP_init, BPP_init=WPP_init, BPI_init=BPI_init,
  1877. gl=gl, gden=gden, gbas=gbas, gapi=gapi, gnI=0.0, gntgt=gntgt,
  1878. eta_fw=eta_fw, eta_bw=eta_bw, eta_PI=eta_PI, eta_IP=eta_IP, seed=seed, WT_noise=WT_noise)
  1879. # overwrite uI_breve to be copy of rP_breve in same layer
  1880. self.uI_breve = [rP_breve for rP_breve in self.rP_breve]
  1881. # linear activation on interneurons
  1882. self.rI_breve = self.uI_breve
  1883. # reshape BPI to be square mat of current layer entries
  1884. self.BPI = [BPI @ WIP for WIP, BPI in zip(self.WIP, self.BPI)]
  1885. def evolve_voltages(self, r0=None, u_tgt=None):
  1886. """
  1887. Overwrites voltage evolution:
  1888. Only difference is that rI_breve = uI_breve = - rP_breve
  1889. """
  1890. self.duP = [np.zeros(shape=uP.shape) for uP in self.uP]
  1891. self.duI = [np.zeros(shape=uI.shape) for uI in self.uI]
  1892. # same for dendritic voltages and rates
  1893. self.rP_breve_old = deepcopy_array(self.rP_breve)
  1894. self.rI_breve_old = deepcopy_array(self.rI_breve)
  1895. if self.r0 is not None:
  1896. self.r0_old = self.r0.copy()
  1897. self.vbas_old = deepcopy_array(self.vbas)
  1898. self.vden_old = deepcopy_array(self.vden)
  1899. self.vapi_old = deepcopy_array(self.vapi)
  1900. # calculate lookahead
  1901. self.uP_breve = [self.prospective_voltage(self.uP[i],
  1902. self.uP_old[i],
  1903. self.taueffP[i]) for i in range(len(self.uP))]
  1904. self.uI_breve = [rP_breve for rP_breve in self.rP_breve] # modification here
  1905. # calculate rate of lookahead: phi(ubreve)
  1906. self.rP_breve = [self.activation[i](self.uP_breve[i])
  1907. for i in range(len(self.uP_breve))]
  1908. self.rI_breve = self.uI_breve # modification here
  1909. self.r0 = r0
  1910. # before modifying uP and uI, we need to save copies
  1911. # for future calculation of u_breve
  1912. self.uP_old = deepcopy_array(self.uP)
  1913. self.uI_old = deepcopy_array(self.uI)
  1914. self.vbas, self.vapi, _ = self.calc_dendritic_updates(r0, u_tgt)
  1915. self.duP, _ = self.calc_somatic_updates(u_tgt)
  1916. return self.duP, self.duI
  1917. def evolve_system(self, r0=None, u_tgt=None, learn_weights=True,
  1918. learn_lat_weights=True, learn_bw_weights=False,
  1919. record=True, testing=False, validation=False, compare_dWPP=False):
  1920. """
  1921. evolves the system by one time step:
  1922. updates synaptic weights and voltages given input rate r0
  1923. Disables application of dWIP and duI
  1924. """
  1925. # increase timer by dt and round float to nearest dt
  1926. self.Time = np.round(self.Time + self.dt,
  1927. decimals=int(np.round(-np.log10(self.dt))))
  1928. if testing or validation:
  1929. self.duP, _ = self.evolve_voltages(r0, u_tgt=None) # includes recalc of rP_breve
  1930. else:
  1931. self.duP, _ = self.evolve_voltages(r0, u_tgt) # includes recalc of rP_breve
  1932. if learn_weights or learn_bw_weights or learn_lat_weights or compare_dWPP:
  1933. # base_model evolve_synapse sets BPI = - BPP for BP.
  1934. # we don't want this here, so we buffer and reset after evolve_synapse
  1935. if self.model == 'BP':
  1936. BPI_buffer = deepcopy_array(self.BPI)
  1937. self.dWPP, _, self.dBPP, self.dBPI = self.evolve_synapses(r0, learn_WIP=False)
  1938. if self.model == 'BP':
  1939. self.set_weights(BPI = BPI_buffer)
  1940. # apply evolution
  1941. for i in range(len(self.duP)):
  1942. self.uP[i] += self.duP[i]
  1943. # for i in range(len(self.duI)):
  1944. # self.uI[i] += self.duI[i]
  1945. if learn_weights:
  1946. for i in range(len(self.dWPP)):
  1947. self.WPP[i] += self.dWPP[i]
  1948. if learn_lat_weights:
  1949. # for i in range(len(self.dWIP)):
  1950. # self.WIP[i] += self.dWIP[i]
  1951. for i in range(len(self.dBPI)):
  1952. self.BPI[i] += self.dBPI[i]
  1953. if learn_bw_weights:
  1954. for i in range(len(self.dBPP)):
  1955. self.BPP[i] += self.dBPP[i]
  1956. # logging.warning("SETTING SPS AFTER dW")
  1957. # self.set_self_predicting_state()
  1958. # record step
  1959. if hasattr(self, 'rec_per_steps') and record:
  1960. self.rec_counter += 1
  1961. if self.rec_counter % self.rec_per_steps == 0:
  1962. self.rec_counter = 0
  1963. self.record_step(target=u_tgt)
  1964. # during testing and validation, we record MSE for all steps
  1965. if (testing or validation) and record:
  1966. self.record_step(target=u_tgt, MSE_only=True, testing=testing, validation=validation)
  1967. def set_self_predicting_state(self):
  1968. for i in range(len(self.WIP)-1):
  1969. self.BPI[i] = - self.gbas/(self.gl + self.gbas + self.gapi) * self.BPP[i].copy() @ self.WPP[i+1].copy()
  1970. if len(self.layers) > 2:
  1971. self.BPI[-1] = - self.gbas/(self.gl + self.gbas) * self.BPP[-1].copy() @ self.WPP[-1].copy()

microcircuit.py at commit 7baef7d, no license · at the source

Overview

Authors: Kevin Max1,2, Ismael Jaras2, Arno Granier2, Katharina A. Wilmes2,3, Mihai A. Petrovici2
  1. Neural Computation Unit, Okinawa Institute of Science and Technology, Onna, Japan
  2. Department of Physiology, Bern University, Bern, Switzerland
  3. Institute of Neuroinformatics, University of Zurich and ETH Zurich, Zurich, Switzerland
Journal: PLoS computational biology, volume 22, issue 4, article e1014164
Dates: received 8 September 2025; accepted 25 March 2026; published online 17 April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pcbi.1014164 · PMID 41996333 · PMCID PMC13089762 · OpenAlex W4412970320
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), human (organism)
MeSH: Brain*, Models, Neurological*, Nerve Net*, Neurons*, Action Potentials, Algorithms, Animals, Computational Biology, Computer Simulation, Humans, Learning, Pyramidal Cells, Visual Cortex (* major topic)
Journal subjects: Biology and Life Sciences, Cell Biology, Cellular Types, Animal Cells, Neurons, Neuroscience, Cellular Neuroscience, Neuronal Dendrites, Engineering and Technology, Electrical Engineering, Electrical Circuits, Microcircuits, Computational Biology, Computational Neuroscience, Coding Mechanisms, Cognitive Science, Cognitive Psychology, Learning, Psychology, Social Sciences, Learning and Memory, Computer and Information Sciences, Neural Networks, Anatomy, Brain, Visual Cortex, Medicine and Health Sciences, People and Places, Population Groupings, Professions, Teachers
Topic: Neural dynamics and brain function (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: cited by 1 paper (Europe PMC); 124 references in the paper

Abstract

Neural responses to mismatches between expected and actual stimuli have been widely reported across different species. How does the brain use such error signals for learning? While global error signals can be useful, their ability to learn complex computation at the scale observed in the brain is lacking. In comparison, more local, neuron-specific error signals enable superior performance, but their computation and propagation remain unclear. Motivated by the breakthrough of deep learning, this has inspired the ‘backpropagation and the brain’ hypothesis, i.e., that the brain implements a form of the error backpropagation algorithm. In this work, we introduce a biologically motivated, multi-area cortical microcircuit model, implementing error backpropagation under consideration of recent physiological evidence. We model populations of cortical pyramidal cells acting as representation and error neurons, with bio-plausible local and inter-area connectivity, guided by experimental observations of connectivity of the primate visual cortex. In our model, all information transfer is biologically motivated, inference and learning occur without phases, and network dynamics demonstrably approximate those of error backpropagation. We show the capabilities of our model on a wide range of benchmarks, and compare to other models, such as dendritic hierarchical predictive coding. In particular, our model addresses shortcomings of other theories in terms of scalability to many cortical areas. Finally, we make concrete predictions, which differentiate it from other theories, and which can be tested experimentally.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repositories

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

Zenodo 16909564

License: CC-BY-4.0
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: “Data Availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (14 files), Matplotlib (8 files), SciPy (5 files), PyTorch (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers (HTTP 200)
  • 29 September 2026: the link answers (HTTP 200)
22 files

kma-code/error-neuron-microcircuits

License: none: the authors keep all their rights
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: 7baef7db7acece086d362ff0b71fe29454e693c9, 2 February 2026
Languages: Python (11), Shell (7), Jupyter (3)
Size: 325 files, 21 scripts
Software Heritage: not archived
Found in: the Zenodo archive record
Holds: README, environment (numpy_model/requirements.txt), 3 notebooks
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (14 files), Matplotlib (8 files), SciPy (5 files), PyTorch (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
22 files

The paper's code and data availability statement is in the Data section.

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;
  • 42 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

Code for the network simulations, generating the datasets and producing all figures is accessible under DOI https://doi.org/10.5281/zenodo.16909564.

Reproduced under the paper's license (CC BY), from the paper cited above.

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

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 13 MeSH terms, 5 funders, 98 references.

Cite

This paper

Max, K., Jaras, I., Granier, A., Wilmes, K. A., & Petrovici, M. A. (2026). 'Backpropagation and the brain' realized in cortical error neuron microcircuits. PLoS computational biology, 22(4), e1014164. https://doi.org/10.1371/journal.pcbi.1014164

BibTeX

@article{max2026backpropagation,
author = {Max, Kevin and Jaras, Ismael and Granier, Arno and Wilmes, Katharina A. and Petrovici, Mihai A.},
title = {{'Backpropagation and the brain' realized in cortical error neuron microcircuits}},
journal = {PLoS computational biology},
year = {2026},
month = apr,
volume = {22},
number = {4},
pages = {e1014164},
publisher = {PLOS},
issn = {1553-734X},
doi = {10.1371/journal.pcbi.1014164},
url = {https://doi.org/10.1371/journal.pcbi.1014164},
pmid = {41996333},
pmcid = {PMC13089762}
}

RIS

TY - JOUR
AU - Max, Kevin
AU - Jaras, Ismael
AU - Granier, Arno
AU - Wilmes, Katharina A.
AU - Petrovici, Mihai A.
TI - 'Backpropagation and the brain' realized in cortical error neuron microcircuits
T2 - PLoS computational biology
J2 - PLoS Comput Biol
PY - 2026
DA - 2026/04/17
VL - 22
IS - 4
SP - e1014164
SN - 1553-734X
PB - PLOS
DO - 10.1371/journal.pcbi.1014164
UR - https://doi.org/10.1371/journal.pcbi.1014164
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pcbi.1014164",
"type": "article-journal",
"title": "'Backpropagation and the brain' realized in cortical error neuron microcircuits",
"container-title": "PLoS computational biology",
"author": [
{
"family": "Max",
"given": "Kevin"
},
{
"family": "Jaras",
"given": "Ismael"
},
{
"family": "Granier",
"given": "Arno"
},
{
"family": "Wilmes",
"given": "Katharina A."
},
{
"family": "Petrovici",
"given": "Mihai A."
}
],
"container-title-short": "PLoS Comput Biol",
"volume": "22",
"issue": "4",
"page": "e1014164",
"DOI": "10.1371/journal.pcbi.1014164",
"PMID": "41996333",
"PMCID": "PMC13089762",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pcbi.1014164",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
17
]
]
}
}

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-74460-8 [code]
Spike-based alignment learning solves the weight transport problem.
Journal: Nature communications
In common: PyTorch, SciPy, Matplotlib, 1 other tool, computational modeling (no new data), 17 references, author Kevin Max
[2] doi:10.1038/s41467-026-70347-w [code]
Desegregation of neuronal predictive processing.
Journal: Nature communications
In common: computational modeling (no new data), 17 references
[3] doi:10.1038/s41467-026-70354-x [code]
Global error signal guides local optimization in mismatch calculation.
Journal: Nature communications
In common: SciPy, Matplotlib, NumPy, computational modeling (no new data), 10 references
[4] doi:10.7554/elife.108941 [code]
Visuomotor mismatch EEG responses over occipital cortex of freely moving human subjects.
Journal: eLife
In common: 9 references
[5] doi:10.7554/elife.105968 [code]
Modeling the hallucinatory effects of classical psychedelics in terms of replay-dependent plasticity mechanisms.
Journal: eLife
In common: PyTorch, Matplotlib, NumPy, 7 references
[6] doi:10.7554/elife.105953 [code]
Top-down feedback in deep neural networks leads to functional differences during audiovisual integration.
Journal: eLife
In common: PyTorch, SciPy, Matplotlib, 1 other tool, 6 references
[7] doi:10.1371/journal.pone.0354021 [code]
A genetic algorithm for self-supervised models of oscillatory neurodynamics.
Journal: PloS one
In common: SciPy, Matplotlib, NumPy, computational modeling (no new data), 6 references
[8] doi:10.1126/sciadv.aed6417 [code]
Intrinsic timing, not temporal prediction, underlies ramping dynamics in visual and parietal cortex during passive behavior.
Journal: Science advances
In common: SciPy, Matplotlib, NumPy, 6 references
[9] doi:10.1371/journal.pcbi.1013138 [code]
Hierarchical recurrent temporal prediction as a model of the mammalian dorsal visual pathway.
Journal: PLoS computational biology
In common: PyTorch, SciPy, Matplotlib, 1 other tool, computational modeling (no new data), 5 references
[10] doi:10.21203/rs.3.rs-10384662/v1 [code]
Temporal Gating by Chandelier Cells Encodes Signed Prediction Errors
Journal: Research Square (preprint)
In common: 7 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.