about
Parallelizing non-linear sequential models over the sequence length (arxiv.org)
57 points by famouswaffles on Sep 22, 2023 | hide | past | pdf | 12 comments on HN

In plain words: Sequential models like recurrent networks process one step at a time; this algorithm spreads the work across the sequence so GPUs can run many steps at once, needing no special architecture. It runs up to 1,000 times faster with identical outputs and training results.

Abstract

Sequential models, such as Recurrent Neural Networks and Neural Ordinary Differential Equations, have long suffered from slow training due to their inherent sequential nature. For many years this bottleneck has persisted, as many thought sequential models could not be parallelized. We challenge this long-held belief with our parallel algorithm that accelerates GPU evaluation of sequential models by up to 3 orders of magnitude faster without compromising output accuracy. The algorithm does not need any special structure in the sequential models' architecture, making it applicable to a wide range of architectures. Using our method, training sequential models can be more than 10 times faster than the common sequential method without any meaningful difference in the training results. Leveraging this accelerated training, we discovered the efficacy of the Gated Recurrent Unit in a long time series classification problem with 17k time samples. By overcoming the training bottleneck, our work serves as the first step to unlock the potential of non-linear sequential models for long sequence problems.

Yi Heng Lim, Qi Zhu, Joshua Selfridge, Muhammad Firmansyah Kasim
arXiv:2309.12252 · cs.LG, cs.DC, physics.comp-ph · submitted Sep 21, 2023 · updated Jan 16, 2024
abstract · pdf · html

add comment on HN

Very interesting, but after a superficial quick read, it looks like the proposed parallelization method has O(n³) time complexity and O(n²) space complexity, with n being the number of sequential time steps, i.e., the number of tokens in generative language models.

If that's right, it means that in practice the proposed parallelization method will likely be much slower and much less efficient than modern implementations of self-attention, which have O(n²) time complexity and O(n) space complexity (for example, with FlashAttention). Ouch.

og_kalu, have you had a chance to look at this closely or tinker with it?

The author of the paper here. The cubic time and quadratic space complexity is with respect to the number of dimensions (n), not the sequential time steps. The time and space complexity w.r.t. the number of time steps (L) are both linear, specifically O(Ln^3) time and O(Ln^2) space. The method gives larger speed up with longer sequence (>1k time steps) but with relatively small number of dimensions (<= 64 dimensions from our experiments).
Thank you for posting on HN and clarifying that. I was so off-the-mark! I'm used to the convention in most deep-learning papers of using N for time steps and D for number of dimensions.

Your work looks more interesting to me now, even though cubic time and quadratic space in the number of dimensions are still a significant drag.

Consider: State-of-the-art models often work with dimensions 2-3 orders of magnitude greater than 64. For example, LLaMA 2 models operate on visible and hidden states on 4096 and 11008 dimensions, respectively.

Anyway, thank you again! I'm adding your paper to my reading list.

Yes, sorry for that. I'm coming from physics, so L for sequence length and n for the number of dimensions make more sense :D

I agree with the cubic time and quadratic space is a big limitation for now and I'm looking for ways to make them linear (or close to linear).

> I'm coming from physics

I figured as much :-) because using n and m for vector and linear map dimensions is actually the older, more established convention.

> I'm looking for ways to make them linear (or close to linear)

The holy grail in AI research right now. I imagine you're looking, or have looked, at mapping models to frequency space to make complexity O(n log n). Take a look at the work the Hazy Research folks have been doing at Stanford -- if you haven't already.

Can you share a link to your code repository? I want to play with your algorithm and observe the speedup on my own. Also I'm a big fan of JAX!
We're preparing (i.e. cleaning up) the code for the repo. Will update you when we release the code.
Yeah this isn't a faster alternative to current options without recurrence.

If I have the method correct, it looks like they set constraints and try to guess the computation performed with Newton's method. The guessing can be parallelized.

Technically it's not guaranteed to converge and non trivial computations may not be faster than sequential methods(or may not be reached at all).

Yes, there is no convergence guarantee, but what we found is that typical untrained RNN units (e.g., GRU, LSTM, or a simple MLP for NeuralODE) can converge within 3-5 iterations which gives them a huge speed up over the sequential method. The non-convergence typically happens after many thousand steps of training, but that can be addressed by saving the RNN output from the previous training step as the initial guess for the next training step.

Adding more context from the paper. Although there is no convergence guarantee in forward calculation, the gradient computation only requires 1 iteration and always converge (see section 3.1.1), so even though the forward calculation still uses sequential method, the acceleration in backward computation might be achieved with our method.

Thank you. That's consistent with my (very quick, superficial) read too.

Technically, modern DNNs (transformers, CNNs, RNNs, etc.) don't have convergence guarantees either.

We've just gotten used to SGD somehow always working! :-P

This could be the step forward in compute efficiency we need.