llm-algo-leetcode |Task01

03_GPU_Architecture_and_Memory.ipynb

【扩展】Part01: 11_KV_Cache_and_Memory_Growth.ipynb

复制代码
                   GPU
                    │
        ┌───────────┴───────────┐
        │                       │
     计算部分                 存储部分
        │                       │
 CUDA Core / Tensor Core      Register
                              Shared Memory
                              L2 Cache
                              HBM

GPU为什么快

比如矩阵乘法 C=AB

里面有大量

cij=∑kaikbkjc_{ij} = \sum_k a_{ik} b_{kj}cij=k∑aikbkj

不同的i,j可以大量并行

GPU特别适合矩阵乘法,卷积,Attention,MLP

CUDA Core和Tensor Core是什么

普通 CUDA Core 做的可以粗略理解成:a × b + c

也就是一次标量乘加。

例如:2 × 3 + 4

Tensor Core不一样,专门为矩阵乘法打造

它的设计目标不是:一个数 × 一个数

而是:

一小块矩阵 × 一小块矩阵 + 一小块矩阵

D=AB+CD=AB+CD=AB+C

MMA: matrix multiply-accumulate

Transformer 里面最多的东西是什么?

复制代码
Q = XW_Q
K = XW_K
V = XW_V

QK^T

MLP:
XW_1
XW_2

Tensor Core 是现代 GPU 跑 Transformer 特别快的一个核心原因。

为什么 A100 → H100 算力暴涨,LLM 不一定同比变快?

假设 GPU:计算速度 = 1000 个数 / 秒

但是内存:只能提供 100 个数 / 秒

那计算单元再快也没有用。

因为它会变成:

复制代码
Tensor Core:
算完了
   ↓
等数据
   ↓
等数据
   ↓
等数据

所以 GPU 不仅需要强大的计算单元,还需要一个非常重要的东西:

Memory Hierarchy,内存层级。

GPU 的内存不是只有"显存"

平时 nvidia-smi 看到:

复制代码
A100 80 GB

这个 80 GB 基本说的是:HBM / Global Memory

但 GPU 内部还有很多更快、更小的存储。

从快到慢:

复制代码
最快、最小
   ↑
Register
   │
Shared Memory
   │
L2 Cache
   │
HBM
   ↓
最慢、最大

Shared Memory 可以理解成:GPU 芯片上的一小块高速 SRAM。

复制代码
Shared Memory ≈ 19 TB/s
HBM           ≈ 1.5 TB/s

GPU 优化里面非常重要的一招就是:

复制代码
大数据在 HBM
      ↓
取一小块
      ↓
Shared Memory
      ↓
疯狂重复计算
      ↓
再换下一块

这个"小块"就是:Tile。

L2 Cache 是干什么的?

L2 在中间:

复制代码
Shared Memory
      ↓
     L2
      ↓
     HBM

它是所有 SM 共享的一层缓存。

作用可以简单理解成:如果某份数据刚从 HBM 拿过,而且之后可能还会用,就暂时留在这里。

时间=数据量带宽时间=\frac{数据量}{带宽}时间=带宽数据量

复制代码
latency_ns = (1024 / bandwidth) * 1e9
python 复制代码
import torch
from typing import Dict

def bytes_to_gb(bytes_val: float) -> float:
    return bytes_val / 1e9

# GPU 内存层级的带宽(字节/秒)
MEMORY_BANDWIDTH = {
    'shared_memory': 19e12,  # 19 TB/s (A100)
    'l2_cache': 1.5e12,      # 1.5 TB/s
    'hbm': 1.5e12,           # 1.5 TB/s (A100)
}

