80 行 PyTorch 从零写 DeepSeek 的 MLA:量一遍 KV cache、踩一遍 absorption,你才会明白 vLLM 为什么要加专用内核

我把 DeepSeek V2/V3 的 Multi-head Latent Attention(下称 MLA)按论文流程在单卡 RTX 3090 上用 80 行 PyTorch 写了出来,然后做了三件事:第一,对比 cache 体积,MLA 比同规模 MHA 小 57 倍第二,验证 cache 正确性,prefill+decode 的输出和一次性 forward 数值对齐到 1e-8第三,把这份朴素实现拿去和 MHA 比 latency,结果 16k 上下文时它比 MHA 慢了两倍

慢不是 bug,慢是在告诉你:MLA 从论文到生产引擎之间,还隔着一个叫 absorption 的线性代数技巧,而这个技巧之所以能做,又是因为 DeepSeek 多写了一条 decoupled RoPE 分支。这篇文章用实验把这条链路走完:先写出能 work 的最小 MLA,再量一遍为什么它慢,再补一版 absorbed decode,把它压回比 MHA 还省的区间。

如果你看完 vLLM / SGLang 的 MLA kernel 觉得它“比论文复杂很多”,这篇文章是帮你把论文读到 kernel 的那一步。

测试平台:RTX 3090 (24 GB),PyTorch 2.9.1 + CUDA 12.8,fp16 forward。所有代码和原始日志在文末给出;欢迎把相同脚本在 4090、A100、H100 上跑一遍,数据会变,但结论的方向不会变。


1. MLA 到底省在哪里:一张每 token 字节数表

讨论 MLA 之前先把对手摆清楚。对于一个 decoder-only 模型,KV cache 每层、每 token 的字节数只和 K/V 的形状 + dtype 有关:

注意力形态 每层、每 token 字节(fp16)
MHA(n_h × d_h K + 同规模 V) 2 × n_h × d_h × 2
GQA(n_kv × d_h K + 同规模 V) 2 × n_kv × d_h × 2
MLA(只缓存 c_kv 和 k_R,所有 head 共享) (d_c + d_r) × 2

代入 DeepSeek-V2 论文给出的 n_h=128, d_h=128, d_c=512, d_r=64,MLA 每层每 token 只要 (512+64)×2 = 1152 字节,MHA 等规模要 2×128×128×2 = 65536 字节——57 倍差距。把 60 层、32k 上下文放进公式:

config              MHA/tok     GQA8/tok    MLA/tok          60L × 32k 总量
DeepSeek-V2 style   64.00 KB     4.00 KB     1.12 KB     MHA=120 GB  GQA=7.5 GB  MLA=2.11 GB
Llama-3-70B-like    32.00 KB     4.00 KB     1.12 KB     MHA=60  GB  GQA=7.5 GB  MLA=2.11 GB

3090 上实测分配验证了这组数字,(d_c+d_r) 两段直接用 fp16 tensor 开出来:

L= 2048  MHA/layer=128.00 MB   MLA/layer= 2.25 MB   ratio=56.9x
L= 8192  MHA/layer=512.00 MB   MLA/layer= 9.00 MB   ratio=56.9x
L=32768  MHA/layer= 2.00 GB    MLA/layer=36.00 MB   ratio=56.9x
L=65536  MHA/layer= 4.00 GB    MLA/layer=72.00 MB   ratio=56.9x

作者判断 1:如果你要横向比较多个注意力方案的“上下文成本”,按每层每 token 的字节数算比按论文声明的压缩倍数靠谱得多。同样是“压缩”,GQA 把 K/V 共享到 n_kv 个分组,MLA 是把 K/V 压进一条共享 latent;在 Llama-3-70B 这种已经 GQA-8 的基线上,MLA 的相对优势从“57 倍”缩到“3.6 倍”,并不是任何模型换 MLA 都能看到论文里的那种差距。


2. 用 80 行 PyTorch 把 MLA 写清楚

MLA 论文里最容易被忽略的是:它的 Q 也走了低秩分解,且把 KV 拆成 content(走低秩 latent)rope(走独立 rotary 分支) 两段。只有理解这个拆法,后面 absorption 和 decoupled RoPE 才讲得通。

先给结构:

