Parallelizing non-linear sequential models over the sequence length
arxiv.org
arxiv.org
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?
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).
Technically, modern DNNs (transformers, CNNs, RNNs, etc.) don't have convergence guarantees either.
We've just gotten used to SGD somehow always working! :-P
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.
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.
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 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.