def analyze_memory_hierarchy() -> Dict[str, Dict[str, float]]:
    """
    分析 GPU 内存层级的性能特性。
    """
    result = {}
    
    for mem_type, bandwidth in MEMORY_BANDWIDTH.items():
        # ==========================================
        # TODO 1.1: 计算访问延迟(访问 1 KB 数据)
        # latency_ns = (1024 / bandwidth) * 1e9
        # ==========================================
        latency_ns = (1024 / bandwidth) * 1e9
        
        result[mem_type] = {
            'bandwidth_tb_s': bandwidth / 1e12,
            'latency_ns': latency_ns,
        }
    
    return result

# 测试
mem_analysis = analyze_memory_hierarchy()
for mem_type, stats in mem_analysis.items():
    print(f"{mem_type:15s}: {stats['bandwidth_tb_s']:6.1f} TB/s, Latency: {stats['latency_ns']:6.2f} ns")

真实 latency 还涉及请求调度、cache hit、并发访问等。

Attention

为什么 Attention 是 O(N²)?

Attention score:

复制代码
        K1 K2 K3 K4
Q1      •  •  •  •
Q2      •  •  •  •
Q3      •  •  •  •
Q4      •  •  •  •

为什么标准Attention会迅速变成OOM热点

python 复制代码
def calculate_attention_vram(
    seq_len: int,
    num_heads: int,
    head_dim: int,
    dtype_bytes: int = 2
) -> Dict[str, float]:
    """
    计算标准 Attention 的显存占用。
    """
    # ==========================================
    # TODO 2.1: 计算 Q、K、V 的显存占用
    # qkv_vram = 3 * seq_len * num_heads * head_dim * dtype_bytes
    # ==========================================
    qkv_vram = 3 * seq_len * num_heads * head_dim * dtype_bytes
#QKV三份 所以需要乘以3    
# 如果FP16 那么 dtype_bytes = 2
    # ==========================================
    # TODO 2.2: 计算 Attention 矩阵的显存占用
    # attention_matrix_vram = seq_len * seq_len * dtype_bytes
    # ==========================================
    attention_matrix_vram = seq_len * seq_len * dtype_bytes
    
    # ==========================================
    # TODO 2.3: 计算输出的显存占用
    # output_vram = seq_len * num_heads * head_dim * dtype_bytes
    # ==========================================
    output_vram = seq_len * num_heads * head_dim * dtype_bytes
    
    # ==========================================
    # TODO 2.4: 总显存占用
    # total_vram = qkv_vram + attention_matrix_vram + output_vram
    # ==========================================
    total_vram = qkv_vram + attention_matrix_vram + output_vram
    
    return {
        'qkv': bytes_to_gb(qkv_vram),
        'attention_matrix': bytes_to_gb(attention_matrix_vram),
        'output': bytes_to_gb(output_vram),
        'total': bytes_to_gb(total_vram),
    }

# 测试
print("标准 Attention 显存占用:")
for seq_len in [512, 4096, 128*1024]:
    vram = calculate_attention_vram(seq_len, 32, 128)
    print(f"  seq_len={seq_len:6d}: {vram['total']:10.2f} GB")

FP浮点数

BF更强调能表示更大的数/更小的数

复制代码
Q, K
 ↓
QKᵀ
 ↓
S
 ↓
Softmax(S)
 ↓
P
 ↓
P × V
 ↓
Output

算 S
↓
把 S 写到 HBM

从 HBM 读 S
↓
Softmax
↓
把 P 写到 HBM

再从 HBM 读 P
↓
读 V
↓
矩阵乘法

Q,K→读 HBMQKT=S→写 HBMS→读 HBMSoftmax(S)→写 HBMP→读 HBMPVQ,K \xrightarrow{\text{读 HBM}} QK^T=S \xrightarrow{\text{写 HBM}} S \xrightarrow{\text{读 HBM}} \mathrm{Softmax}(S) \xrightarrow{\text{写 HBM}} P \xrightarrow{\text{读 HBM}} PVQ,K读 HBM QKT=S写 HBM S读 HBM Softmax(S)写 HBM P读 HBM PV

Flash Attention