h  ─┐                 ┌─► c_q  ──► q_c (per head, d_h)
    ├─ W_DQ/W_DKV ─── ┤
        │                 ├─► c_q  ──► q_r (per head, d_r)   ┐
            │                                                     ├─► RoPE
                │                                                     │
                    ├─ W_KR ────────────────────────► k_r (shared, d_r)  ┘
                        │
                            └─ W_DKV ─────────► c_kv (shared latent, d_c)
                                                       │
                                                                                  ├─► W_UK ──► k_c (per head, d_h)
                                                                                                             └─► W_UV ──► v   (per head, d_h)
                                                                                                             ```
cache 只保留两条线:**共享的 `c_kv ∈ [B, L, d_c]`** 和 **共享的 `k_r ∈ [B, L, d_r]`**。其余的 Q、K_C、V 都是运行时算出来的。

下面是我这版最小 MLA 的核心代码(完整版见附录)。变量名尽量沿用论文:`d_c` 是 KV latent 维度,`d_r` 是 rotary 分支维度,`d_cq` 是 Q latent 维度。

```python
class MLA(nn.Module):
    def __init__(self, d_model=1024, n_h=16, d_h=64, d_c=128, d_r=32, d_cq=256):
            super().__init__()
                    self.n_h, self.d_h, self.d_c, self.d_r, self.d_cq = n_h, d_h, d_c, d_r, d_cq
        self.Wdq  = nn.Linear(d_model, d_cq, bias=False)
                self.Wuq  = nn.Linear(d_cq,    n_h * d_h, bias=False)   # content Q
                        self.Wqr  = nn.Linear(d_cq,    n_h * d_r, bias=False)   # rope Q
        self.Wdkv = nn.Linear(d_model, d_c, bias=False)
                self.Wuk  = nn.Linear(d_c,     n_h * d_h, bias=False)   # content K
                        self.Wuv  = nn.Linear(d_c,     n_h * d_h, bias=False)   # V
                                self.Wkr  = nn.Linear(d_model, d_r, bias=False)         # shared K_R
        self.Wo   = nn.Linear(n_h * d_h, d_model, bias=False)
    def forward(self, h, past=None):
            B, L, _ = h.shape
                    n_h, d_h, d_c, d_r = self.n_h, self.d_h, self.d_c, self.d_r
        c_kv = self.Wdkv(h)                # [B, L, d_c]   ← cache
                k_r  = self.Wkr(h)                 # [B, L, d_r]   ← cache
                        if past is not None:
                                    c_kv = torch.cat([past[0], c_kv], dim=1)
                                                k_r  = torch.cat([past[1], k_r ], dim=1)
                                                        Lt = c_kv.shape[1]
        qz  = self.Wdq(h)
                q_c = self.Wuq(qz).view(B, L, n_h, d_h)
                        q_r = self.Wqr(qz).view(B, L, n_h, d_r)
        k_c = self.Wuk(c_kv).view(B, Lt, n_h, d_h)         # 运行时重建
                v   = self.Wuv(c_kv).view(B, Lt, n_h, d_h)         # 运行时重建
                        k_r_h = k_r.unsqueeze(2).expand(B, Lt, n_h, d_r)   # 所有 head 共享
        # decoupled RoPE: 只作用在 d_r 段上
                pos_q = torch.arange(Lt - L, Lt, device=h.device)
                        pos_k = torch.arange(Lt,        device=h.device)
                                cos_q, sin_q = build_rope(pos_q, d_r, device=h.device, dtype=h.dtype)
                                        cos_k, sin_k = build_rope(pos_k, d_r, device=h.device, dtype=h.dtype)
                                                q_r   = apply_rope(q_r,   cos_q[None, :, None, :], sin_q[None, :, None, :])
                                                        k_r_h = apply_rope(k_r_h, cos_k[None, :, None, :], sin_k[None, :, None, :])
        q = torch.cat([q_c, q_r], dim=-1)      # [B, L,  H, d_h+d_r]
                k = torch.cat([k_c, k_r_h], dim=-1)    # [B, Lt, H, d_h+d_r]
                        q = q.transpose(1, 2); k = k.transpose(1, 2); vv = v.transpose(1, 2)
                                o = F.scaled_dot_product_attention(q, k, vv, is_causal=(past is None))
                                        return self.Wo(o.transpose(1, 2).contiguous().view(B, L, n_h * d_h)), (c_kv, k_r)
                                        ```
正确性先过一遍。把同一段序列分两次输入(前面 L-1 做 prefill,最后一个 token 带着 cache 进来)和整段一次性 forward 对比:

== cache equivalence test ==
prefill segment max diff: 0.00e+00
decoded last token max diff: 4.47e-08


数值基本在 fp32 的 round-off 噪声上;cache 合并逻辑 OK。到这里 MLA 已经能 work 了。

**作者判断 2**:很多人复现 MLA 时的第一个坑,是把 `k_r` 和 `q_r` 也放到 per-head 里去,然后忘了 `k_r` 其实是“所有 head 共享”的——这会让你的 cache 里多存 `n_h × d_r` 倍的冗余,也会让 absorption 没法做。上面 `k_r = self.Wkr(h)`(维度 `[B, L, d_r]`,与 n_h 解耦)是 DeepSeek 原始设计的关键。

---

## 3. 先给一个反直觉结果:我的 80 行 MLA 比 MHA 还慢

把同一套超参(`d_model=2048, n_h=16, d_h=128`)下的 MHA 和上面这版朴素 MLA(`d_c=384, d_r=32, d_cq=1024`)塞进 fp16、单层、单卡的 prefill+decode 测试:

== prefill=1024, decode=256 ==
MHA prefill= 1.65 ms decode/tok=0.449 ms cache=10.00 MB
MLA prefill= 2.77 ms decode/tok=1.144 ms cache= 1.02 MB

== prefill=4096, decode=256 ==
MHA prefill= 3.86 ms decode/tok=0.245 ms cache=34.00 MB
MLA prefill= 4.91 ms decode/tok=1.057 ms cache= 3.45 MB

== prefill=16384, decode=256 ==
MHA prefill= 23.22 ms decode/tok=0.584 ms cache=130.00 MB
MLA prefill= 46.74 ms decode/tok=2.211 ms cache= 13.20 MB


结论直白:**cache 是省了,latency 是赔的**。16k 上下文时 MLA 每个 decode token 要 2.2 ms,几乎是 MHA 的 4 倍。

为什么?看 decode 步骤里 MLA 多干的活:

- `self.Wuk(c_kv)`:把整条 `c_kv ∈ [B, Lt, d_c]` 升维到 `[B, Lt, n_h·d_h]`,这是一个 `Lt × d_c × (n_h·d_h)` 的 matmul,随 Lt 线性增长;
- - `self.Wuv(c_kv)`:同上,再做一次;
- - 之后再把 K/V 展平成 `[B, H, Lt, d_h]` 做 attention。
也就是说,**每产生一个 token,你都在把整条历史 cache 做两次 up-projection**。MHA 没这个负担:K 和 V 本来就是 `[B, Lt, H, d_h]`,append 一下就能用。

所以 DeepSeek V2 的 paper 并没有骗你。它真实的意思是:“cache 能压 57 倍,但是朴素实现的 decode kernel 会慢,所以我们论文里顺手给了一个 absorption 技巧来把这两次升维吸进 Q 端的常数矩阵”。

**作者判断 3**:如果你在博客或知乎看到“把 MLA 换进 MHA 代码里就能加速长上下文”的说法,应该警惕。没有 absorption 的 MLA 在 decode 阶段是被带宽 + 计算两头打的,它只在 cache 极端紧张(比如 3090/4090 跑 65k 上下文,MHA 直接 OOM 的场景)下才一定更优。工程团队要么接 vLLM/SGLang 的 MLA kernel,要么自己把 absorption 写对,才能吃到它的红利。

---

## 4. 为什么是“decoupled RoPE”?一行 einsum 就能看清

论文里 `K = concat(K_C, K_R)` 的拆法第一眼看上去像工程妥协。真正的原因是:MLA 的压缩技巧依赖 `W_UK` 和 Q 之间做矩阵合并(absorption),而 RoPE 不允许这种合并。

先写两条路径:

```python
# Path A: 朴素做法——先从 c_kv 重建每 head 的 K_C,再和 Q_C 算分数
K_C = torch.einsum('bld,hdk->blhk', c_kv, W_UK)              # [B,L,H,Dh]
scores_A = torch.einsum('blhk,bmhk->bhlm', Q_C, K_C)

