OSCR

How the visual brain can learn to parse images using a multiscale, incremental grouping process.

Code ↔ Paper

12 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 12 matches
  1. [1] § Model and task › Recurrent network › Hidden layers ( ). ↔ src/models/recurrent_network.py, lines 83–212 · score 0.72 · progressively coarser spatial, spatial resolution, Hidden layers, hierarchical, VIP, RF
  2. [2] § Model and task › Recurrent network › Input layer ( ). ↔ src/models/layers.py, lines 293–349 · score 0.69 · lateral inhibition, layer receives, higher layer, sensory, signals, modulated
  3. [3] § Model and task › Recurrent network › Input layer ( ). ↔ src/models/layers.py, lines 293–349 · score 0.64 · lateral inhibition, feedback modulated, sensory, signal, gated, Pyramidal
  4. [4] § Model and task ↔ src/models/layers.py, lines 351–438 · score 0.62 · lower layers, horizontal connections, higher layers, gated, feedback, model
  5. [5] § Model and task › Recurrent network ↔ src/models/recurrent_network.py, lines 83–212 · score 0.61 · pyramidal neurons, horizontal connections, recurrent networks, VIP, feedback, tracing
  6. [6] § Model and task › Feedforward network ↔ src/tasks/tasks.py, lines 110–248 · score 0.59 · distractor curves, curve tracing task, target curve, selection, RF, stimulus
  7. [7] § Model and task › Training › Training of the recurrent network. ↔ src/models/recurrent_network.py, lines 324–342 · score 0.57 · reward prediction error, recurrent network, inspired, Weights, Training
  8. [8] § Model and task › Recurrent network › Hidden layers ( ). ↔ src/models/layers.py, lines 351–438 · score 0.57 · hidden layers integrated, VIP activity, horizontal, feedback
  9. [9] § Model and task ↔ src/tasks/tasks.py, lines 110–248 · score 0.54 · distractor curve, curve tracing task, overlapping, coarser, loop, selection
  10. [10] § Model and task › Recurrent network › Input layer ( ). ↔ src/models/layers.py, lines 144–170 · score 0.53 · lateral inhibitory, inhibitory weights, kernel, connections, layer
  11. [11] § Model and task › Recurrent network › Hidden layers ( ). ↔ src/models/layers.py, lines 440–553 · score 0.53 · receptive field, feedforward weights, stride, selection, Hidden, layer
  12. [12] § Model and task › Training › Training of feedforward networks. ↔ src/models/recurrent_network.py, lines 16–22 · score 0.52 · reinforcement learning, recurrent network, curve tracing, architecture

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 · 600 lines · 28 KB · no license · 6 matches

  1. # -*- coding: utf-8 -*-
  2. """
  3. Created on Fri Oct 29 12:29:28 2021
  4. @author: Sami
  5. """
  6. from typing import Optional, List, Tuple
  7. import numpy as np
  8. import torch
  9. import torch.nn as nn
  10. import torch.nn.functional as F
  11. class CustomLayer(nn.Module):
  12. """Base class for custom neural network layers with specialized initialization and update rules.
  13. This class provides common functionality for input, hidden, and output layers including
  14. weight initialization, activation functions, and biologically-inspired learning rules.
  15. """
  16. def __init__(self) -> None:
  17. """Initialize the custom layer with default parameters."""
  18. self.initialisation_range = 0.1
  19. super().__init__()
  20. def activation_function(self, x: torch.Tensor) -> torch.Tensor:
  21. """Apply ReLU activation function.
  22. Args:
  23. x (torch.Tensor): Input tensor.
  24. Returns:
  25. torch.Tensor: Activated output.
  26. """
  27. return torch.relu(x)
  28. def step_function(self, x: torch.Tensor) -> torch.Tensor:
  29. """Apply a piecewise linear step function with saturation.
  30. Args:
  31. x (torch.Tensor): Input tensor.
  32. Returns:
  33. torch.Tensor: Transformed tensor with linear region and saturation.
  34. """
  35. eps = 1
  36. phi = 1
  37. x[x <= eps/phi] = phi * x[x <= eps/phi]
  38. x[x <= 0] = 0
  39. x[x >= eps/phi] = eps
  40. return x
  41. def gating_function(self, x: torch.Tensor) -> torch.Tensor:
  42. """Apply a smooth gating function for modulation.
  43. Args:
  44. x (torch.Tensor): Input tensor.
  45. Returns:
  46. torch.Tensor: Gated output with smooth saturation.
  47. """
  48. param = 100
  49. return torch.relu(param*x/(torch.sqrt(1+(param**2)*(x**2))))
  50. def initialize_feedforward_weights(self, layer: nn.Module, one_to_one: bool = False, change_scale: bool = False, receptive_field_size: int = 1) -> None:
  51. """Initialize feedforward connection weights with specific patterns.
  52. Args:
  53. layer (nn.Module): The layer whose weights to initialize.
  54. one_to_one (bool): If True, use one-to-one connectivity pattern.
  55. change_scale (bool): If True, initialize for scale-changing connections.
  56. receptive_field_size (int): Size of the receptive field for spatial connections.
  57. """
  58. # Initialize weight tensor with zeros
  59. weights = torch.zeros(layer.weight.shape)
  60. if layer.bias is not None:
  61. bias = torch.zeros(layer.bias.shape)
  62. if change_scale:
  63. # For scale-changing connections, initialize all weights uniformly
  64. for f in range(weights.shape[0]):
  65. for f2 in range(weights.shape[1]):
  66. weights[f, f2, :, :] = self.initialisation_range * np.random.rand() + 1
  67. else:
  68. # For same-scale connections
  69. for f in range(weights.shape[0]):
  70. for f2 in range(weights.shape[1]):
  71. if not one_to_one:
  72. # Localized receptive field: only connect nearby spatial locations
  73. lower_bound = self.grid_size - receptive_field_size//2 - 1
  74. upper_bound = self.grid_size + receptive_field_size//2
  75. weights[f, f2, lower_bound:upper_bound,lower_bound:upper_bound] = 0.001 * np.random.rand()
  76. else:
  77. # One-to-one connections: single weight per feature pair
  78. if self.layer_type != "output":
  79. weights[f,f2,0] = self.initialisation_range * np.random.rand() + 1
  80. else:
  81. # Output layer uses smaller initial weights
  82. weights[f, f2, 0] = 0.001 + 0.001 * np.random.rand()
  83. # Initialize biases to 1 (positive baseline)
  84. if layer.bias is not None:
  85. for f in range(bias.shape[0]):
  86. bias[f] = 1
  87. # Set as learnable parameters
  88. layer.weight = torch.nn.Parameter(weights)
  89. if layer.bias is not None:
  90. layer.bias = torch.nn.Parameter(bias)
  91. def initialize_feedback_weights(self, layer: nn.Module, change_scale: bool = False) -> None:
  92. """Initialize feedback connection weights.
  93. Args:
  94. layer (nn.Module): The layer whose weights to initialize.
  95. change_scale (bool): If True, initialize for scale-changing connections.
  96. """
  97. # Initialize weight tensor with zeros
  98. weights = torch.zeros(layer.weight.shape)
  99. if layer.bias is not None:
  100. bias = torch.zeros(layer.bias.shape)
  101. # Set up feedback connections (top-down modulation)
  102. for f in range(weights.shape[0]):
  103. for f2 in range(weights.shape[1]):
  104. if self.layer_type != "output":
  105. if change_scale:
  106. # Scale-changing feedback: full spatial connectivity
  107. weights[f, f2, :, :] = self.initialisation_range * np.random.rand() + 0.5
  108. else:
  109. # Local feedback: connect to 4 nearest neighbors (cross pattern)
  110. # This creates a plus-shaped connectivity pattern
  111. weights[f, f2, 0,1] = self.initialisation_range * np.random.rand() + 0.1 # Top
  112. weights[f, f2, 1,0] = self.initialisation_range * np.random.rand() + 0.1 # Left
  113. weights[f, f2, 1,2] = self.initialisation_range * np.random.rand() + 0.1 # Right
  114. weights[f, f2, 2,1] = self.initialisation_range * np.random.rand() + 0.1 # Bottom
  115. # Set as learnable parameters
  116. layer.weight = torch.nn.Parameter(weights)
  117. if layer.bias is not None:
  118. layer.bias = torch.nn.Parameter(bias)
  119. def initialize_inhibitory_weights(self, layer: nn.Module) -> None:
  120. """Initialize inhibitory connection weights with negative biases.
  121. Args:
  122. layer (nn.Module): The layer whose weights to initialize.
  123. """
  124. # Initialize weight tensor with zeros
  125. weights = torch.zeros(layer.weight.shape)
  126. if layer.bias is not None:
  127. bias = torch.zeros(layer.bias.shape)
  128. # Set up lateral inhibition: only within same feature channel (f == f2)
  129. for f in range(weights.shape[0]):
  130. for f2 in range(weights.shape[1]):
  131. if f == f2:
  132. # Self-connection at center of 3x3 kernel
  133. weights[f, f2, 1, 1] = 1
  134. # Initialize biases to negative values for inhibitory effect
  135. if layer.bias is not None:
  136. for f in range(bias.shape[0]):
  137. bias[f] = -1 - self.initialisation_range * np.random.rand()
  138. # Set as learnable parameters
  139. layer.weight = torch.nn.Parameter(weights)
  140. if layer.bias is not None:
  141. layer.bias = torch.nn.Parameter(bias)
  142. def average_traces(self, traces: torch.Tensor, mask: torch.Tensor, receptive_field_size: int = 1) -> torch.Tensor:
  143. """Average gradient traces over spatial dimensions for weight updates.
  144. Args:
  145. traces (torch.Tensor): Gradient traces to average.
  146. mask (torch.Tensor): Mask indicating valid connections.
  147. receptive_field_size (int): Size of receptive field for output layers.
  148. Returns:
  149. torch.Tensor: Averaged traces.
  150. """
  151. traces *= mask
  152. if self.layer_type != "output":
  153. m = torch.mean(traces[:,:,:,:], axis=(2,3))
  154. m = m[:, :, None,None]
  155. traces[:,:] = m
  156. else:
  157. intermmask = torch.ones((2 * self.grid_size - 1, 2 * self.grid_size - 1))
  158. lower_bound = self.grid_size - receptive_field_size//2 - 1
  159. upper_bound = self.grid_size + receptive_field_size//2
  160. intermmask = torch.zeros((2 * self.grid_size - 1, 2 * self.grid_size - 1))
  161. intermmask[lower_bound:upper_bound,lower_bound:upper_bound] = 1
  162. m = torch.mean(traces[:, :, intermmask == 1], axis=2)
  163. traces[:, :, intermmask == 1] = m[:,None]
  164. return(traces)
  165. def update_weight(self, layer: nn.Module, upper: torch.Tensor, beta: float, delta: float, mask: Optional[torch.Tensor] = None, z: Optional[torch.Tensor] = None, average: bool = True, receptive_field_size: int = 1, inhib: bool = False, change_scale: bool = False) -> nn.Module:
  166. """Update layer weights using reward prediction error.
  167. Args:
  168. layer (nn.Module): The layer to update.
  169. upper (torch.Tensor): Upper layer activity.
  170. beta (float): Learning rate.
  171. delta (float): Reward prediction error.
  172. mask (torch.Tensor, optional): Connection mask.
  173. z (torch.Tensor, optional): Gradient output.
  174. average (bool): Whether to average traces.
  175. receptive_field_size (int): Receptive field size.
  176. inhib (bool): If True, update inhibitory connections only.
  177. change_scale (bool): If True, handle scale-changing connections.
  178. Returns:
  179. nn.Module: Updated layer.
  180. """
  181. with torch.no_grad():
  182. # Compute gradients of output with respect to weights and biases
  183. # This implements eligibility traces for credit assignment
  184. if layer.bias is not None:
  185. delta_weight = torch.autograd.grad(upper, [layer.weight, layer.bias], grad_outputs=z, retain_graph=True, allow_unused=True)
  186. delta_bias = delta_weight[1]
  187. # Clamp gradient magnitudes to prevent instability
  188. delta_bias = torch.clamp(delta_bias, None, 1)
  189. else:
  190. delta_weight = torch.autograd.grad(upper, layer.weight, grad_outputs=z, retain_graph=True, allow_unused=True)
  191. delta_weight = delta_weight[0]
  192. # Clamp weight gradients to prevent large updates
  193. delta_weight = torch.clamp(delta_weight, None, 1)
  194. if inhib == False:
  195. # Update excitatory weights using the reward prediction error
  196. # delta_weight acts as eligibility trace, delta is reward prediction error
  197. if (self.layer_type == 'output' and average) or change_scale:
  198. # For output and scale-changing layers, average gradients spatially
  199. delta_weight = self.average_traces(delta_weight, mask, receptive_field_size=receptive_field_size)
  200. elif mask is not None:
  201. # Apply connection mask to enforce sparse connectivity
  202. delta_weight = delta_weight * mask
  203. # Weight update: w = w + learning_rate * RPE * eligibility_trace
  204. weight_update = layer.weight + beta * delta * delta_weight
  205. layer.weight.copy_(weight_update)
  206. else:
  207. # Update inhibitory biases only (weights remain fixed)
  208. bias_update = layer.bias + beta * delta * delta_bias
  209. layer.bias.copy_(bias_update)
  210. return(layer)
  211. def make_mask(self, weight: torch.Tensor) -> torch.Tensor:
  212. """Create a binary mask from weight tensor.
  213. Args:
  214. weight (torch.Tensor): Weight tensor.
  215. Returns:
  216. torch.Tensor: Binary mask (1 where weights are non-zero, 0 elsewhere).
  217. """
  218. mask = torch.clone(weight.detach())
  219. mask[mask != 0] = 1
  220. return mask
  221. def to(self, device: torch.device) -> None:
  222. """Move layer parameters to specified device.
  223. Args:
  224. device (torch.device): Target device (CPU or CUDA).
  225. """
  226. if self.layer_type == 'input':
  227. self.FB.to(device)
  228. self.lateral_inhibition.to(device)
  229. self.feedback_mask = self.feedback_mask.to(device)
  230. self.inhibition_mask = self.inhibition_mask.to(device)
  231. elif self.layer_type == 'hidden':
  232. self.FF.to(device)
  233. self.feedforward_mask = self.feedforward_mask.to(device)
  234. self.horizontal_mask = self.horizontal_mask.to(device)
  235. self.H.to(device)
  236. if self.has_feedback:
  237. self.FB.to(device)
  238. self.feedback_mask = self.feedback_mask.to(device)
  239. elif self.layer_type == 'output':
  240. for layer in range(len(self.skip_weights)):
  241. self.skip_weights[layer].to(device)
  242. for layer in range(1,len(self.skip_masks)):
  243. self.skip_masks[layer] = self.skip_masks[layer].to(device)
  244. class InputLayer(CustomLayer):
  245. """Input layer with feedback modulation and lateral inhibition.
  246. This layer receives sensory input and is modulated by feedback from higher layers
  247. through VIP and SOM interneuron populations.
  248. """
  249. def __init__(self, feature_in: int, feature_out: int) -> None:
  250. """Initialize input layer.
  251. Args:
  252. feature_in (int): Number of input features.
  253. feature_out (int): Number of output features from higher layer.
  254. """
  255. super().__init__()
  256. self.layer_type = 'input'
  257. K_size = 3
  258. self.FB = nn.Conv2d(feature_out, feature_in, K_size, stride=1, padding='same',bias=False)
  259. self.lateral_inhibition = nn.Conv2d(feature_in, feature_in, K_size, stride=1, padding='same',bias=True)
  260. # Initializing weights
  261. self.initialize_feedback_weights(self.FB)
  262. self.initialize_inhibitory_weights(self.lateral_inhibition)
  263. # Setting the mask to have connections only between neighboours
  264. self.feedback_mask = self.make_mask(self.FB.weight)
  265. self.inhibition_mask = self.make_mask(self.lateral_inhibition.weight)
  266. def forward(self, upper_y: torch.Tensor, input_stimulus: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
  267. """Forward pass through input layer.
  268. Args:
  269. upper_y (torch.Tensor): Modulation signal from upper layer.
  270. input_stimulus (torch.Tensor): Sensory input stimulus.
  271. Returns:
  272. tuple: (modulated_output, vip_activity, som_activity)
  273. """
  274. Y = input_stimulus
  275. vip_activity = self.step_function(self.FB(upper_y))
  276. som_activity = self.activation_function(1 - vip_activity)
  277. modulated_output = self.activation_function((-self.lateral_inhibition(som_activity)) * self.gating_function(Y))
  278. return(modulated_output, vip_activity, som_activity)
  279. def update_layer(self, upper: List[torch.Tensor], z: List[torch.Tensor], beta: float, delta: float) -> None:
  280. """Update layer weights based on reward prediction error.
  281. Args:
  282. upper (list): Upper layer activities [pyramidal, som, vip].
  283. z (list): Gradients [pyramidal, som, vip].
  284. beta (float): Learning rate.
  285. delta (float): Reward prediction error.
  286. """
  287. self.FB = self.update_weight(self.FB, upper[2], beta, delta, self.feedback_mask, z[2],change_scale=False)
  288. self.lateral_inhibition = self.update_weight(self.lateral_inhibition, upper[0], beta, delta, self.inhibition_mask, z[0],inhib=True)
  289. class HiddenLayer(CustomLayer):
  290. """Hidden layer with feedforward, feedback, and horizontal connections.
  291. This layer integrates information from lower layers (feedforward), higher layers
  292. (feedback), and within the same layer (horizontal connections).
  293. """
  294. def __init__(self, feature_in_lower: int, feature_out: int, feature_in_higher: int, big_pixels_size: int, grid_size: int, has_feedback: bool = True, change_scale_fb: bool = False, change_scale_ff: bool = False, higher_scale: bool = False) -> None:
  295. """Initialize hidden layer.
  296. Args:
  297. feature_in_lower (int): Number of features from lower layer.
  298. feature_out (int): Number of output features.
  299. feature_in_higher (int): Number of features from higher layer.
  300. big_pixels_size (int): Stride for scale-changing connections.
  301. grid_size (int): Size of spatial grid.
  302. has_feedback (bool): Whether to include upper layer modulation.
  303. change_scale_fb (bool): Whether feedback changes scale.
  304. change_scale_ff (bool): Whether feedforward changes scale.
  305. higher_scale (bool): Whether this is a higher-scale layer.
  306. """
  307. super().__init__()
  308. self.layer_type = 'hidden'
  309. self.has_feedback = has_feedback # If the higher layer has a modulated group
  310. self.change_scale_fb = change_scale_fb
  311. self.change_scale_ff = change_scale_ff
  312. self.higher_scale = higher_scale
  313. self.big_pixels_size = big_pixels_size
  314. self.grid_size = grid_size
  315. # Making the weights
  316. if self.change_scale_ff:
  317. self.FF = nn.Conv2d(feature_in_lower, feature_out, self.big_pixels_size, stride=self.big_pixels_size, bias=True)
  318. else:
  319. self.FF = nn.Conv2d(feature_in_lower, feature_out, 1, stride=1, padding='same', bias=True)
  320. if has_feedback:
  321. self.FB = nn.ConvTranspose2d(feature_in_higher, feature_out, self.big_pixels_size, stride=self.big_pixels_size, padding=0,bias=False)
  322. self.H = nn.Conv2d(feature_out,feature_out,3,stride = 1,padding='same',bias=False)
  323. # Initializing the weights
  324. self.initialize_feedforward_weights(self.FF,change_scale = change_scale_ff,one_to_one=True)
  325. if has_feedback:
  326. self.initialize_feedback_weights(self.FB,change_scale = change_scale_fb)
  327. self.initialize_feedback_weights(self.H,change_scale=False)
  328. # Making the masks
  329. if has_feedback:
  330. self.feedback_mask = self.make_mask(self.FB.weight)
  331. self.feedforward_mask = self.make_mask(self.FF.weight)
  332. self.horizontal_mask = self.make_mask(self.H.weight)
  333. def forward(self, current_y: torch.Tensor, lower_ymod: Optional[torch.Tensor] = None, upper_y: Optional[torch.Tensor] = None, horiz: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
  334. """Forward pass through hidden layer.
  335. Args:
  336. current_y (torch.Tensor): Current layer activity.
  337. lower_ymod (torch.Tensor, optional): Input from lower layer.
  338. upper_y (torch.Tensor, optional): Modulation from upper layer.
  339. horiz (torch.Tensor, optional): Horizontal connections.
  340. Returns:
  341. tuple: (modulated_output, vip_activity, som_activity)
  342. """
  343. feedforward_input = self.FF(lower_ymod)
  344. if upper_y is not None:
  345. vip_activity = self.step_function(self.FB(upper_y) + self.H(horiz))
  346. else:
  347. vip_activity = self.step_function(self.H(horiz))
  348. som_activity = self.activation_function(1 - vip_activity)
  349. modulated_output = self.activation_function((feedforward_input - som_activity) * self.gating_function(current_y))
  350. return(modulated_output, vip_activity, som_activity)
  351. def update_layer(self, upper: List[torch.Tensor], z: List[torch.Tensor], beta: float, delta: float, train_v: bool = True) -> None:
  352. """Update layer weights based on reward prediction error.
  353. Args:
  354. upper (list): Upper layer activities [pyramidal, som, vip].
  355. z (list): Gradients [pyramidal, som, vip].
  356. beta (float): Learning rate.
  357. delta (float): Reward prediction error.
  358. train_v (bool): Whether to train VIP connections.
  359. """
  360. self.FF = self.update_weight(self.FF, upper[0], beta, delta, self.feedforward_mask, z[0],change_scale=self.change_scale_ff)
  361. if self.has_feedback:
  362. self.FB = self.update_weight(self.FB, upper[2], beta, delta, self.feedback_mask, z[2],change_scale = self.change_scale_fb)
  363. self.H = self.update_weight(self.H, upper[2], beta, delta, self.horizontal_mask, z[2],change_scale = False)
  364. class OutputLayer(CustomLayer):
  365. """Output layer that aggregates information from all hierarchical levels.
  366. This layer combines skip connections from input, hidden, and high-level layers
  367. to produce action values (Q-values).
  368. """
  369. def __init__(self, high_features: int, hidden_features: int, input_features: int, feature_out: int, grid_size: int, RF_size_list: List[int]) -> None:
  370. """Initialize output layer.
  371. Args:
  372. high_features (int): Number of high-level features.
  373. hidden_features (int): Number of hidden features.
  374. input_features (int): Number of input features.
  375. feature_out (int): Number of output features.
  376. grid_size (int): Size of spatial grid.
  377. RF_size_list (list): List of receptive field sizes for each scale.
  378. """
  379. super().__init__()
  380. self.grid_size = grid_size
  381. K_size = 2 * self.grid_size - 1
  382. self.layer_type = 'output'
  383. self.grid_size = grid_size
  384. self.RF_size_list = RF_size_list
  385. self.hidden_features = hidden_features
  386. self.high_features = high_features
  387. self.feature_out = feature_out
  388. self.skip_weights = []
  389. for layer in range(len(self.RF_size_list)):
  390. if layer == 0:
  391. self.skip_weights.append(nn.Conv2d(input_features, feature_out, 1, stride=self.RF_size_list[layer], padding='same', bias=False))
  392. elif layer == 1:
  393. self.skip_weights.append(nn.Conv2d(hidden_features, feature_out, K_size, stride=self.RF_size_list[layer], padding='same',bias=False))
  394. else:
  395. self.skip_weights.append(nn.ConvTranspose2d(high_features, feature_out,K_size, stride=self.RF_size_list[layer], padding=int(0.5*(2 * self.grid_size - 1 - self.RF_size_list[layer])),bias=False))
  396. self.skip_masks = [None]
  397. for layer in range(1,len(self.RF_size_list)):
  398. self.skip_masks.append(torch.zeros_like(self.skip_weights[layer].weight))
  399. lower_bound = self.grid_size-(self.RF_size_list[layer]//2)-1
  400. upper_bound = self.grid_size+(self.RF_size_list[layer]//2)
  401. self.skip_masks[layer][:,:,lower_bound:upper_bound,lower_bound:upper_bound] = 1/50
  402. # Initializing weights
  403. self.initialize_feedforward_weights(self.skip_weights[0], one_to_one=True)
  404. for layer in range(1,len(self.RF_size_list)):
  405. self.initialize_feedforward_weights(self.skip_weights[layer], receptive_field_size=self.RF_size_list[layer])
  406. self.skip_weights = nn.ModuleList(self.skip_weights)
  407. def forward(self, pyramidal_recurrent: List[torch.Tensor]) -> torch.Tensor:
  408. """Forward pass through output layer.
  409. Args:
  410. pyramidal_recurrent (list): List of pyramidal activities from all layers.
  411. Returns:
  412. torch.Tensor: Output Q-values for action selection.
  413. """
  414. Y = self.skip_weights[0](pyramidal_recurrent[0])
  415. for layer in range(1,len(pyramidal_recurrent)):
  416. Y = Y + self.skip_weights[layer](pyramidal_recurrent[layer])
  417. return(Y)
  418. def rescale(self, new_grid_size: int, device: torch.device) -> None:
  419. """Rescale output layer for different grid size.
  420. Args:
  421. new_grid_size (int): New spatial grid size.
  422. device (torch.device): Device to move parameters to.
  423. """
  424. K_size = 2 * new_grid_size - 1
  425. skip_weights = [self.skip_weights[0]]
  426. for layer in range(1,len(self.RF_size_list)):
  427. if layer == 1:
  428. skip_weights.append(nn.Conv2d(self.hidden_features, self.feature_out, K_size, stride=self.RF_size_list[layer], padding='same',bias=False))
  429. else:
  430. skip_weights.append(nn.ConvTranspose2d(self.high_features, self.feature_out,K_size, stride=self.RF_size_list[layer], padding=int(0.5*(2 * new_grid_size - 1 - self.RF_size_list[layer])),bias=False))
  431. for layer in range(1,len(self.RF_size_list)):
  432. weight = torch.zeros_like(skip_weights[layer].weight)
  433. lower_bound = new_grid_size - self.RF_size_list[layer]//2 - 1
  434. upper_bound = new_grid_size + self.RF_size_list[layer]//2
  435. non_zero_weights = torch.unique(self.skip_weights[layer].weight[self.skip_weights[layer].weight!=0])
  436. non_zero_weights = non_zero_weights[:,None,None,None]
  437. weight[:,:,lower_bound:upper_bound,lower_bound:upper_bound] = non_zero_weights
  438. skip_weights[layer].weight = torch.nn.Parameter(weight)
  439. self.skip_weights[layer] = skip_weights[layer]
  440. self.skip_weights[layer].to(device)
  441. self.skip_masks = [None]
  442. for layer in range(1,len(self.RF_size_list)):
  443. self.skip_masks.append(torch.zeros_like(self.skip_weights[layer].weight))
  444. lower_bound = new_grid_size-(self.RF_size_list[layer]//2)-1
  445. upper_bound = new_grid_size+(self.RF_size_list[layer]//2)
  446. self.skip_masks[layer][:,:,lower_bound:upper_bound,lower_bound:upper_bound] = 1/50
  447. self.grid_size = new_grid_size
  448. def update_layer(self, upper: torch.Tensor, beta: float, delta: float) -> None:
  449. """Update all skip connection weights.
  450. Args:
  451. upper (torch.Tensor): Upper layer activity.
  452. beta (float): Learning rate.
  453. delta (float): Reward prediction error.
  454. """
  455. for layer in range(len(self.skip_weights)):
  456. if layer == 0:
  457. self.skip_weights[layer] = self.update_weight(self.skip_weights[layer],upper,beta,delta,average=False)
  458. else:
  459. self.skip_weights[layer] = self.update_weight(self.skip_weights[layer],upper,beta,delta,self.skip_masks[layer],receptive_field_size = self.RF_size_list[layer])
  460. class FeedforwardLayer(CustomLayer):
  461. """Wrapper for pretrained feedforward networks.
  462. This layer encapsulates pretrained feedforward networks for object and curve detection
  463. at multiple scales.
  464. """
  465. def __init__(self, feedforward: nn.ModuleList, feedforward_interm: nn.ModuleList, num_scales: int) -> None:
  466. """Initialize feedforward layer.
  467. Args:
  468. feedforward (nn.ModuleList): List of feedforward layers.
  469. feedforward_interm (nn.ModuleList): List of intermediate layers.
  470. num_scales (int): Number of spatial scales.
  471. """
  472. super().__init__()
  473. self.num_scales = num_scales
  474. self.feedforward = feedforward
  475. self.feedforward_interm = feedforward_interm
  476. self.sig = nn.Sigmoid()
  477. def forward(self, x: torch.Tensor) -> List[torch.Tensor]:
  478. """Forward pass through feedforward network.
  479. Args:
  480. x (torch.Tensor): Input image.
  481. Returns:
  482. list: Multi-scale feature representations.
  483. """
  484. intern_representation = [None] * self.num_scales
  485. for layer in range(self.num_scales):
  486. if layer == 0:
  487. x = F.relu(self.feedforward[0](x))
  488. intern_representation[layer] = x
  489. else:
  490. interm = F.relu(self.feedforward_interm[layer](x))
  491. intern_representation[layer] = self.sig(self.feedforward[layer](interm))
  492. intern_representation[0] = intern_representation[0].detach()
  493. for layer in range(1,len(intern_representation)):
  494. intern_representation[layer] = torch.relu(intern_representation[layer] - 0.7)
  495. intern_representation[layer] = intern_representation[layer].detach()
  496. return intern_representation

layers.py at commit cb55334, no license · at the source

Overview

Authors: Sami Mollard1, Sander M. Bohte2,3, Pieter R. Roelfsema1,4,5,6
  1. Department of Vision & Cognition, Netherlands Institute for Neuroscience, Amsterdam, The Netherlands
  2. Machine Learning Group, Centrum Wiskunde & Informatica, Amsterdam, The Netherlands
  3. Swammerdam Institute for Life Sciences, University of Amsterdam, Amsterdam, Netherlands
  4. Laboratory of Visual Brain Therapy, Sorbonne Université, Institut National de la Santé et de la Recherche Médicale, Centre National de la Recherche Scientifique, Institut de la Vision, Paris, France
  5. Department of Integrative Neurophysiology, Center for Neurogenomics and Cognitive Research, VU University, Amsterdam, The Netherlands
  6. Department of Neurosurgery, Amsterdam University Medical Center, Amsterdam, The Netherlands
Journal: PLoS computational biology, volume 22, issue 4, article e1014193
Dates: received 17 June 2025; accepted 1 April 2026; published online 15 April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pcbi.1014193 · PMID 41984938 · PMCID PMC13095124 · OpenAlex W7154471241
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), non-human primate (organism)
Methods: Preprocessing
MeSH: Learning*, Models, Neurological*, Visual Cortex*, Animals, Computational Biology, Humans, Neurons, Visual Perception (* major topic)
Journal subjects: Biology and Life Sciences, Cell Biology, Cellular Types, Animal Cells, Neurons, Neuroscience, Cellular Neuroscience, Cognitive Science, Cognitive Psychology, Learning, Learning Curves, Psychology, Social Sciences, Learning and Memory, Anatomy, Brain, Visual Cortex, Medicine and Health Sciences, Perception, Sensory Perception, Vision, Computer and Information Sciences, Neural Networks, Organisms, Eukaryota, Animals, Vertebrates, Amniotes, Mammals, Primates, Monkeys, Zoology, Recurrent Neural Networks
Topic: Visual perception and processing mechanisms (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 60 references in the paper

Abstract

Natural scenes usually contain many objects that need to be segregated from each other and the background. Object-based attention is the process that groups image fragments belonging to the same objects. Curve-tracing tasks provide a special case, testing our ability to group image elements of an elongated curve. In the brain, curve-tracing is associated with the gradual spread of enhanced neuronal activity over the representation of the traced curve. Previous studies demonstrated that the tracing speed is higher if curves are far apart than if they are nearby. One hypothesis is that a larger distance between curves permits activity propagation in higher visual cortical areas. In these higher areas receptive fields are larger and connections exist between neurons representing image regions that are farther apart (Pooresmaeili et al., 2014). We propose a recurrent architecture for the scale-invariant tracing of curves and objects. The architecture is composed of a feedforward pathway that dynamically selects the appropriate scale for tracing, and a recurrent pathway for propagating enhanced neuronal activity through horizontal and feedback connections, enabled by a disinhibitory loop involving VIP and SOM interneurons. We trained the network using a biologically plausible reinforcement learning scheme and observed that training on short curves allowed the networks to generalize to longer curves and 2D-objects. The network chose the scale based on the distance between curves and the width of objects, just as in human psychophysics and the visual cortex of monkeys. The results provide a mechanistic account of the learning and execution of multiscale perceptual grouping in the brain.

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 12 matches between paragraphs and lines of code.

samimol/multiscale_tracing

License: none: the authors keep all their rights
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: cb55334b4075ad35b3e1c891f64cab10a477c93f, 2 March 2026
Languages: Python (12)
Size: 28 files, 12 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, environment (requirements.txt)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (8 files), PyTorch (8 files), scikit-image (2 files), SciPy (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
13 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;
  • 12 scripts, each with its path and the digest of its content;
  • 12 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

All the code used to train the networks and to analyze the data is available on the following GitHub address: https://github.com/samimol/multiscale_tracing.

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, 3 authors, 8 MeSH terms, 5 funders, 52 references.

Cite

This paper

Mollard, S., Bohte, S. M., & Roelfsema, P. R. (2026). How the visual brain can learn to parse images using a multiscale, incremental grouping process. PLoS computational biology, 22(4), e1014193. https://doi.org/10.1371/journal.pcbi.1014193

BibTeX

@article{mollard2026how,
author = {Mollard, Sami and Bohte, Sander M. and Roelfsema, Pieter R.},
title = {{How the visual brain can learn to parse images using a multiscale, incremental grouping process}},
journal = {PLoS computational biology},
year = {2026},
month = apr,
volume = {22},
number = {4},
pages = {e1014193},
publisher = {PLOS},
issn = {1553-734X},
doi = {10.1371/journal.pcbi.1014193},
url = {https://doi.org/10.1371/journal.pcbi.1014193},
pmid = {41984938},
pmcid = {PMC13095124}
}

RIS

TY - JOUR
AU - Mollard, Sami
AU - Bohte, Sander M.
AU - Roelfsema, Pieter R.
TI - How the visual brain can learn to parse images using a multiscale, incremental grouping process
T2 - PLoS computational biology
J2 - PLoS Comput Biol
PY - 2026
DA - 2026/04/15
VL - 22
IS - 4
SP - e1014193
SN - 1553-734X
PB - PLOS
DO - 10.1371/journal.pcbi.1014193
UR - https://doi.org/10.1371/journal.pcbi.1014193
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pcbi.1014193",
"type": "article-journal",
"title": "How the visual brain can learn to parse images using a multiscale, incremental grouping process",
"container-title": "PLoS computational biology",
"author": [
{
"family": "Mollard",
"given": "Sami"
},
{
"family": "Bohte",
"given": "Sander M."
},
{
"family": "Roelfsema",
"given": "Pieter R."
}
],
"container-title-short": "PLoS Comput Biol",
"volume": "22",
"issue": "4",
"page": "e1014193",
"DOI": "10.1371/journal.pcbi.1014193",
"PMID": "41984938",
"PMCID": "PMC13095124",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pcbi.1014193",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
15
]
]
}
}

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.1371/journal.pbio.3003831 [code]
Disinhibitory signaling enables flexible coding of top-down information in cortical networks.
Journal: PLoS biology
In common: PyTorch, SciPy, NumPy, 6 references
[2] doi:10.1371/journal.pbio.3003915 [code]
Noise-invariant representations of sound emerge along the canonical cortical hierarchy.
Journal: PLoS biology
In common: scikit-image, PyTorch, SciPy, 1 other tool, 3 references
[3] doi:10.1038/s41598-026-45730-8 [code]
Gestalt laws enhance the representation of figures over backgrounds in the visual cortex and influence contrast perception.
Journal: Scientific reports
In common: non-human primate, 4 references
[4] doi:10.1038/s41467-026-72146-9 [code]
Modeling attention and binding in the brain through bidirectional recurrent gating.
Journal: Nature communications
In common: PyTorch, SciPy, NumPy, 3 references
[5] doi:10.1038/s42003-026-10418-2 [code]
Cortical PV and VIP interneurons similarly influence SST neuron output despite distinct unitary properties.
Journal: Communications biology
In common: 5 references
[6] doi:10.1038/s41467-026-72619-x [code]
An inhibitory brainstem pathway reduces visual detection during background motion.
Journal: Nature communications
In common: 4 references
[7] 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, PyTorch, SciPy, 1 other tool, 1 reference
[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: scikit-image, SciPy, NumPy, 2 references
[9] doi:10.1038/s41467-026-71331-0 [code]
A multimodal approach for visualizing and identifying electrophysiological cell types in vivo.
Journal: Nature communications
In common: PyTorch, SciPy, NumPy, 2 references
[10] doi:10.1038/s41467-026-71918-7 [code]
Developmental disinhibition gates language lateralization in childhood.
Journal: Nature communications
In common: PyTorch, SciPy, NumPy, 2 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.