复制代码
          HBM
           │
      取 Q/K/V 小块
           ↓
         SRAM
           │
         QKᵀ
           │
       Softmax
           │
          × V
           │
     都在 SRAM 里
           ↓
     最终结果写 HBM

QiKjT→Softmax 的局部统计更新→乘 Vj→累加输出Q_iK_j^T \rightarrow \text{Softmax 的局部统计更新} \rightarrow \text{乘 }V_j \rightarrow \text{累加输出}QiKjT→Softmax 的局部统计更新→乘 Vj→累加输出

FlashAttention 使用在线算法,分块处理过程中维护类似:

text 复制代码
当前最大值 m
当前指数和 l
当前输出累积值

所以:

复制代码
处理 block 1
↓
更新 m、l


处理 block 2
↓
修正 m、l


处理 block 3
↓
继续更新

最后得到与正常 Softmax 对应的结果,而无需保存完整:

N×NN×NN×Nscore matrix。

这里最关键的是 Online Softmax。因为 Softmax 本来需要一整行:

softmax(xi)=exi∑jexj \mathrm{softmax}(x_i)= \frac{e^{x_i}}{\sum_j e^{x_j}} softmax(xi)=∑jexjexi

看起来必须先获得整行 (S) 才能算,但 FlashAttention 会维护当前已经看到元素的:

m=当前最大值,l=当前指数和 m=\text{当前最大值},\qquad l=\text{当前指数和} m=当前最大值,l=当前指数和

每来一个新 block,就更新 (m,l) 和输出累加值。因此即使一行 Attention 被拆成很多块,也能最终得到和普通 Softmax 等价的结果,不需要保存完整的一行 (S)。

传统 Attention 是"算一个大中间矩阵 → 写回显存 → 再读回来继续算";FlashAttention 是"把数据分块搬进片上 SRAM → 中间结果就地计算和归约 → 只把必要的最终结果写回 HBM"。

复制代码
Q/K/V
+
Online Softmax 少量状态
+
Output

FlashAttention的核心是IO优化

python 复制代码
def calculate_flash_attention_vram(
    seq_len: int,
    num_heads: int,
    head_dim: int,
    dtype_bytes: int = 2
) -> Dict[str, float]:
    """
    计算 FlashAttention 的显存占用。
    """
    # ==========================================
    # TODO 3.1: 计算 Q、K、V 的显存占用
    # qkv_vram = 3 * seq_len * num_heads * head_dim * dtype_bytes
    # ==========================================
    qkv_vram = 3 * seq_len * num_heads * head_dim * dtype_bytes
    
    # ==========================================
    # TODO 3.2: FlashAttention 只需存储 Online Softmax 的中间值
    # online_softmax_vram = seq_len * num_heads * 2 * 4
    # ==========================================
    online_softmax_vram = seq_len * num_heads * 2 * 4
    
    # ==========================================
    # TODO 3.3: 计算输出的显存占用
    # output_vram = seq_len * num_heads * head_dim * dtype_bytes
    # ==========================================
    output_vram = seq_len * num_heads * head_dim * dtype_bytes
    
    # ==========================================
    # TODO 3.4: 总显存占用
    # total_vram = qkv_vram + online_softmax_vram + output_vram
    # ==========================================
    total_vram = qkv_vram + online_softmax_vram + output_vram
    
    return {
        'qkv': bytes_to_gb(qkv_vram),
        'online_softmax': bytes_to_gb(online_softmax_vram),
        'output': bytes_to_gb(output_vram),
        'total': bytes_to_gb(total_vram),
    }

# 测试
print("\nFlashAttention 显存占用:")
for seq_len in [512, 4096, 128*1024]:
    vram_std = calculate_attention_vram(seq_len, 32, 128)
    vram_flash = calculate_flash_attention_vram(seq_len, 32, 128)
    ratio = vram_std['total'] / vram_flash['total'] if vram_flash['total'] > 0 else float('inf')
    print(f"  seq_len={seq_len:6d}: 标准={vram_std['total']:10.2f} GB, Flash={vram_flash['total']:10.2f} GB, 节省={ratio:8.0f}x")

