acceptodds
Under review as a conference paper at ICLR 2027

When Does Probabilistic Inference Improve Task Learning for Stochastic RNNs?

Abstract

Recurrent neural networks (RNNs) are used both to solve temporal tasks and to model the dynamics of physical and biological systems. In both settings, recurrent dynamics are often subject to noise, but how this stochasticity is handled in learning algorithms differs. Task optimization typically minimizes output error along noisy rollouts of network dynamics using backpropagation through time (BPTT). System identification instead infers latent trajectories conditioned on observed outputs and learns from this posterior, for example, via expectation–maximization (EM). Can such inference-based learning of recurrent weights also benefit task optimization? We compare EM-based posterior-inference-based EM with prior-rollout BPTT training and show that the two coincide in the low-noise limit, in both objective and optimization dynamics, but diverge when noise is non-negligible. In this regime, we identify theoretically and empirically three task families where posterior inference improves task learning: long-delay tasks, where it avoids the exponential, delay-dependent decay of BPTT's learning signal; tasks with distribution shifts, where conditioning latent activity on targets reduces catastrophic interference; and tasks with ambiguous targets, where EM captures temporal structure that output-error minimization ignores. Finally, we show that several alternatives to BPTT can be interpreted as different approximate E-steps within a unified EM framework, suggesting that advances in state-space-model inference may offer principled routes to new BPTT alternatives. Together, these results illuminate when inference-based training is advantageous over prior-rollout training in stochastic learning systems.

Then back it, or bet against it.

Related papers

Open the market on this paper to see 7 more related papers.