23 ms·
This is extremely wrong. The attention component that's quadratic is a relatively small portion of compute.
by chillee 3y ago
This is extremely wrong.
The attention component that's quadratic is a relatively small portion of compute.
- cochne 3y agoI think they are correct, do you have a source? From my knowledge the only other components are the fully connected networks which are not big contributors.
- lumost 3y agoFor small values of N, the linear terms of the transformer dominate. At the end of the day, a double layer of 764*2048 is still north of 3.1 MM flops/token/layer.
- imtringued 3y agoIt's quadratic, because of the dot product in the attention mechanism. You can use K-V Caching to get rid of a lot of the quadratic runtime that comes from redundant matrix multiplications, but after you have cached everything, you still need to calculate the dot product k_i * q_j with i,j being index of the tokens. With n tokens, you will get O(n*n). But you have to remember that this is only n^2 multiplications. It's not exactly the end of the world at context sizes of 32k, for example. It only gets nasty in the hundred thousands to millions. Here is the source I used: https://sebastianraschka.com/blog/2023/self-attention-from-scratch.html https://sebastianraschka.com/blog/2023/self-attention-from-s...
- chillee 3y agohttps://news.ycombinator.com/item?id=39475528 https://news.ycombinator.com/item?id=39475528 My source is me :) I work at PyTorch on ML compilers. If you don't believe me, perhaps you'll believe Karpathy's diagram (and the general discussion in the thread): https://twitter.com/karpathy/status/1658161721251602432 https://twitter.com/karpathy/status/1658161721251602432
- dTal 3y agoTo anyone doubting this, note that llama.cpp does not slow down by a factor of 16 when you pass -c 2048 instead of -c 512.
- cs702 3y agoI'm the author of the parent comment, and just saw your response. First of all, by "a factor of 100,000×" I mean a model-specific factor, call it c, of 100,000x, but looking at what I wrote I can see how it could be misinterpreted. Second, as sequences become longer, the quadractic component will become more important.
- chillee 3y ago> For a sequence with N = 100,000 tokens, it would mean cost dropping by a factor of 100,000× I'm not sure I understand the intended interpretation of this. Concretely speaking, if it cost CoolAI 100k seconds of compute to process a sequence of length 100k, it would not take them 1 second now. I agree that as sequences become longer, the quadratic component will become more important. But as models get bigger, the attention component also becomes less important. For example, to take a concrete model (say Llama-70B), it takes about 1.4e16 MLP FLOPs (70 billion * 100000 * 2) to process 100k tokens. The attention component takes about 6.5e15 FLOPS (80 [layers] * 100k [sequence length]^2 * 8192 [hidden dim]). So even if attention turned constant it would reduce runtime by about 30% with today's model at 100k sequence length.
- cs702 3y ago> I'm not sure I understand the intended interpretation of this. Concretely speaking, if it cost CoolAI 100k seconds of compute to process a sequence of length 100k, it would not take them 1 second now. Cost would be lower by a factor of sequence length N = 100,000. If you call the factor C, cost would be lower by C × N. For most models and context lengths today, C is below 1. As we continue to increase sequence length N -- say, as we go from 100K to 100M tokens, incorporating multiple modalities -- factor C will increase toward 1 for all Transformers. Again, I see how what I wrote could be misinterpreted. Sorry about that!
- casercaramel144 3y agoHuh? I thought the issue before ringattention is the memory requirement of the softmax layer, since you have to load the whole matrix in at once? It's O(s^2) no? Also hi horace.
- chillee 3y agoWho is this :think: But no, FlashAttention already solved the memory requirements of attention. RingAttention is primarily useful for parallelizing across the sequence component.
- casercaramel144 3y agoIt's camel. How do you do matrix vector attention without keeping the full matrix in cache, surely you don't just load unload it a million times