def test_gpu_memory_practice():
    mem = analyze_memory_hierarchy()
    assert 'shared_memory' in mem and 'hbm' in mem
    assert mem['shared_memory']['bandwidth_tb_s'] > mem['hbm']['bandwidth_tb_s']

    attn = calculate_attention_vram(512, 32, 128)
    flash = calculate_flash_attention_vram(512, 32, 128)
    assert attn['total'] > flash['total']
    assert attn['qkv'] > 0 and flash['online_softmax'] > 0
    print('✅ 03 GPU Architecture and Memory tests passed')

test_gpu_memory_practice()

多卡时为什么会被通信卡住

PCIe 和 NVLink

如果模型太大,需要:

复制代码
GPU 0
GPU 1
GPU 2
GPU 3

那又出现一个问题:GPU 之间怎么交换数据?

普通方式:PCIe

高速专用方式:NVLink

Notebook 给的示意数量级:

复制代码
PCIe Gen4 双向 ≈ 64 GB/s
H100 NVLink 总双向 ≈ 900 GB/s

例如 Tensor Parallel:

复制代码
GPU 0 算一部分矩阵
GPU 1 算一部分矩阵
       ↓
需要交换结果

于是整个系统可能变成:

复制代码
计算很快
↓
等通信
↓
计算
↓
等通信

这时候就不再是:Memory Bound

而可能是:Communication Bound

第一种问题:

复制代码
Tensor Core
█████      █████
     ↑
   等 HBM
 Memory Bound

第二种问题:

复制代码
Tensor Core
一直很忙
Compute Bound

第三种问题,多 GPU:

复制代码
GPU0 █████    █████
           ↑
       等 GPU1 数据
Communication Bound

Tensor Core
    ↓
矩阵计算很快

HBM
    ↓
容量大但相对计算速度不够快

Shared Memory / SRAM
    ↓
小但非常快

Tiling
    ↓
把大问题切成能放 SRAM 的小块

Arithmetic Intensity
    ↓
每搬一个 Byte 能做多少计算

Memory Bound
    ↓
数据搬运速度决定性能

FlashAttention
    ↓
利用 tiling + SRAM + online softmax
减少 HBM 往返

NVLink
    ↓
多 GPU 之间的高速通信
python 复制代码
def pcie_vs_nvlink(payload_mb, pcie_gbps=64, nvlink_gbps=900):
    # 带宽差异真正影响的是把一块数据搬过去要花多少时间。
    pcie_ms = payload_mb * 8 / pcie_gbps
    nvlink_ms = payload_mb * 8 / nvlink_gbps
    return {'pcie_ms': round(pcie_ms, 2), 'nvlink_ms': round(nvlink_ms, 2), 'speedup': round(pcie_ms / nvlink_ms, 1)}

for payload in [64, 256, 1024]:
    print(payload, 'MB ->', pcie_vs_nvlink(payload))
print('higher bandwidth only matters when the transfer is on the critical path')

PCIe

PCIe 可以先理解成:计算机里面通用的高速总线。

GPU、网卡、SSD 等设备都可以通过 PCIe 和系统连接。典型结构可以粗略画成:

复制代码
           CPU
            │
          PCIe
            │
      ┌─────┴─────┐
      │           │
    GPU 0       GPU 1

或者:

复制代码
             CPU
              │
        PCIe Switch
        /     |     \
     GPU0   GPU1   GPU2

所以 PCIe 本来的设计目的不是:"专门让 GPU 和 GPU 疯狂交换几百 GB 数据。"而是一个通用互连标准。

问题是 GPU 本身的计算和显存带宽已经非常高。

比如你前面看到:HBM:TB/s 级别

而你给的材料用 PCIe Gen4 双向约:64 GB/s

说明数量级差距。

复制代码
GPU 内部搬数据:
██████████████████ 非常快

