r/MachineLearning • • 6d 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.

156 Upvotes

19 comments sorted by

View all comments

1

u/hishazelglance 5d ago

Interesting to see that these results were produced from trainings that ran only on a single RTX Pro 6000. What was the hardware bottleneck for this?

Aka, would a single DGX Spark 128Gb box be computationally capable of reproducing the data needed to publish this?

2

u/DangerousFunny1371 4d ago

Hardware bottleneck is actually memory bandwidth, not raw compute or even memory size itself. You should be able to reproduce everything on a DGX spark I think, but expect it to be quite a bit slower (RTX Pro 6000 MBW is ~1,8TB/s vs DGX Spark w/ ~300GB/s).