OSCR

Hierarchical recurrent temporal prediction as a model of the mammalian dorsal visual pathway.

Code ↔ Paper

8 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 8 matches
  1. [1] § Materials and methods › Data analysis ↔ virtual_physiology/VirtualNetworkPhysiology.py, lines 490–585 · score 0.95 · partial correlation, plaid angle, plaid component, units responded, individual components, Plaid responses
  2. [2] § Materials and methods › Data analysis ↔ analysis/surround_suppression.ipynb, lines 506–643 · score 0.76 · Surround suppression, preferred orientation, tuning curve, grating stimulus, temporal frequency, mask
  3. [3] § Results › Model units mirror the visual system’s hierarchy of 2D motion sensitivity ↔ virtual_physiology/VirtualNetworkPhysiology.py, lines 490–585 · score 0.70 · plaid angles, plaid response, tuning curves, grating stimuli, Plaid pattern, Fisher
  4. [4] § Materials and methods › Model-brain alignment ↔ neural_fitting/get_neural_data.py, lines 9–29 · score 0.58 · signal power, noise power, Neural
  5. [5] § Results › Model units mirror the visual system’s hierarchy of 2D motion sensitivity ↔ analysis/plaid_motion.ipynb, lines 68–98 · score 0.57 · plaid motion, selective units, plaid pattern, DSI, tuned
  6. [6] § Materials and methods › The hierarchical recurrent temporal prediction model ↔ models/network_hierarchical_recurrent.py, lines 149–163 · score 0.56 · enforce Dale, Law, excitatory, inhibitory, weight, recurrent
  7. [7] § Results › The hierarchical recurrent temporal prediction model matches cortical stimulus representations better than other models ↔ analysis/surround_suppression.ipynb, lines 506–643 · score 0.54 · surround suppression, tuning curve, full models, hierarchical recurrent, feedback, stimulus
  8. [8] § Results › The hierarchical recurrent temporal prediction model matches cortical stimulus representations better than other models ↔ analysis/plaid_motion.ipynb, lines 308–419 · score 0.54 · KS distance, plaid pattern, iv, iii, G3, G2

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 · 709 lines · 30 KB · no license · 2 matches

  1. import pickle, math, os
  2. import numpy as np
  3. import torch
  4. import scipy
  5. import scipy.stats
  6. from scipy import ndimage
  7. import scipy.optimize as opt
  8. from scipy import signal
  9. import scipy.fft as fft
  10. class VirtualPhysiology:
  11. # model trained model
  12. # hyperparameters model hyperparameters
  13. # hidden_unit_range range object of hidden units for analysis
  14. # device device for tensors (cpu or cuda)
  15. def __init__ (self, model, hyperparameters, frame_shape, hidden_units, device):
  16. self.data = []
  17. self.model = model
  18. self.model.eval()
  19. self.hyperparameters = hyperparameters
  20. self.warmup = hyperparameters["warmup"]
  21. self.frame_size = hyperparameters["frame_size"]
  22. self.t_steps = 50
  23. self.frame_shape = frame_shape
  24. self.hidden_units = hidden_units
  25. self.device = device
  26. self.data = []
  27. for group in self.hidden_units:
  28. self.data.append([])
  29. self.osi_thresh = 0.4
  30. self.dsi_thresh = 0.3
  31. self.mean_response_offset = 5
  32. min_n, max_n = min(frame_shape), max(frame_shape)
  33. self.spatial_frequencies = [i/max_n for i in range(1, min_n//2+1)] # Cycles / pixel
  34. self.orientations = np.arange(0, 360, 5) # Degrees
  35. self.temporal_frequencies = np.linspace(1/self.t_steps, 1/4, self.t_steps//4) # Cycles / frame
  36. @classmethod
  37. def load (cls, data_path, model, hyperparameters, frame_shape, hidden_units, device):
  38. virtual_physiology = cls(
  39. model=model,
  40. hyperparameters=hyperparameters,
  41. frame_shape=frame_shape,
  42. hidden_units=hidden_units,
  43. device=device
  44. )
  45. with open(data_path, 'rb') as handler:
  46. virtual_physiology.data = pickle.load(handler)
  47. return virtual_physiology
  48. # relative_path = use file name as part of data_path relative to model dir
  49. def save (self, data_path):
  50. with open(data_path, 'wb') as p:
  51. pickle.dump(self.data, p, protocol=4)
  52. with open(data_path + '.params', 'wb') as p:
  53. params = {
  54. "t_steps": self.t_steps,
  55. "spatial_frequencies": self.spatial_frequencies,
  56. "orientations": self.orientations,
  57. "temporal_frequencies": self.temporal_frequencies
  58. }
  59. pickle.dump(params, p, protocol=4)
  60. def get_unit_data (self, unit_idx):
  61. for group in self.data:
  62. for unit_data in group:
  63. if unit_data["hidden_unit_index"] == unit_idx:
  64. return unit_data
  65. return False
  66. def get_group_from_unit_idx (self, unit_idx):
  67. for group_idx in range(len(self.hidden_units)):
  68. if unit_idx < np.sum(self.hidden_units[:group_idx+1]):
  69. return group_idx
  70. return False
  71. def get_response_weighted_average (self, n_rand_stimuli=100):
  72. # Pre-allocate lists of lists to hold gaussian noise stimuli and associated unit activity
  73. stimuli = []
  74. unit_activity = []
  75. for _ in range (np.sum(self.hidden_units)):
  76. stimuli.append([])
  77. unit_activity.append([])
  78. noise_shape = (n_rand_stimuli, self.warmup+self.t_steps, self.frame_size)
  79. noise = np.random.normal(loc=0, scale=1, size=noise_shape)
  80. noise = torch.Tensor(noise).to(self.device)
  81. with torch.no_grad():
  82. _, hidden_state = self.model(noise)
  83. response = hidden_state.detach().numpy().reshape(n_rand_stimuli*(self.warmup+self.t_steps), -1)
  84. noise = noise.detach().numpy().reshape(-1, self.frame_size)
  85. rwa_arr = []
  86. for unit_idx, unit_responses in enumerate(response.T):
  87. group_idx = self.get_group_from_unit_idx(unit_idx)
  88. if unit_idx % 100 == 0:
  89. print('Processing RWA for unit', unit_idx)
  90. if np.sum(unit_responses):
  91. if group_idx == 0:
  92. rwa = np.average(noise, axis=0, weights=unit_responses)
  93. else:
  94. rwa = np.average(noise[:-group_idx], axis=0, weights=unit_responses[group_idx:])
  95. self.data[group_idx].append({
  96. "hidden_unit_index": unit_idx,
  97. "response_weighted_average": rwa
  98. })
  99. print('Finished averaging stimuli')
  100. return self
  101. # sf = cycles per pixel
  102. # tf = cycles per second
  103. # speed = tf/sf = pixels per second
  104. def get_grating_stimuli(self, spatial_frequency, orientation, temporal_frequency, grating_amplitude, frames):
  105. y_size, x_size = self.frame_shape
  106. theta = (orientation-90) * np.pi/180
  107. x, y = np.meshgrid(np.arange(0, x_size), np.arange(0, y_size))
  108. x_theta = x * np.cos(theta) + y * np.sin(theta)
  109. phase_shift = 2*np.pi*temporal_frequency
  110. phases = np.arange(frames)*phase_shift
  111. grating_frames = []
  112. for phase in phases:
  113. grating_frames.append( grating_amplitude * np.sin(2*spatial_frequency*np.pi*x_theta - phase) )
  114. gratings = np.array(grating_frames).reshape(1, frames, y_size*x_size)
  115. gratings = (gratings-np.mean(gratings))/np.std(gratings) ##
  116. gratings = torch.Tensor(gratings).to(self.device)
  117. return gratings
  118. def get_grating_responses (self):
  119. # Add array to data dictionary structures containing
  120. # complete response (for each grating phase) and mean response (averaged across phases)
  121. # for each spatial frequency/orientation/tf combination
  122. for group_data in self.data:
  123. for unit_data in group_data:
  124. unit_data["grating_responses"] = np.zeros((
  125. len(self.spatial_frequencies),
  126. len(self.orientations),
  127. len(self.temporal_frequencies),
  128. self.t_steps
  129. ))
  130. unit_data["mean_grating_responses"] = np.zeros((
  131. len(self.spatial_frequencies),
  132. len(self.orientations),
  133. len(self.temporal_frequencies)
  134. ))
  135. # Keep track of progress for display purposes
  136. param_count = 0
  137. try:
  138. total_params = self.data[0][0]["mean_grating_responses"].size
  139. except:
  140. total_params = self.data[1][0]["mean_grating_responses"].size
  141. # Loop through each parameter combination for each unit
  142. for sf_idx, sf in enumerate(self.spatial_frequencies):
  143. for ori_idx, ori in enumerate(self.orientations):
  144. for tf_idx, tf in enumerate(self.temporal_frequencies):
  145. # Generate gratings for particular param combination
  146. gratings = self.get_grating_stimuli(sf, ori, tf, grating_amplitude=1, frames=self.warmup+self.t_steps)
  147. # Feedforward pass through network
  148. with torch.no_grad():
  149. _, hidden_state = self.model(gratings)
  150. # Loop through unit responses at each time step
  151. for group_data in self.data:
  152. for unit_data in group_data:
  153. unit_idx = unit_data["hidden_unit_index"]
  154. # Discard warm up period of network's response to gratings
  155. unit_responses = hidden_state[0, self.warmup:, unit_idx].cpu().numpy()
  156. unit_data["grating_responses"][sf_idx, ori_idx, tf_idx] = unit_responses
  157. unit_data["mean_grating_responses"][sf_idx, ori_idx, tf_idx] = np.mean(unit_responses)
  158. if param_count % 100 == 99:
  159. print("Finished param combination {}/{}".format(param_count+1, total_params))
  160. param_count += 1
  161. print("Finished tuning curve")
  162. return self
  163. # Takes orientation tuning curve at max tf and sf
  164. # Returns direction selectivity (DSI)
  165. def get_DSI (self, tuning_curve):
  166. orient_pref_idx = np.where(tuning_curve == np.max(tuning_curve))[0][0]
  167. orient_pref = self.orientations[orient_pref_idx]
  168. orient_pref_resp = tuning_curve[orient_pref_idx]
  169. orient_opp = (orient_pref + 180) % 360
  170. orient_opp_idx = np.where(self.orientations == orient_opp)[0][0]
  171. orient_opp_resp = tuning_curve[orient_opp_idx]
  172. DSI = (orient_pref_resp - orient_opp_resp) / (orient_pref_resp + orient_opp_resp)
  173. #DSI = 1 - (orient_opp_resp/orient_pref_resp)
  174. return DSI
  175. # Takes orientation tuning curve at max tf and sf
  176. # Returns orientation selectivity (OSI)
  177. def get_OSI (self, tuning_curve):
  178. orient_pref_idx = np.where(tuning_curve == np.max(tuning_curve))[0][0]
  179. orient_pref = self.orientations[orient_pref_idx]
  180. orient_pref_resp = tuning_curve[orient_pref_idx]
  181. orient_orth1 = (orient_pref + 90) % 360
  182. orient_orth1_idx = np.where(self.orientations == orient_orth1)[0][0]
  183. orient_orth2 = (orient_pref - 90) % 360
  184. orient_orth2_idx = np.where(self.orientations == orient_orth2)[0][0]
  185. orient_orth_resp = (tuning_curve[orient_orth1_idx]+tuning_curve[orient_orth2_idx]) / 2
  186. OSI = (orient_pref_resp - orient_orth_resp) / (orient_pref_resp + orient_orth_resp)
  187. return OSI
  188. def get_orientation_tuning_curve(self, unit_data):
  189. mean_grating_responses = unit_data["mean_grating_responses"]
  190. # Get indices of max grating response (mean across time steps)
  191. max_mean_grating_response = np.max(mean_grating_responses)
  192. unit_data["max_mean_grating_response"] = max_mean_grating_response
  193. sf_idx, ori_idx, tf_idx = [idx[0] for idx in np.where(mean_grating_responses == max_mean_grating_response)]
  194. orientation_tuning_curve = mean_grating_responses[sf_idx, :, tf_idx]
  195. return orientation_tuning_curve
  196. # Gets OSI, DSI, CV and modulation ratio for each unit
  197. def get_grating_responses_parameters (self):
  198. for group_idx, group_data in enumerate(self.data):
  199. for unit_i, unit_data in enumerate(group_data):
  200. grating_responses = unit_data["grating_responses"]
  201. mean_grating_responses = unit_data["mean_grating_responses"]
  202. # Get indices of max grating response (mean across time steps)
  203. max_mean_grating_response = np.max(mean_grating_responses)
  204. unit_data["max_mean_grating_response"] = max_mean_grating_response
  205. sf_idx, ori_idx, tf_idx = [idx[0] for idx in np.where(mean_grating_responses == max_mean_grating_response)]
  206. # Convert these indices into the underlying parameters
  207. max_sf = unit_data["preferred_sf"] = self.spatial_frequencies[sf_idx]
  208. max_ori = unit_data["preferred_orientation"] = self.orientations[ori_idx]
  209. max_tf = unit_data["preferred_tf"] = self.temporal_frequencies[tf_idx]
  210. # Response to moving grating for parameters that give maximum response
  211. optimum_grating_response = grating_responses[sf_idx, ori_idx, tf_idx]
  212. unit_data["optimum_grating_response"] = optimum_grating_response
  213. # Get CV and DSI measures from curve
  214. # given max spatial frequency and temporal frequency parameters
  215. orientation_curve = mean_grating_responses[sf_idx, :, tf_idx]
  216. unit_data["OSI"] = self.get_OSI(orientation_curve)
  217. unit_data["DSI"] = self.get_DSI(orientation_curve)
  218. if unit_i % 50 == 49:
  219. print("Finished unit {} / {}, group {} / {}".format(
  220. unit_i+1, len(group_data), group_idx+1, len(self.data)
  221. ))
  222. return self
  223. # Filter unresponsive units and units where curve fitting failed
  224. def filter_unit_data (self, group_idx):
  225. group_data = self.data[group_idx]
  226. # Reject low mean response units (< 1% of mean, max mean response)
  227. all_max_mean = []
  228. for group in self.data:
  229. for u in group:
  230. all_max_mean.append(u["max_mean_grating_response"])
  231. mean_responses = [u["max_mean_grating_response"] for u in group_data]
  232. response_threshold = 0.01 * np.mean(all_max_mean)
  233. n_filtered = len(np.where(mean_responses < response_threshold)[0])
  234. print("{} / {} units below response threshold".format(n_filtered, len(group_data)))
  235. # Reject units where curve fitting failed (modulation_ratio set as False)
  236. n_filtered = len([u for u in group_data if not u["modulation_ratio"]])
  237. print("{} / {} units failed to fit curve for modulation ratio estimate".format(n_filtered, len(group_data)))
  238. # Now actually filter those units out
  239. return [u for u in group_data if u["modulation_ratio"] and u["max_mean_grating_response"] >= response_threshold]
  240. # Filter unresponsive units (in place method)
  241. def filter_nonresponding_units (self):
  242. for group_idx, group in enumerate(self.data):
  243. filtered = []
  244. for u in group:
  245. max_mean = np.max(u["mean_grating_responses"])
  246. if max_mean > 0:
  247. filtered.append(u)
  248. self.data[group_idx] = filtered
  249. original_n = self.hidden_units[group_idx]
  250. filtered_n = len(self.data[group_idx])
  251. print(f"{filtered_n} / {original_n} units kept after filtering non-responsive units")
  252. return self
  253. def get_moving_bar_stimuli (self, direction, x, y, bar_amplitude=1, bar_size=5, frames_len = 20):
  254. # Include warmup period for network in stimulus
  255. # Where first n frames will be gray (all 0's)
  256. total_frames = self.warmup + frames_len
  257. stimuli = np.zeros((total_frames, self.frame_shape[0], self.frame_shape[1]))
  258. for i in range(frames_len):
  259. bar_position = 0
  260. square = np.ones((bar_size, bar_size))*bar_amplitude
  261. if direction == 0 or direction == 270:
  262. bar_position = (bar_size - i) % bar_size
  263. else:
  264. bar_position = i%(bar_size)
  265. if direction == 0 or direction == 180:
  266. square[bar_position, :] = -bar_amplitude
  267. else:
  268. square[:, bar_position] = -bar_amplitude
  269. stimulus = np.zeros((self.frame_shape[0], self.frame_shape[1]))
  270. stimulus[y:y+bar_size, x:x+bar_size] = square
  271. stimuli[self.warmup+i, :, :] = stimulus
  272. # Reshape frame into flat array and convert into Tensor object
  273. stimuli = stimuli.reshape(total_frames, self.frame_size)
  274. stimuli = torch.Tensor(stimuli).unsqueeze(0).to(self.device)
  275. return stimuli
  276. def fit_sine (self, x, y):
  277. # Fit to sine
  278. def func(x, a, b, c, d):
  279. return a*np.sin(b*x + c) + d
  280. def get_mse_loss (y, y_est):
  281. return np.sum((y-y_est)**2)/ len(y)
  282. # Get r_squared from https://stackoverflow.com/a/37899817
  283. def get_rsq (y, y_est):
  284. residuals = y - y_est
  285. ss_res = np.sum(residuals**2)
  286. ss_tot = np.sum((y-np.mean(y))**2)
  287. r_squared = 1 - (ss_res / ss_tot)
  288. return r_squared
  289. best_params = []
  290. for iteration in range(5):
  291. scale = 1-0.2*iteration
  292. n_random_guesses = 10000 if iteration == 0 else 1000
  293. params = []
  294. loss_list = []
  295. for _ in range(n_random_guesses):
  296. if iteration == 0:
  297. rand_params = [
  298. np.random.uniform(low=np.mean(y)-2.5, high=np.mean(y)+2.5),
  299. np.random.uniform(low=0, high=10),
  300. np.random.uniform(low=0, high=len(y)),
  301. np.random.uniform(low=np.min(y)-2.5, high=np.max(y)+2.5)
  302. ]
  303. else:
  304. prev_best_params = best_params[-1]
  305. rand_params_ = [
  306. np.random.uniform(low=-2*scale, high=2*scale),
  307. np.random.uniform(low=-1*scale, high=1*scale),
  308. np.random.uniform(low=-2*scale, high=2*scale),
  309. np.random.uniform(low=-2*scale, high=2*scale)
  310. ]
  311. rand_params = [p+prev_best_params[idx] for idx, p in enumerate(rand_params_)]
  312. # Get the estimated curve based on fitted parameters
  313. y_est = func(x, *rand_params)
  314. loss = get_mse_loss(y, y_est) #get_rsq(y, y_est)
  315. params.append(rand_params)
  316. loss_list.append(loss)
  317. # Get the index of the lowest RSQ, use this to find the
  318. # corresponding parameters used
  319. best_params.append(params[np.argmin(loss_list)])
  320. final_params = best_params[-1]
  321. final_y_est = func(x, *final_params)
  322. final_loss = get_mse_loss(y, final_y_est)
  323. final_rsq = get_rsq(y, final_y_est)
  324. return final_params, final_y_est, final_rsq, final_loss
  325. # Takes list of response as well as a start and
  326. # end offset for where curve fitting should occur
  327. # Returns modulation ratio, estimated curve and RSQ of curve fit
  328. def get_modulation_ratio (self, activity, start_offset, end_offset):
  329. x = np.arange(start_offset, end_offset)
  330. y = activity[start_offset:end_offset]
  331. final_params, final_y_est, final_rsq, final_loss = self.fit_sine (x, y)
  332. # Try one more time if it fails
  333. if final_loss < 0.05 or final_rsq > 0.5:
  334. final_params, final_y_est, final_rsq, final_loss = self.fit_sine (x, y)
  335. # Average unit activity
  336. f0 = np.mean(activity[self.warmup:])
  337. # Absolute of the amplitude of the fitted sine
  338. f1 = (abs(final_params[0]))
  339. mod_ratio = f1/f0
  340. # Reject f values for those units with poor sine fits
  341. #if (final_loss < 0.05 or final_rsq > 0.5) and f0 != 0:
  342. cc = scipy.stats.pearsonr(final_y_est, y)[0]
  343. if cc>0.9 and f0 != 0:
  344. return {
  345. 'modulation_ratio' : mod_ratio,
  346. 'modulation_ratio_all' : mod_ratio,
  347. 'modulation_ratio_cc' : cc,
  348. 'modulation_ratio_y' : final_y_est,
  349. 'modulation_ratio_rsq' : final_rsq,
  350. 'modulation_ratio_loss' : final_loss,
  351. 'modulation_ratio_params' : final_params
  352. }
  353. else:
  354. return {
  355. 'modulation_ratio' : False,
  356. 'modulation_ratio_all' : mod_ratio,
  357. 'modulation_ratio_cc' : cc,
  358. 'modulation_ratio_y' : final_y_est,
  359. 'modulation_ratio_y_true' : final_y_est,
  360. 'modulation_ratio_rsq' : final_rsq,
  361. 'modulation_ratio_loss' : final_loss,
  362. 'modulation_ratio_params' : final_params
  363. }
  364. def get_modulation_ratio_all (self):
  365. for g_idx, g in enumerate(self.data):
  366. for unit_idx, unit_data in enumerate(g):
  367. if unit_idx % 10 == 0:
  368. print(f'Starting group {g_idx}, unit {unit_idx}')
  369. modulation_data = self.get_modulation_ratio(
  370. unit_data['optimum_grating_response'], self.warmup, self.t_steps
  371. )
  372. unit_data["modulation_ratio"] = modulation_data["modulation_ratio"]
  373. unit_data["modulation_ratio_cc"] = modulation_data["modulation_ratio_cc"]
  374. unit_data["modulation_ratio_all"] = modulation_data["modulation_ratio_all"]
  375. unit_data["modulation_ratio_y"] = modulation_data["modulation_ratio_y"]
  376. unit_data["modulation_ratio_rsq"] = modulation_data["modulation_ratio_rsq"]
  377. unit_data["modulation_ratio_loss"] = modulation_data["modulation_ratio_loss"]
  378. unit_data["modulation_ratio_params"] = modulation_data["modulation_ratio_params"]
  379. # https://www.nature.com/articles/nn1786#Sec11
  380. # plaid_angle = spread of plaid components
  381. def get_plaid_pattern_index (self, grating_amplitude=1):
  382. # Test if true response is more similar to pattern or component predictions
  383. # Partial correlation of x and y controlling for z
  384. def partial_correlation (x, y, z):
  385. r_xy, _ = scipy.stats.pearsonr(x, y)
  386. r_xz, _ = scipy.stats.pearsonr(x, z)
  387. r_yz, _ = scipy.stats.pearsonr(y, z)
  388. return (r_xy - r_xz*r_yz) / ((1-r_xz**2)*(1-r_yz**2))**0.5
  389. # Converts from distribution of r values to normal distribution
  390. # (to compare across units)
  391. def fisher_r_to_z (a, df):
  392. return 0.5*np.log((1+a)/(1-a)) / (1/df)**0.5
  393. orientation_step = 6
  394. orientations = self.orientations[::6]
  395. for group_idx, group_data in enumerate(self.data):
  396. for unit_i, unit_data in enumerate(group_data):
  397. max_sf_idx = np.where(np.array(self.spatial_frequencies) == unit_data['preferred_sf'])[0][0]
  398. max_tf_idx = np.where(np.array(self.temporal_frequencies) == unit_data['preferred_tf'])[0][0]
  399. # Response if unit responds to plaid the same as a grating of same angle
  400. pattern_prediction = unit_data['mean_grating_responses'][max_sf_idx, ::orientation_step, max_tf_idx]
  401. z_p_arr = []
  402. z_c_arr = []
  403. for plaid_angle in [60, 90, 120, 150]:
  404. # Response if unit responds to individual components of the plaid
  405. # Shift is the amount to rotate the tuning curve (+/- half the plaid angle)
  406. shift = int( (plaid_angle/2) / (orientations[1]-orientations[0]) )
  407. component_prediction_a = np.roll(pattern_prediction, shift)
  408. component_prediction_b = np.roll(pattern_prediction, -shift)
  409. component_prediction = (component_prediction_a + component_prediction_b) / 2
  410. # Now get 'true' response to the plaid (as a function of orientation)
  411. plaid_response = []
  412. with torch.no_grad():
  413. for ori_idx, ori in enumerate(orientations):
  414. # Plaid at unit's preferred tf and sf
  415. grating_a = self.get_grating_stimuli(
  416. unit_data['preferred_sf'],
  417. ori+plaid_angle/2,
  418. unit_data['preferred_tf'],
  419. grating_amplitude/2,
  420. 50
  421. )
  422. grating_b = self.get_grating_stimuli(
  423. unit_data['preferred_sf'],
  424. ori-plaid_angle/2,
  425. unit_data['preferred_tf'],
  426. grating_amplitude/2,
  427. 50
  428. )
  429. plaid = grating_a + grating_b
  430. # Feed into model
  431. _, hidden_state = self.model(plaid)
  432. # Get mean response for current unit
  433. mean_plaid_response = hidden_state[0, self.warmup:, unit_data['hidden_unit_index']] \
  434. .mean() \
  435. .detach() \
  436. .numpy()
  437. plaid_response.append(mean_plaid_response)
  438. r_p = partial_correlation(plaid_response, pattern_prediction, component_prediction)
  439. r_c = partial_correlation(plaid_response, component_prediction, pattern_prediction)
  440. z_p = fisher_r_to_z(r_p, len(pattern_prediction)-3)
  441. z_c = fisher_r_to_z(r_c, len(component_prediction)-3)
  442. z_p_arr.append(z_p)
  443. z_c_arr.append(z_c)
  444. z_p_mean = np.mean(z_p_arr)
  445. z_c_mean = np.mean(z_c_arr)
  446. pattern_idx = z_p_mean - z_c_mean
  447. # Save all these results
  448. unit_data['plaid_rp'] = r_p
  449. unit_data['plaid_rc'] = r_c
  450. unit_data['plaid_zp'] = z_p_mean
  451. unit_data['plaid_zc'] = z_c_mean
  452. unit_data['plaid_pattern_index'] = pattern_idx
  453. unit_data['plaid_response'] = plaid_response
  454. print("Finished unit {} / {}, group {} / {}".format(
  455. unit_i+1, len(group_data), group_idx+1, len(self.data)
  456. ))
  457. return self
  458. def get_moving_bar_responses (self):
  459. bar_size = 20
  460. # -5 because we are effectively convolving with moving bar "kernel" without padding
  461. convolved_size_r = self.frame_shape[0] - bar_size
  462. convolved_size_pixels_r = np.arange(convolved_size_r)
  463. convolved_size_c = self.frame_shape[1] - bar_size
  464. convolved_size_pixels_c = np.arange(convolved_size_c)
  465. # Pre-allocate array to hold mean response to stimulus at each location/direction
  466. for group_data in self.data:
  467. for unit_data in group_data:
  468. unit_data["moving_bar_responses"] = np.zeros((4, convolved_size_r, convolved_size_c))
  469. # Loop through each orientation and spatial location
  470. for orient_idx, orient in enumerate([0, 90, 180, 270]):
  471. for row_idx, row in enumerate(convolved_size_pixels_r):
  472. for col_idx, col in enumerate(convolved_size_pixels_c):
  473. stimuli = self.get_moving_bar_stimuli(orient, col, row, bar_size=bar_size)
  474. with torch.no_grad():
  475. _, hidden_state = self.model(stimuli)
  476. for group_data in self.data:
  477. for unit_data in group_data:
  478. # Get response at each time step of stimulus
  479. unit_activity = hidden_state[0, :, unit_data["hidden_unit_index"]].cpu().numpy()
  480. # Take baseline response as mean of warmup period
  481. #baseline_response = np.mean(unit_activity[:self.warmup])
  482. # Take overlal mean response as non-warmup period subtracted from baseline
  483. mean_response = np.mean(unit_activity[:]) # - baseline_response)
  484. unit_data["moving_bar_responses"][orient_idx, row_idx, col_idx] = mean_response
  485. print("Got responses for {} degrees, row {}".format(orient, row+1))
  486. # Also compute average response across orientations to produce 'heatmap'
  487. for group_data in self.data:
  488. for unit_data in group_data:
  489. unit_data["mean_moving_bar_responses"] = np.mean(unit_data["moving_bar_responses"], axis=0)
  490. return self
  491. def get_receptive_field_centres (self):
  492. def twoD_Gaussian(xy, amplitude, xo, yo, sigma_x, sigma_y, theta, offset):
  493. xy = (16, 16)
  494. x = np.linspace(0, int(xy[0]-1), int(xy[0]))
  495. y = np.linspace(0, int(xy[1]-1), int(xy[1]))
  496. x,y = np.meshgrid(x, y)
  497. xo = float(xo)
  498. yo = float(yo)
  499. a = (np.cos(theta)**2)/(2*sigma_x**2) + (np.sin(theta)**2)/(2*sigma_y**2)
  500. b = -(np.sin(2*theta))/(4*sigma_x**2) + (np.sin(2*theta))/(4*sigma_y**2)
  501. c = (np.sin(theta)**2)/(2*sigma_x**2) + (np.cos(theta)**2)/(2*sigma_y**2)
  502. g = offset + amplitude*np.exp( - (a*((x-xo)**2) + 2*b*(x-xo)*(y-yo)
  503. + c*((y-yo)**2)))
  504. return g.ravel()
  505. def weighted_centre (resp):
  506. (X,Y) = np.meshgrid(np.arange(0, resp.shape[0]), np.arange(0, resp.shape[1]))
  507. x_coord = (X*resp).sum() / resp.sum().astype("float")
  508. y_coord = (Y*resp).sum() / resp.sum().astype("float")
  509. return (y_coord, x_coord)
  510. for group_idx, group_data in enumerate(self.data):
  511. poor_fit_count = 0
  512. for unit_i, unit_data in enumerate(group_data):
  513. resp = unit_data["mean_moving_bar_responses"]
  514. max_val = np.max(resp)
  515. mean_centre = weighted_centre(resp)
  516. initial_guess = (max_val, mean_centre[1], mean_centre[0], 3, 3, 0, 0)
  517. try:
  518. popt, pcov = opt.curve_fit(
  519. twoD_Gaussian,
  520. resp.shape,
  521. resp.reshape(-1),
  522. p0=initial_guess,
  523. bounds=(
  524. [0, -np.inf, -np.inf, 0, 0, 0, 0],
  525. [np.inf, np.inf, np.inf, np.inf, np.inf, np.pi*2, np.inf]
  526. )
  527. )
  528. resp_gaussian_fit = twoD_Gaussian(resp.shape, *popt)
  529. centre = popt[2], popt[1]
  530. cc = scipy.stats.pearsonr(resp.reshape(-1), resp_gaussian_fit)[0]
  531. except Exception as e:
  532. centre = (np.nan, np.nan)
  533. # 'Convolution' with moving bar means this center needs to be scaled back up to
  534. # full size of frame
  535. scale_factor = self.frame_shape[0]/resp.shape[0]
  536. centre = (centre[0] * scale_factor, centre[1] * scale_factor)
  537. if np.isnan(centre[0]) or np.isnan(centre[1]) or cc < 0.6:
  538. unit_data["receptive_field_centre"] = (-1, -1)
  539. unit_data["receptive_field_gaussian_fit"] = np.nan
  540. unit_data["receptive_field_response"] = resp
  541. poor_fit_count += 1
  542. else:
  543. unit_data["receptive_field_centre"] = centre
  544. unit_data["receptive_field_gaussian_fit"] = resp_gaussian_fit.reshape(16, 16)
  545. if unit_i % 100 == 99:
  546. print("Finished unit {} / {}, group {} / {}".format(
  547. unit_i+1, len(group_data), group_idx+1, len(self.data)
  548. ))
  549. print("{}/{} units were rejected as Gaussian could not be fit to receptive fields in group {}".format(
  550. poor_fit_count, len(group_data), group_idx+1
  551. ))
  552. return self

VirtualNetworkPhysiology.py at commit dc9bb5f, no license · at the source

Overview

Authors: Sebastian Klavinskis-Whiting1, Andrew J. King1, Nicol S. Harper1
  1. Department of Physiology, Anatomy and Genetics, University of Oxford, Oxford, United Kingdom
Institutions: University of Oxford (United Kingdom)
Journal: PLoS computational biology, volume 22, issue 5, article e1013138
Dates: received 18 May 2025; accepted 7 May 2026; published online 28 May 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pcbi.1013138 · PMID 42207826 · PMCID PMC13252843 · OpenAlex W4410450465
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), systems (subfield)
Methods: Spectral & time-frequency, Connectivity, Statistics, Machine learning
MeSH: Models, Neurological*, Visual Cortex*, Visual Pathways*, Animals, Computational Biology, Computer Simulation, Neurons, Recurrent Neural Networks (* major topic)
Journal subjects: Research and Analysis Methods, Mathematical and Statistical Techniques, Statistical Methods, Forecasting, Physical Sciences, Mathematics, Statistics, Biology and life sciences, Organisms, Eukaryota, Animals, Vertebrates, Amniotes, Mammals, Primates, Monkeys, Old World monkeys, Macaque, Zoology, Anatomy, Brain, Visual Cortex, Medicine and Health Sciences, Neuroscience, Neuronal Tuning, Cognitive Science, Cognitive Psychology, Perception, Sensory Perception, Vision, Psychology, Social Sciences, Cell Biology, Cellular Types, Animal Cells, Neurons, Cellular Neuroscience, Physiology, Sensory Physiology, Visual System, Sensory Systems, Computer and Information Sciences, Neural Networks
Topic: Morphological variations and asymmetry (Geometry and Topology, Mathematics), according to OpenAlex
Funding: Wellcome Trust (WT108369/Z/2015/Z); Nuffield Department of Clinical Neurosciences, University of Oxford (Not applicable)
Citations: not cited yet (Europe PMC); 71 references in the paper

Abstract

A major goal of neuroscience is to identify general principles that can explain the diverse structures and functions of the brain. The principle of temporal prediction provides one approach, arguing that the sensory brain is optimized to represent stimulus features that efficiently predict the immediate future input. Previous work has demonstrated that feedforward hierarchical temporal prediction models can capture the tuning properties of neurons along the visual pathway, and that recurrent temporal prediction models can explain local functional connectivity within primary visual cortex. However, the visual system is also characterized by extensive inter-areal feedback recurrency, which existing models lack. We aimed to better account for the dynamic features of neurons in the visual cortex by incorporating both local recurrency and inter-areal feedback connectivity into a hierarchical temporal prediction model. The resulting model captured tuning properties along the dorsal visual pathway, including pattern motion selectivity and surround suppression, and the contribution of inter-areal connectivity to these properties. Moreover, compared with several alternative normative models, the hierarchical recurrent temporal prediction model provided the closest fit to these tuning properties and was best able to explain neuronal responses to natural stimuli. Accordingly, temporal prediction accounts well for information processing along the visual pathway.

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

Repository

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

sebbkw/hierarchical_temporal_prediction

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: dc9bb5f145bd1969d261b1af1531ab6cc6c4d078, 2 May 2026
Languages: Python (26), Jupyter (7), Shell (1)
Size: 81 files, 34 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, 7 notebooks
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (22 files), PyTorch (21 files), SciPy (12 files), Matplotlib (9 files), Pingouin (6 files), pandas (5 files), OpenCV (3 files), statsmodels (2 files), AllenSDK (1 file), h5py (1 file), scikit-image (1 file), scikit-learn (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
35 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:

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

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

Data

Datasets cited

Data Availability

All code supporting this study is publicly available at https://github.com/sebbkw/hierarchical_temporal_prediction. The model was trained using a novel dataset drawn from around 2.5 hours of wildlife videos, which were obtained from public sources, and are available at https://osf.io/hf2pj/ (comprising https://doi.org/10.5281/zenodo.20041063; https://doi.org/10.5281/zenodo.20041567; https://doi.org/10.5281/zenodo.20041596).

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

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

Cite

This paper

Klavinskis-Whiting, S., King, A. J., & Harper, N. S. (2026). Hierarchical recurrent temporal prediction as a model of the mammalian dorsal visual pathway. PLoS computational biology, 22(5), e1013138. https://doi.org/10.1371/journal.pcbi.1013138

BibTeX

@article{klavinskiswhiting2026hierarchical,
author = {Klavinskis-Whiting, Sebastian and King, Andrew J. and Harper, Nicol S.},
title = {{Hierarchical recurrent temporal prediction as a model of the mammalian dorsal visual pathway}},
journal = {PLoS computational biology},
year = {2026},
month = may,
volume = {22},
number = {5},
pages = {e1013138},
publisher = {PLOS},
issn = {1553-734X},
doi = {10.1371/journal.pcbi.1013138},
url = {https://doi.org/10.1371/journal.pcbi.1013138},
pmid = {42207826},
pmcid = {PMC13252843}
}

RIS

TY - JOUR
AU - Klavinskis-Whiting, Sebastian
AU - King, Andrew J.
AU - Harper, Nicol S.
TI - Hierarchical recurrent temporal prediction as a model of the mammalian dorsal visual pathway
T2 - PLoS computational biology
J2 - PLoS Comput Biol
PY - 2026
DA - 2026/05/28
VL - 22
IS - 5
SP - e1013138
SN - 1553-734X
PB - PLOS
DO - 10.1371/journal.pcbi.1013138
UR - https://doi.org/10.1371/journal.pcbi.1013138
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pcbi.1013138",
"type": "article-journal",
"title": "Hierarchical recurrent temporal prediction as a model of the mammalian dorsal visual pathway",
"container-title": "PLoS computational biology",
"author": [
{
"family": "Klavinskis-Whiting",
"given": "Sebastian"
},
{
"family": "King",
"given": "Andrew J."
},
{
"family": "Harper",
"given": "Nicol S."
}
],
"container-title-short": "PLoS Comput Biol",
"volume": "22",
"issue": "5",
"page": "e1013138",
"DOI": "10.1371/journal.pcbi.1013138",
"PMID": "42207826",
"PMCID": "PMC13252843",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pcbi.1013138",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
28
]
]
}
}

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-76098-y [code]
A single computational objective can produce specialization of streams in visual cortex.
Journal: Nature communications
In common: scikit-image, h5py, statsmodels, 6 other tools, 4 references
[2] doi:10.1016/j.isci.2026.116825 [code]
Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.
Journal: iScience
In common: Pingouin, OpenCV, scikit-image, 8 other tools, systems
[3] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: Pingouin, OpenCV, scikit-image, 8 other tools, systems
[4] doi:10.7554/elife.108408 [code]
Frequency and laminar profile of feature-specific visual activity revealed by interleaved EEG-fMRI.
Journal: eLife
In common: Pingouin, scikit-image, h5py, 5 other tools, systems, 2 references
[5] doi:10.1371/journal.pbio.3003915 [code]
Noise-invariant representations of sound emerge along the canonical cortical hierarchy.
Journal: PLoS biology
In common: OpenCV, scikit-image, h5py, 6 other tools, 2 references
[6] doi:10.1038/s41597-026-07248-6 [code]
A large-scale fMRI dataset for vision-language semantic association.
Journal: Scientific data
In common: OpenCV, scikit-image, h5py, 6 other tools, 2 references
[7] doi:10.1038/s41593-026-02314-z [code]
Low-dimensional population dynamics in the brainstem gate REM sleep.
Journal: Nature neuroscience
In common: Pingouin, OpenCV, h5py, 6 other tools, systems, 1 reference
[8] doi:10.1038/s41467-026-76939-w [code]
HIPPIE: a generative model for electrophysiological analysis across species, technologies, and modalities.
Journal: Nature communications
In common: AllenSDK, h5py, statsmodels, 6 other tools, 1 reference
[9] doi:10.1038/s42003-026-10169-0 [code]
Shared representations in brains and models reveal a two-route cortical organization during scene perception.
Journal: Communications biology
In common: h5py, statsmodels, PyTorch, 5 other tools, 3 references
[10] doi:10.1016/j.isci.2026.117375 [code]
Motor priming is associated with widespread recruitment into neural ensembles and more rapid ensemble transitions.
Journal: iScience
In common: Pingouin, OpenCV, h5py, 7 other tools, systems

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.