GPU 间 PCIe:
██ 相对慢

因此如果模型频繁交换 tensor:

复制代码
算
↓
等 PCIe
↓
算
↓
等 PCIe
↓
算
↓
等 PCIe

GPU 就大量时间没在真正计算。

所以 NVIDIA 做了 NVLink,NVLink 可以粗略理解为:专门为 NVIDIA GPU 之间高速交换数据设计的互连。

复制代码
PCIe:
通用道路

NVLink:
GPU 之间的高速专线

NVLink 为什么更快?因为它从设计开始就是针对:GPU ↔ GPU这种大规模数据交换。

材料列的数量级是:

复制代码
PCIe Gen4
≈ 64 GB/s 双向
A100 NVLink
≈ 600 GB/s 总双向
H100 NVLink
≈ 900 GB/s 总双向

一张 GPU 可以拥有多条 NVLink,把所有链路聚合起来看总通信能力,可以达到很高的带宽。

那 NVSwitch 又是什么?这又往前走一步。假设有 8 张 GPU。如果只是两两 NVLink:

复制代码
GPU0 ─ GPU1
GPU2 ─ GPU3
GPU4 ─ GPU5
GPU6 ─ GPU7

那么:GPU0 想和 GPU7 通信。就不一定方便。

于是引入:NVSwitch

你可以把它理解成:专门连接多张 NVIDIA GPU 的高速交换机。

复制代码
GPU0 ─┐
GPU1 ─┤
GPU2 ─┤
GPU3 ─┤
      ├── NVSwitch
GPU4 ─┤
GPU5 ─┤
GPU6 ─┤
GPU7 ─┘

于是任何 GPU:GPU0

都可以很高效地和:

复制代码
GPU1
GPU2
GPU3
...
GPU7

交换数据。

⚠️ 常见误区

  • Shared MemoryL2 快,不代表可以把所有数据都塞进去;它更适合做局部块内复用。
  • HBM 带宽已经很高,不代表就不会 Memory Bound;在高算力 GPU 上,带宽反而更容易成为瓶颈。
  • FlashAttention 主要减少的是 HBM 访问,不是把主要计算量"变没了"。
  • NVLink 很快,但仍然需要正确的通信库、拓扑和并行策略配合,否则并不会自动接近跑满。

练习

task 1 小测验

【题目 1】 B

关于 Transformer 自回归推理中的 KV-Cache 机制,下列说法正确的是:

A. KV-Cache 仅在模型训练阶段使用,推理时无需启用

B. KV-Cache 为每一层缓存已生成 token 的 Key 和 Value 向量,避免重复计算

C. KV-Cache 占用的显存大小与序列长度无关,只取决于模型隐层维度

D. 每次生成新 token 时,KV-Cache 需要重新计算所有历史 token 的 Key 和 Value

KV Cache 缓存每一层历史 token 的 K、V,新 token 生成时不用重新计算历史 K/V

【题目 2】 B

在大型语言模型推理中,Paged Attention 机制的主要设计目标是:

A. 将注意力分数按页划分,减少矩阵乘法的计算量

B. 将 KV-Cache 划分为固定大小的物理块(page),通过非连续存储实现按需分配,减少显存浪费和碎片

C. 利用分页技术将 KV-Cache 从 GPU 显存交换到 CPU 内存,以支持无限长序列

D. 将输入序列分页,使每个 token 只与相邻页内的 token 计算注意力

PagedAttention 把 KV Cache 分成固定大小的 block/page,允许非连续存储 + 按需分配,减少显存碎片和浪费

【题目 3】 C

关于 NVIDIA GPU 的内存层次结构,以下描述正确的是:

A. 全局内存(Global Memory)的访问延迟最低,且具有最高的带宽

B. 共享内存(Shared Memory)位于 GPU 芯片外的 DRAM 中,容量较大

C. L1 数据缓存与共享内存通常位于同一个物理单元(SM 内),两者容量可配置分配

