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 Memory比L2快,不代表可以把所有数据都塞进去;它更适合做局部块内复用。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 信息