# Path B: absorption——把 W_UK 吸到 Q 端,直接和 c_kv 打分
Q_abs = torch.einsum('blhk,hdk->blhd', Q_C, W_UK)            # [B,L,H,Dc]
scores_B = torch.einsum('blhd,bmd->bhlm', Q_abs, c_kv)

两条路径在线性代数上是同一件事,实测也是:

content-path absorption max diff: 2.38e-06

误差就是 fp32 的舍入噪声。这意味着 decode 步骤可以跳过 W_UK 矩阵、直接让 Q_abs 去和 cache 打分,维度从 [B,L,H,Dh] · [B,Lt,H,Dh] 变成 [B,L,H,Dc] · [B,Lt,Dc],后者只需一份共享 c_kv 而不是 per-head K——这正是 vLLM / SGLang MLA kernel 的真实形态。

然后我们把 RoPE 加回 K_C 再看:

K_C_rope = K_C * rot                                  # 简化的 RoPE 位置因子
scores_A_rope = torch.einsum('blhk,bmhk->bhlm', Q_C, K_C_rope)
scores_B_rope_broken = torch.einsum('blhd,bmd->bhlm', Q_abs, c_kv)  # 还是未旋转的 c_kv

此时:

if we RoPE K_C, absorption path diverges by: 1.03e+01

