skip to content
← All Projects

Decision Making in Low-Rank Recurrent Neural Networks

Building on Dubreuil et al. (2022)

What this is about

You can train an RNN to solve almost any cognitive task, but it's usually hard to say how it solves it. Low-rank RNNs make that question tractable. If the recurrent weight matrix is constrained to rank R, the activity of every neuron is confined to a small subspace, and the whole network can be described by a handful of latent variables.

Dubreuil et al. (2022) take this one step further. Treat each neuron as a point in a "loading space" defined by its connectivity weights, and the covariance of that point cloud tells you what the latent dynamics will be. In this project I trained 512-neuron low-rank RNNs on two classic tasks, perceptual decision-making (rank one) and parametric working memory (rank two), and asked whether those connectivity statistics were enough to explain, and rebuild, what each network had learned.

Latent trajectories of the rank-one network on the decision task, colored by mean stimulus strength. Positive trials flow to one attractor and negative trials to the other.
The rank-one network on the decision task. Each line is a trial, projected onto the input and recurrent directions and colored by stimulus strength. Evidence accumulates, then the network commits to one of two attractors.

What I found

The decision task worked out cleanly. The rank-one network settles into two stable fixed points with an unstable one at the origin, and three covariances do all the work: one brings the input into the latent, one provides positive feedback, and one lines the latent up with the readout. If you fit a Gaussian to the trained network's connectivity and sample brand new networks from it, they also solve the task with 100% accuracy. A one-dimensional equivalent circuit built from the same covariances behaves just like the full network.

Working memory was messier, and more interesting because of it. The rank-two network learned the task with low error, but not the way the paper describes. Instead of holding on to the first stimulus in a persistent latent, its two latents rotate around each other, and the readout depends on the oscillation being at the right phase when the second stimulus arrives. In other words, it learned to keep time rather than to remember. Change the delay between stimuli and performance swings up and down with it, and networks resampled from its connectivity statistics mostly fail.

Mean squared error of the rank-two network as a function of the delay between stimuli. Error is near zero at the trained delay of 50 steps and oscillates at longer delays.
Error of the rank-two network as the delay between stimuli changes. It was trained at 50 steps; away from that, performance comes and goes with the phase of its internal oscillation.

Plugging the covariances reported in the paper into the equivalent circuit gives the solution I was expecting: one latent acts as a line attractor that holds the first stimulus, the other carries a transient response to the second, and the result doesn't care how long the delay is. My guess is that training with variable delays, as Dubreuil et al. did, would push the network toward that solution instead of the oscillatory shortcut.

How it's built

The model, training loop, and task generators are a small PyTorch package. Only the recurrent connectivity vectors are trained; the input and readout weights stay fixed at their random initialization, so everything interesting about a solution lives in the recurrent structure.

The analysis side covers projecting activity into the latent space, PCA, fitting and resampling the loading-space distribution, simulating the reduced Gaussian circuits, and finding fixed points by minimizing the speed of the dynamics (following Sussillo & Barak, 2013). The whole thing runs end to end from a single notebook with fixed seeds.

Reference

Dubreuil, A., Valente, A., Beiran, M., Mastrogiuseppe, F., & Ostojic, S. (2022). The role of population structure in computations through neural dynamics. Nature Neuroscience, 25, 783–794.