Attention: Isn't it quadratic in context length? I dunno, this feels like the crude first iteration of something that will get inevitably passed by something that scales better.
Attention: Isn't it quadratic in context length? I dunno, this feels like the crude first iteration of something that will get inevitably passed by something that scales better.
"""
Complexity is quadratic in sequence length. For 512 tokens it is 262K, but for 4000 tokens it becomes 16M and goes OOM on a single GPU. We need about 100K-1M tokens to load whole books at once.
Since 2017 there have been hundreds of attempts to bring O(N^2) to O(N), but none of them replaced the vanilla attention yet in large models. They lose on accuracy. Maybe Flash attention has a shot (https://arxiv.org/abs/2205.14135).
"""
Keep in mind: Once you go into precision as low as 4 bits (or lower?), all sorts of optimizations can become practical. Off the top of my head, maybe you could cache and reuse common attention sub-matrices (e.g., a 16×16 sub-matrix with 4-bit elements occupies only 16×16×4÷8=128 bytes of space)?
My sense is there's so much money at stake here, that whoever does this first will win big even if they end up having to replace or augment it with something better down the road. Hypothetical example: Imagine Intel or AMD coming out with a $1K or $2K card that has "built-in 4-bit attention," enabling you to run transformers of much greater scale on a run-of-the-mill desktop PC. I'd buy that in a heartbeat.
[a] Here's a recent post about a new approach from a group at Stanford that looks promising to me, although I don't fully understand all its details yet: https://news.ycombinator.com/item?id=35502187
Alternative models like S4 have been able to get transformer level performance with O(N) sequence length scaling.
[0]: https://twitter.com/typedfemale/status/1609867110695735296