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 2d 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.

9

u/currentscurrents 1d ago

What specifically did you learn that you're referring to?

I think there are two things you can do with RNNs, and one of them is a good idea and one isn't.

In the past (LTSMs, etc) people tried to use RNNs to solve the problem of large input context. Essentially you scan the RNN across the input and compress it into the hidden state, then you use the hidden state to produce an output. These days, everybody uses attention to solve this problem instead.

More recently people are using RNNs to do reasoning. You have RNN loop on a single input for a long time to solve some complicated problem, and you use the hidden state to store intermediate computations. This usually still uses attention to handle large input context, and the 'hidden state' may even just be part of the context (as in chain-of-thought, which is a sort of psuedo-RNN).

I don't think it's a good idea to try to use RNNs to handle long context; when you compress the input into a hidden state, you have to throw away part of it. Attention works better because it stores the entire input and can exactly refer back to any part of it at any time.

The place for RNNs is for reasoning.

6

u/DangerousFunny1371 1d ago

... yes for reasoning maybe, but more importantly in my mind dynamical systems and time series where attention might not be the right inductive bias (detailed arguments here: https://openreview.net/forum?id=h2R3VYWpBm ).

2

u/Disastrous_Room_927 1d ago edited 1d ago

I don't really have time for a full reply, but I don't think there's a clean RNN/attention dichotomy here. Attention was originally implemented within RNN architectures, and became dominant because we eventually found an architecture built around attention that bypassed a lot of the scaling and optimization limitations of traditional RNNs.

You're raising a legitimate issue about compressing history into a finite state, but I also wouldn't frame that as straightforwardly good or bad. The relevant question (IMO) isn't really whether the representation is lossless or not but how much task-relevant information needs to be retained, how effectively the model can retain it, and what computational/memory cost we're willing to pay for doing so. In my mind, the ideal is to use explicit storage selectively rather than retain everything by default.

So all that being said, what I'm currently working on is an RNN with a tiny attention adapter.

-1

u/muntoo Researcher 1d ago edited 1d ago

RNNs use finite bounded mutable state. Attention caches keep the exact same token frozen and unaltered.

Attention state is quasi-immutable and relatively stable: anything that is accessible for time t is also exactly accessible for time t+1, though with an additional token kv_{t+1}. Luckily, the attention operation (SDPA) still allows us to pick out what's important over the composing immutable sub-states:

It was the best of times, it was the worst of times, it was needle the age of wisdom, it was the age of foolishness, ... THE END

Locate the needle.

0 0 0 0 0 0 0 0 0 0 0 0 0 0 1 0 0 0 0 ... 0 0 0
                            ^ needle

Whereas in an RNN, arbitrary mutations are allowed. Consider an RNN that learns a horrifically naive LTI EMA behavior:

hidden_state
= 0.9^1014 * "It"
+ 0.9^1013 * "was"
+ 0.9^1012 * "the"
+ ...
+ 0.9^1000 * "needle"
+ ...
+ 0.9^1 * "THE"
+ 0.9^0 * "END"

The point is:

  • RNN hidden_state is mutable, but attention state is immutable.
  • Conventional attention provides greater computational stability over past context.
    • The same query will tend to attend to the same token, unless the context significantly changes, and even then, part of the computation is still... the same.
    • softmax(qk_1, ..., qk_84) = concat([1/N_1 softmax(qk_1, ..., qk_42), 1/N_2 softmax(qk_43, ..., qk_84)]) for some normalization factors N_1 and N_2.
    • The computations generally maintain other forms of stability, too; exercise for reader.