Paged Attention
Normal attention kernels waste both compute and memory.
Compute Wastage
GPUs work on tensors in uniform shape. To enable continuous batching, we just put all requests as one 2D matrix. The smaller requests are padded to make the matrix uniform. This means some resources are wasted at the end where the actual data doesn't exist.
Memory Wastage
All KV caches have same size. This means, even the smallest request must reserve the same amount of memory.
How paged attention works?
Instead of always allocating one single block of KV cache in GPU memory, we allocate only small blocks of GPU memory. The blocks are added as the context grows.
When the context grows, the KV cache must be reallocated to a larger block of memory. This is a problem because the GPU memory is already fragmented and there may not be a contiguous block