D. 寄存器文件(Register File)的总容量远大于全局内存,且对所有线程可见

Shared Memory 和 L1 Cache 都位于 Shared Memory 内部的片上存储体系,现代 NVIDIA GPU 上通常共享/划分相关片上资源

复制代码
SM
└── 一块较大的片上高速存储资源
       │
       ├── 一部分给 Shared Memory
       │
       └── 一部分给 L1 Cache
┌────────────────────────────┐
│            SM              │
│                            │
│   Register File            │
│                            │
│   ┌──────────┬──────────┐  │
│   │ Shared   │ L1 Cache │  │
│   │ Memory   │          │  │
│   └──────────┴──────────┘  │
│      Unified Data Cache    │
└────────────────────────────┘
              │
              ↓
           L2 Cache
              │
              ↓
        Global Memory

【题目 4】C

在 GPU 高性能计算编程(如 CUDA)中,对内核(kernel)性能影响最大且通常需要程序员显式手动优化的内存类型是:

A. 全局内存(Global Memory)

B. 纹理内存(Texture Memory)

C. 共享内存(Shared Memory)与寄存器(Registers)

D. 常量内存(Constant Memory)

Shared Memory 和 Registers 是 CUDA kernel 优化中程序员需要重点控制的片上高速存储


KV cache and Memory Growth

KV Cache 为什么越来越大?

先只考虑 一层、一个 KV head

假设已经有 4 个 token:

复制代码
KV Cache
token1 → K1  V1
token2 → K2  V2
token3 → K3  V3
token4 → K4  V4

生成 token5:

复制代码
KV Cache
token1 → K1  V1
token2 → K2  V2
token3 → K3  V3
token4 → K4  V4
token5 → K5  V5    ← 新增

注意:每生成一个 token,就必须多保存一份 K 和 V

所以:

复制代码
token 数:
1  2  3  4  5  6 ...

对应:

复制代码
KV Cache:
小 → → → → → → 大

KV Cache Bytes≈2×L×B×Hkv×D×S \text{KV Cache Bytes} \approx 2 \times L \times B \times H_{kv} \times D \times S KV Cache Bytes≈2×L×B×Hkv×D×S

其中:

  • LLL 是层数

  • BBB 是 batch size

  • HkvH_{kv}Hkv 是 KV 头数

  • DDD 是 head dim

  • SSS 是上下文长度

  • 前面的 222 表示同时存 K 和 V

    Transformer:

    Layer 1
    Layer 2
    Layer 3
    ...
    Layer 32

    每一层都有自己的 Attention。

python 复制代码
def kv_cache_bytes(seq_len, num_layers, num_kv_heads, head_dim, batch_size=1, dtype_bytes=2):
    return 2 * seq_len * num_layers * num_kv_heads * head_dim * batch_size * dtype_bytes

examples = [(1024, 32, 32, 128), (2048, 32, 32, 128), (4096, 32, 32, 128)]
for seq_len, layers, kv_heads, head_dim in examples:
    size_gb = kv_cache_bytes(seq_len, layers, kv_heads, head_dim) / 1e9
    print(f"seq_len={seq_len:4d} -> KV cache ≈ {size_gb:5.2f} GB")

MHA / GQA / MQA 为什么能改变 KV Cache 大小?

  • MHA (Multi-Head Attention):每个 query head 都有自己对应的一组 K/V,缓存压力最大。
  • MQA (Multi-Query Attention):多个 query head 共享同一组 K/V,KV cache 立刻变小。
  • GQA (Grouped-Query Attention):介于 MHA 和 MQA 之间,把 query heads 分组共享 K/V,在显存和表达能力之间做折中。

从缓存角度看,真正决定显存大小的不是 query heads 有多少,而是 要存多少组 K/V。所以只要 KV 头数下降,缓存就会按比例下降。

这也是为什么很多长上下文模型会采用 MQA 或 GQA:它们不只是"改了 attention 的形式",而是在直接压低推理时的 KV cache 成本。

