Md. Asif Uddin

    LectionesPart VII

    Position, memory and the cost of attention

    Attention is quadratic and the cache does not fit. Five papers about that.

    Nothing in this part changes what attention computes. Every one of these papers returns the same numbers a naive implementation would, or very nearly, and the contribution is entirely in where those numbers live and how often they move.

    That is a genuinely different kind of research and it is worth reading a run of it together. The constraint being optimised is not FLOPs — it is memory bandwidth, cache size, and the number of bytes you must keep per token per layer for the whole of a conversation.

    RoPE is the exception: it changes how position enters the model, and it is here because everything about long context depends on it.

    The reading

    1. RoPE

      RoFormer: Enhanced Transformer with Rotary Position Embedding

      Su et al. · Neurocomputing · 2021

      Claim
      Rotating query and key vectors by an angle proportional to position makes attention depend on relative position, with no learned position embeddings at all.
      Why
      In nearly every open model you will run. The rotation is also what explains why extending context by rescaling frequencies works — you cannot reason about that trick without this paper.
      Read
      Section 3.2. Do the two-dimensional case first; the general case is that, repeated.
    2. FlashAttention

      FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness

      Dao, Fu, Ermon, Rudra & Re · NeurIPS · 2022

      Claim
      Tiling the attention computation so it stays in on-chip SRAM makes it two to four times faster and linear in memory, returning exactly the same result.
      Why
      The paper that made reading the hardware a respectable research method. Nothing about the mathematics changes; the entire gain is in avoiding round trips to high-bandwidth memory.
      Read
      Section 3.1 and the memory-hierarchy figure.
    3. MQA

      Fast Transformer Decoding: One Write-Head is All You Need

      Shazeer · arXiv · 2019

      Claim
      Sharing a single key-value head across all query heads collapses the decoding cache and makes generation an order of magnitude faster.
      Why
      Nine pages, one author, and the reason your inference server fits in memory. Read it to understand what the key-value cache actually is before anyone optimises it further.
      Read
      All of it.
    4. GQA

      GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints

      Ainslie et al. · EMNLP · 2023

      Claim
      Grouping query heads over a small number of key-value heads recovers most of multi-query speed without its quality loss, and existing checkpoints can be converted rather than retrained.
      Why
      The compromise that shipped. Llama, Mistral and most of what came after use it, and the conversion recipe is why adoption was immediate.
    5. Ring Attention

      Ring Attention with Blockwise Transformers for Near-Infinite Context

      Liu, Zaharia & Abbeel · ICLR 2024 · 2023

      Claim
      Distributing blocks across devices arranged in a ring, and overlapping communication with computation, makes context length scale with the number of devices.
      Why
      Read it directly after FlashAttention. It is the same tiling argument, moved from one GPU's memory hierarchy to a cluster's.