误差从 1e-6 跳到 10.3。原因:RoPE 的位置因子 rot(pos) 是和 W_UK 之后的那个 Dh 维绑在一起的,W_UK · rot(pos) 不等于 rot(pos) · W_UK,absorption 步骤把 W_UK 预吸到 Q_abs 里之后就再也对不上了。

这就是为什么 MLA 必须把 K 拆成两段:content 分支是纯线性映射,允许做 absorption;rope 分支用很小的 d_r(DeepSeek-V2 是 64),单独挂在一边,大家都 RoPE 一下再拼起来,绕开矩阵合并冲突

作者判断 4:“decoupled RoPE” 的真正含义不是“多一个额外的 RoPE”,而是“用很小的 d_r 换来 W_UK 可以被吸收”。这是 MLA 从“cache 更小”进化到“decode 也不慢”的那把钥匙,很多教程在讲 MLA 时把它讲成工程手感,其实是数学硬约束。


5. 把 absorption 补上:decode/tok 从 2.2 ms 回到 1.4 ms

光证明 absorption 在数学上等价没意义,我把它写成一个只在 decode 时生效的版本(MLAAbsorbDecode,完整代码见附录)。核心差异:

  • decode 时不再调用 Wuk(c_kv) / Wuv(c_kv),而是
    • W_UK 重排成 [H, d_c, d_h],和 Q 做 einsum('blhd,hcd->blhc'),得到 Q_abs;
    • content 分数 = Q_abs · c_kv^T,rope 分数 = q_r · k_r^T(共享 k_r 广播到各 head);
    • attention 先在 latent 维得到 ctx_latent [B, H, 1, d_c],再用 W_UV 投到 [B, H, 1, d_h]
      在相同模型、相同 prefill 后续 256 token decode 的条件下:
prefill=  1024  base-MLA/tok=1.071 ms   absorbed/tok=1.384 ms   diff=1.53e-05   speedup=0.77x
prefill=  4096  base-MLA/tok=1.085 ms   absorbed/tok=1.343 ms   diff=7.63e-06   speedup=0.81x
prefill= 16384  base-MLA/tok=2.200 ms   absorbed/tok=1.364 ms   diff=3.81e-06   speedup=1.61x

三点值得看:

  1. correctness:fp16 下和朴素版本最大差 1.5e-5,和 fp16 精度同量级。
    1. 短上下文 absorption 反而更慢(0.77x):因为 absorbed 版本每个 decode 都要多做一次 W_UK 重排和 einsum,在 Lt 小时那点常数开销盖过了收益。
    1. 长上下文 absorption 变成 1.61x:16k 时朴素版本被 Wuk(c_kv) 拖成 2.2 ms/tok,absorbed 版稳定在 1.36 ms/tok——这就是生产 kernel 要把整件事搬到 CUDA 层的原因,Python einsum 都能看出趋势,手写 kernel 把常数摁下去之后这条曲线会更陡。
      作者判断 5:在 Python + PyTorch 层写 MLA,最多只能看到“absorption 在长上下文变快”这种趋势;真正的低延迟必须靠 CUDA kernel(例如 SGLang 的 flash_mla、vLLM 的 mla_fwd)把 Q_abs、c_kv、k_r 的 gather + attention 融进一个 pass。如果你的推理场景是中短上下文(<4k)、batch 很小、模型层数不多,MHA/GQA 路径反而更快,换 MLA 不一定划算。

6. 如果你要把 MLA 放进真实项目,先决定这几件事

按使用路径分三类:

