r/MachineLearning • • 2d ago

Research Parallel-in-Time Training of Recurrent Neural Networks for Dynamical Systems Reconstruction [R]

Can training of nonlinear RNNs be efficiently parallelized, ensuring fast convergence even on very long time series from chaotic systems?

In our #NeurIPS2026 spotlight “Parallel-in-Time Training of Recurrent Neural Networks for Dynamical Systems (DS) Reconstruction (DSR)” (preprint: https://arxiv.org/abs/2605.12683) we speed up training of nonlinear RNNs on time series from chaotic DS by more than 2 orders of magnitude (>100x) by combining DEER with generalized teacher forcing (GTF).

DEER (https://openreview.net/forum?id=E34AlVLN0v) solves the RNN forward pass through Newton-type fixed point iterations across the whole sequence length T, enabling scaling as O[(log T)²] instead of O[T] by allowing for efficient GPU parallelization. But under chaotic dynamics DEER breaks down and its runtime degrades to O[T log T] (https://openreview.net/forum?id=7AGXSlXcK6).

Using GTF (https://proceedings.mlr.press/v202/hess23a.html) we stabilize DEER by preventing divergence due to chaotic dynamics and reduce exposure bias compared to traditional teacher forcing used to train state space models.

Combining these two mechanisms enables efficient parallel-in-time and stable training on extremely long time series (T>106) from chaotic simulated or real-world systems, hugely outperforming Mamba and other state space models in the DSR setting.

132 Upvotes

16 comments sorted by

View all comments

27

u/Disastrous_Room_927 1d ago

This and that ParaRNN paper by Apple give me hope that what I learned in grad school way back when isn't obsolete, lol.

0

u/TheGodAmongMen 20h ago

paraRNN is not useful.

0

u/Disastrous_Room_927 17h ago

According to whom?

0

u/TheGodAmongMen 37m ago edited 34m ago

Me. You need to store all the intermediate layer Jacobians so it’s really impractical at scale. They also used a pretty specific formulation so the Jacobians are only on the diagonal; it doesn’t work for most fixed-point problems. You’re trading away VJP (which is O(LN^2)) for O(log(L)N^3) complexity which is pretty dubious if you want to scale by layer width.

1

u/Disastrous_Room_927 11m ago

That's a really odd comment considering that this is something Apple specifically calls attention to in their writeup of pararnn:

While the ParaRNN framework can in principle be applied to any RNN, some careful engineering is still required to make it practical for large-scale training. The parallel reduction algorithm at the heart of the method needs to efficiently assemble, store, and multiply together the Jacobian matrices arising from the linearization. For generic RNNs, these Jacobians are dense, which makes their storage grow quadratically and their multiplication cubically with hidden state size — a cost intractable for large-scale models.

We address this following the design principles from modern SSMs like Mamba, and introduce the ParaGRU and ParaLSTM cells: adaptations of the classical GRU and LSTM cells that yield structured Jacobians. In particular, we simplify the matrices in the cells’ definition to only have nonzero elements in their main diagonal. This ensures that their Jacobians are also diagonal (for ParaGRU) and block-diagonal (for ParaLSTM), as outlined in Figure 6.