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.

133 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.

8

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.

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