一、你直接用 DeepSeek-V2 / V3 / R1 官方 checkpoint

  • 如果是 vLLM 推理:升到支持 MLA 专用 kernel 的版本(MLA fast path 在 2024 年底开始合进主线,2025 年以后已是默认)。跑起来就能享受 57x cache 压缩和长上下文 throughput。

    • 如果是 Hugging Face Transformers 纯 eager forward:你拿到的是我第 2 节那版的性能画风——cache 省了,但 decode 慢,长序列才打平 MHA。别以此为依据评估 MLA。
    • 如果你写单卡 demo 脚本(3090/4090),请记得 d_c=512 意味着每层每 token 1.1 KB;32k 上下文 60 层不过 2 GB,这是 MLA 在消费卡跑长上下文 LLM 的核心价值。
      二、你要把 MLA 塞进自己的非 DeepSeek 架构
  • 给出 d_c(通常 d_model/2 ~ d_model/4 之间)和 d_r(64 是个稳妥值,别调到 8 这种极端值)。

    • 训练时 Q/K/V 各一条 down+up 投影都得带上;只做 KV 投影而不做 Q 投影,参数量省不了多少但会把收敛搞复杂。
    • 预训练初期 MLA 模型收敛速度和 MHA 接近,真正的差异在长上下文 eval 时才出现;不要拿 1k 上下文的 perplexity 判断它优劣。
      三、你只想自己推理时加载 MLA 模型
  • 不要尝试把 MLA 模型手工“反”成 MHA 再跑,因为 W_UK/W_UV 拼起来的等效 K/V 维度是 n_h × d_h = 16384,单个 K 矩阵就 256 MB,反过来算既慢又占显存。

    • 真的想 debug 精度,就关掉 absorption,用第 2 节那个朴素版跑一遍;它和 absorption 版在 fp32 上对得上 1e-6、fp16 对得上 1e-5,是你的 ground truth。

7. 参考与延伸阅读

  • DeepSeek-AI, DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model, arXiv:2405.04434(MLA 首次系统提出,含完整超参与对比实验)。
    • DeepSeek-AI, DeepSeek-V3 Technical Report, arXiv:2412.19437(MLA 在更大模型上的配置与训练 recipe,d_c=512, d_r=64 的来源)。
    • vLLM 源码:vllm/attention/backends/mla.py 中的 apply_forward_flash_mla_fwd,对应本文第 5 节 absorption + decoupled RoPE 的 CUDA 版本。
    • SGLang 源码:sglang/srt/layers/attention/flashinfer_mla_backend.py,可以对比 FlashInfer 的 MLA kernel 与本文 einsum 版在结构上的异同。
    • 本文实验脚本(mla_min.py / run_bench.py / absorb_demo.py / absorbed_decode.py):复现结果只需要 PyTorch ≥ 2.1 + 任意 10+ GB 显存 GPU;所有命令和日志见末尾附录。

附录:完整实验脚本

A. mla_min.py(本文第 2 节)

import math, torch
import torch.nn as nn
import torch.nn.functional as F


def rotate_half(x):
    x1, x2 = x[..., ::2], x[..., 1::2]
        return torch.stack((-x2, x1), dim=-1).flatten(-2)

def apply_rope(x, cos, sin):
    return x * cos + rotate_half(x) * sin

def build_rope(pos, dim, base=10000.0, device=None, dtype=torch.float32):
    inv = 1.0 / (base ** (torch.arange(0, dim, 2, device=device, dtype=dtype) / dim))
        freqs = torch.outer(pos.to(inv.dtype), inv)
            emb = torch.repeat_interleave(freqs, 2, dim=-1)
                return emb.cos(), emb.sin()