python 复制代码
def kv_cache_gb(seq_len, num_layers, num_kv_heads, head_dim, batch_size=1, dtype_bytes=2):
    return kv_cache_bytes(seq_len, num_layers, num_kv_heads, head_dim, batch_size, dtype_bytes) / 1e9

seq_len = 4096
num_layers = 32
head_dim = 128
for name, kv_heads in [("MHA", 32), ("GQA", 8), ("MQA", 1)]:
    print(f"{name:>3s}: kv_heads={kv_heads:2d}, KV cache ≈ {kv_cache_gb(seq_len, num_layers, kv_heads, head_dim):5.2f} GB")

PagedAttention 和 MLA 又分别怎么解决 KV Cache 问题?

因为它们两个都可以说:"缓解 KV Cache 显存问题。"。但是解决的不是同一个问题。Notebook 也特意把它们分成"缓存管理"和"表示压缩"两个方向。

  • PagedAttention 主要解决的是 缓存分配和访问组织 问题。

    • 它把 KV cache 按页组织,避免长序列和多请求场景里出现连续大块显存分配困难。
    • 它的重点不是把 K/V 表示本身压缩掉,而是让缓存的存储、搬运和复用更稳定。
  • MLA (Multi-Head Latent Attention) 主要解决的是 表示压缩 问题。

    • 它把原本需要长期保存的 KV 表示压到更低维的潜变量空间里。
    • 这样做的核心收益是直接降低每个 token 需要保留的缓存体积。

可以把它们理解成两种不同方向的优化:

  • PagedAttention 是在优化"怎么管 cache"。
  • MLA 是在优化"cache 本身有多大"。

前者偏系统实现,后者偏表示结构。两者都在缓解长上下文下的显存压力,但切入点不同。

复制代码
PagedAttention
      ↓
"东西这么多我承认,
 但我想办法把仓库管理得更好。"

MLA
      ↓
"我直接让每件东西体积更小。"
python 复制代码
def paged_attention_pages(seq_len, page_size):
    return (seq_len + page_size - 1) // page_size

# Multi-Head Latent Attention
#把原本需要长期保存的 KV 表示压缩成更低维的 latent representation。
def mla_cache_bytes(seq_len, num_layers, latent_dim, batch_size=1, dtype_bytes=2):
    return seq_len * num_layers * latent_dim * batch_size * dtype_bytes

seq_len = 4096
page_size = 128
print(f"PagedAttention pages: {paged_attention_pages(seq_len, page_size)}")
print(f"MLA cache example: {mla_cache_bytes(seq_len, 32, 64) / 1e9:.2f} GB (latent_dim=64)")
  • KV cache 不是只和 token 数有关,它还和层数、batch size、KV 头数一起增长。

  • MQA / GQA 不是单纯改名字,而是在实打实地压低缓存体积。

  • PagedAttention 解决的是缓存管理和碎片化,不等于表示压缩。

  • MLA 解决的是表示体积,不等于把调度和分配问题也一并解决。

    复制代码
                 自回归生成
                      │
                历史 K/V 不能丢
                      │
                      ↓
                   KV Cache
                      │
          ┌───────────┼───────────┐
          ↓           ↓           ↓
        seq_len    batch       layers
          ↑           ↑           ↑
          └────── 都会线性放大 ────┘
                      │
                      ↓
               num_kv_heads
                      │
          ┌───────────┼──────────┐
          ↓           ↓          ↓
         MHA         GQA        MQA
         32           8          1
       最大          中间        最小
    
                      ↓
    
          KV Cache 还是很难管理?
                      │
                      ↓
              PagedAttention
              怎么管理 cache
    
          KV Cache 本身还是太大?
                      │
                      ↓
                   MLA
              cache 表示压缩

MLA核心确实是 low-rank KV joint compression(K/V 联合低秩压缩)

复制代码
             Down Projection