class MLA(nn.Module):
    def __init__(self, d_model=1024, n_h=16, d_h=64, d_c=128, d_r=32, d_cq=256):
            super().__init__()
                    self.n_h, self.d_h, self.d_c, self.d_r, self.d_cq = n_h, d_h, d_c, d_r, d_cq
                            self.Wdq  = nn.Linear(d_model, d_cq, bias=False)
                                    self.Wuq  = nn.Linear(d_cq, n_h * d_h, bias=False)
                                            self.Wqr  = nn.Linear(d_cq, n_h * d_r, bias=False)
                                                    self.Wdkv = nn.Linear(d_model, d_c, bias=False)
                                                            self.Wuk  = nn.Linear(d_c, n_h * d_h, bias=False)
                                                                    self.Wuv  = nn.Linear(d_c, n_h * d_h, bias=False)
                                                                            self.Wkr  = nn.Linear(d_model, d_r, bias=False)
                                                                                    self.Wo   = nn.Linear(n_h * d_h, d_model, bias=False)
    def forward(self, h, past=None):
            B, L, _ = h.shape
                    n_h, d_h, d_c, d_r = self.n_h, self.d_h, self.d_c, self.d_r
                            c_kv = self.Wdkv(h)
                                    k_r = self.Wkr(h)
                                            if past is not None:
                                                        c_kv = torch.cat([past[0], c_kv], dim=1)
                                                                    k_r = torch.cat([past[1], k_r], dim=1)
                                                                            Lt = c_kv.shape[1]
                                                                                    qz = self.Wdq(h)
                                                                                            q_c = self.Wuq(qz).view(B, L, n_h, d_h)
                                                                                                    q_r = self.Wqr(qz).view(B, L, n_h, d_r)
                                                                                                            k_c = self.Wuk(c_kv).view(B, Lt, n_h, d_h)
                                                                                                                    v = self.Wuv(c_kv).view(B, Lt, n_h, d_h)
                                                                                                                            k_r_h = k_r.unsqueeze(2).expand(B, Lt, n_h, d_r)
                                                                                                                                    pos_q = torch.arange(Lt - L, Lt, device=h.device)
                                                                                                                                            pos_k = torch.arange(Lt, device=h.device)
                                                                                                                                                    cos_q, sin_q = build_rope(pos_q, d_r, device=h.device, dtype=h.dtype)
                                                                                                                                                            cos_k, sin_k = build_rope(pos_k, d_r, device=h.device, dtype=h.dtype)
                                                                                                                                                                    q_r = apply_rope(q_r, cos_q[None, :, None, :], sin_q[None, :, None, :])
                                                                                                                                                                            k_r_h = apply_rope(k_r_h, cos_k[None, :, None, :], sin_k[None, :, None, :])
                                                                                                                                                                                    q = torch.cat([q_c, q_r], dim=-1)
                                                                                                                                                                                            k = torch.cat([k_c, k_r_h], dim=-1)
                                                                                                                                                                                                    q = q.transpose(1, 2); k = k.transpose(1, 2); vv = v.transpose(1, 2)
                                                                                                                                                                                                            o = F.scaled_dot_product_attention(q, k, vv, is_causal=(past is None))
                                                                                                                                                                                                                    return self.Wo(o.transpose(1, 2).contiguous().view(B, L, n_h * d_h)), (c_kv, k_r)
                                                                                                                                                                                                                    ```
### B. KV 字节表和 3090 分配(本文第 1 节实验日志)

== cache equivalence test ==
prefill segment max diff: 0.00e+00
decoded last token max diff: 4.47e-08

== per-token KV bytes ==
config MHA/tok GQA8/tok MLA/tok (60L × 32k)
DeepSeek-V2 style 64.00 KB 4.00 KB 1.12 KB MHA=120.00 GB GQA=7.50 GB MLA=2.11 GB
Llama-3-70B-like 32.00 KB 4.00 KB 1.12 KB MHA= 60.00 GB GQA=7.50 GB MLA=2.11 GB

== live 3090 allocation ==
L= 2048 MHA/layer= 128.00 MB MLA/layer= 2.25 MB ratio=56.9x
L= 8192 MHA/layer= 512.00 MB MLA/layer= 9.00 MB ratio=56.9x
L= 32768 MHA/layer= 2.00 GB MLA/layer= 36.00 MB ratio=56.9x
L= 65536 MHA/layer= 4.00 GB MLA/layer= 72.00 MB ratio=56.9x


### C. absorption 和 decoupled RoPE 的等价/冲突检查(第 4 节)

content-path absorption max diff: 2.38e-06
if we RoPE K_C, absorption path diverges by: 1.03e+01 (this is the reason decoupled RoPE is mandatory)


### D. absorbed decode 与朴素版本 latency 对比(第 5 节)

prefill= 1024 base-MLA/tok=1.071 ms absorbed/tok=1.384 ms correctness max-diff=1.53e-05 speedup=0.77x
prefill= 4096 base-MLA/tok=1.085 ms absorbed/tok=1.343 ms correctness max-diff=7.63e-06 speedup=0.81x
prefill= 16384 base-MLA/tok=2.200 ms absorbed/tok=1.364 ms correctness max-diff=3.81e-06 speedup=1.61x


> 测试平台:RTX 3090 (24 GB) + Driver 590.44.01 + CUDA 13.1,PyTorch 2.9.1+cu128,fp16 forward,batch=1,`d_model=2048, n_h=16, d_h=128, d_c=384, d_r=32, d_cq=1024`。
> 
Logo

欢迎加入DeepSeek 技术社区。在这里,你可以找到志同道合的朋友,共同探索AI技术的奥秘。

更多推荐