h_t ─────────────────────────→ c_t^KV
                                │
                       ┌────────┴────────┐
                       ↓                 ↓
                     W_UK              W_UV
                       ↓                 ↓
                      K_t               V_t

不用给每个历史 C_i 单独恢复完整 K_i。

不用给每个历史 token 单独恢复完整 V。


复制代码
第一步
K/V 联合低秩压缩

h_t
 ↓
c_t^KV
只缓存 latent

然后:

复制代码
第二步
利用矩阵结合律 / 权重吸收

避免 decode 时
把所有 latent
完整恢复为 K/V

512 维只够"存",但 Attention 最终还是需要多个 head 各自的 K/V 表示来参与计算。

复制代码
存储阶段:只保存 512 维 latent
计算阶段:从这 512 维 latent 变换出各个 head 需要的表示

数学定义里有 Up Projection,但真正计算时可以利用矩阵结合律,把它"挪到另一边",于是不用真的把每个历史 token 的 512 维 latent 展开成完整 K/V。

利用我们前面提到的 weight absorption.把部分 up-projection 权重吸收到 Query/输出侧。

实际高效 decode 可以变形为:

复制代码
Q
 ↓ 先变换
直接和历史 cKV latent
做核心 Attention 计算
复制代码
第 t 个 token
    │
    ↓
Token Embedding
    │
    ↓
前面的 Transformer Layers
    │
    ↓
当前第 l 层的 hidden state
h_t^(l)
    │
    ├───────────────────────────────┐
    │                               │
    │                               │
    ↓                               ↓
Query 路径                       KV 路径
    │                               │
    ↓                               ↓
Q Projection                  Down Projection
    │                         W_DKV
    ↓                               │
Q_t                               ↓
                              c_t^KV
                           低维 latent
                           例如 512 维
                                │
                                │
                 ┌──────────────┴──────────────┐
                 │                             │
                 ↓                             ↓
          Up Projection                 Up Projection
              W_UK                          W_UV
                 │                             │
                 ↓                             ↓
          K 的内容表示                     V 表示
              K_t^C                         V_t
                 │
                 │
                 │      另外还有位置编码分支
                 │              │
                 │              ↓
                 │          RoPE Key
                 │            K_t^R
                 │              │
                 └───────┬──────┘
                         ↓
                    完整 Key
                 [K_t^C ; K_t^R]

普通 MHA 会缓存:

复制代码
历史 token 1 → K1 + V1
历史 token 2 → K2 + V2
历史 token 3 → K3 + V3
...
历史 token t → Kt + Vt

MLA 则主要缓存:

复制代码
历史 token 1 → c1^KV + 少量 RoPE 信息
历史 token 2 → c2^KV + 少量 RoPE 信息
历史 token 3 → c3^KV + 少量 RoPE 信息
...
历史 token t → ct^KV + 少量 RoPE 信息
相关推荐
地平线开发者1 小时前
征程6|YOLOv5x 在 Horizon 征程6 上的端到端部署实践(下)
算法·自动驾驶
艾为电子2 小时前
【应用方案】电视沉浸式音频升级: 电视音频 awinic“芯片 + 算法” 一体化解决方案
算法·音视频
大熊背2 小时前
树莓派相机自动白平衡详解(二)
算法·白平衡·isppipeline
Scabbards_2 小时前
面试Leetcode - 算法合集
算法·leetcode·面试
chuan.bai2 小时前
Java RAG 实战附录:qwen3 与 bge-m3 模型切换指南
java·人工智能·算法
wuyk5552 小时前
7.AVL 树:第一个自平衡二叉搜索树
开发语言·stm32·单片机·算法
程序猿炎义2 小时前
【llm-algo-leetcode学习笔记】显存与性能认知底座
笔记·学习·leetcode
rannn_1112 小时前
【力扣hot100】二叉树专题+总结
java·算法·leetcode·二叉树
啊嘞嘞?2 小时前
力扣(岛屿数量)
算法·leetcode