File size: 8,231 Bytes
5e6d9f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
"""MLA:Multi-head Latent Attention(多头潜在注意力)。

这是 DeepSeek-V2 起最核心的省显存设计。

标准 MHA 的 KV cache 每 token 要存 2 * n_heads * head_dim 个数;
MLA 先把每个 token 压成一个低维潜在向量 c_kv(kv_lora_rank 维),
用的时候再用 W_UK / W_UV 升维还原出 K 和 V。于是 cache 每 token 只需要存
    kv_lora_rank + qk_rope_head_dim
个数,nano 档就是 64+16=80,而同规模 MHA 需要 2*8*32=512,省了 6.4 倍。

一个麻烦:RoPE 是位置相关的旋转,如果 K 是从 c_kv 现算出来的,
旋转矩阵没法和 W_UK 交换位置,"矩阵吸收"技巧就失效了。
DeepSeek 的解法是「解耦 RoPE」:额外切出 qk_rope_head_dim 维专门承载位置信息,
这部分所有 head 共享、直接缓存;剩下的 qk_nope 部分完全不带位置信息。

两条实现路径:
  naive  —— 把 K/V 显式还原出来,走 F.scaled_dot_product_attention(能吃到 flash 内核),
            训练时用它最快。
  absorb —— 把 W_UK 吸收进 Q、W_UV 吸收进输出投影,全程在潜在空间里算注意力,
            长上下文解码时中间激活小得多。
两条路径数学上完全等价,tests.py 里有对拍。
"""

import math
from typing import Optional

import torch
import torch.nn as nn
import torch.nn.functional as F

from .layers import RMSNorm, apply_rope


class MLA(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.dim = cfg.dim
        self.n_heads = cfg.n_heads
        self.q_lora_rank = cfg.q_lora_rank
        self.kv_lora_rank = cfg.kv_lora_rank
        self.qk_nope_head_dim = cfg.qk_nope_head_dim
        self.qk_rope_head_dim = cfg.qk_rope_head_dim
        self.qk_head_dim = cfg.qk_head_dim
        self.v_head_dim = cfg.v_head_dim
        self.attn_impl = cfg.attn_impl
        self.softmax_scale = self.qk_head_dim ** -0.5

        # ---- Query 投影:可选低秩分解(大模型上能省不少参数)----
        if self.q_lora_rank == 0:
            self.wq = nn.Linear(self.dim, self.n_heads * self.qk_head_dim, bias=False)
        else:
            self.wq_a = nn.Linear(self.dim, self.q_lora_rank, bias=False)
            self.q_norm = RMSNorm(self.q_lora_rank, cfg.norm_eps)
            self.wq_b = nn.Linear(self.q_lora_rank, self.n_heads * self.qk_head_dim, bias=False)

        # ---- KV 联合压缩:一次投影同时产出潜在向量 c_kv 和共享的 k_pe ----
        self.wkv_a = nn.Linear(self.dim, self.kv_lora_rank + self.qk_rope_head_dim, bias=False)
        self.kv_norm = RMSNorm(self.kv_lora_rank, cfg.norm_eps)
        self.wkv_b = nn.Linear(self.kv_lora_rank,
                               self.n_heads * (self.qk_nope_head_dim + self.v_head_dim), bias=False)

        self.wo = nn.Linear(self.n_heads * self.v_head_dim, self.dim, bias=False)
        self.dropout_p = cfg.dropout

        # 推理缓存(训练时为 None)
        self.kv_cache: Optional[torch.Tensor] = None
        self.pe_cache: Optional[torch.Tensor] = None

    # ------------------------------------------------------------------
    def setup_cache(self, max_batch: int, max_seq_len: int, device, dtype):
        self.kv_cache = torch.zeros(max_batch, max_seq_len, self.kv_lora_rank,
                                    device=device, dtype=dtype)
        self.pe_cache = torch.zeros(max_batch, max_seq_len, self.qk_rope_head_dim,
                                    device=device, dtype=dtype)

    def clear_cache(self):
        self.kv_cache = None
        self.pe_cache = None

    def cache_bytes(self, seq_len: int) -> int:
        if self.kv_cache is None:
            return 0
        per_token = self.kv_lora_rank + self.qk_rope_head_dim
        return per_token * seq_len * self.kv_cache.element_size()

    # ------------------------------------------------------------------
    def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor,
                start_pos: int = 0, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
        B, T, _ = x.shape
        end = start_pos + T

        # ---------------- Query ----------------
        if self.q_lora_rank == 0:
            q = self.wq(x)
        else:
            q = self.wq_b(self.q_norm(self.wq_a(x)))
        q = q.view(B, T, self.n_heads, self.qk_head_dim).transpose(1, 2)   # (B,H,T,qk)
        q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
        q_pe = apply_rope(q_pe, cos, sin)

        # ---------------- KV 压缩 ----------------
        kv = self.wkv_a(x)                                                  # (B,T,c+rope)
        c_kv, k_pe = kv.split([self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
        c_kv = self.kv_norm(c_kv)                                           # 缓存的就是它
        k_pe = apply_rope(k_pe.unsqueeze(1), cos, sin).squeeze(1)           # (B,T,rope),所有 head 共享

        if self.kv_cache is not None:
            self.kv_cache[:B, start_pos:end] = c_kv.to(self.kv_cache.dtype)
            self.pe_cache[:B, start_pos:end] = k_pe.to(self.pe_cache.dtype)
            c_kv_all = self.kv_cache[:B, :end].to(x.dtype)
            k_pe_all = self.pe_cache[:B, :end].to(x.dtype)
        else:
            c_kv_all, k_pe_all = c_kv, k_pe

        if self.attn_impl == "naive":
            out = self._attn_naive(q_nope, q_pe, c_kv_all, k_pe_all, mask)
        else:
            out = self._attn_absorb(q_nope, q_pe, c_kv_all, k_pe_all, mask)

        out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.v_head_dim)
        return self.wo(out)

    # ------------------------------------------------------------------
    def _attn_naive(self, q_nope, q_pe, c_kv_all, k_pe_all, mask):
        """显式还原 K/V,交给 SDPA(可命中 flash-attention 内核)。"""
        B, H, T, _ = q_nope.shape
        S = c_kv_all.shape[1]

        kv = self.wkv_b(c_kv_all).view(B, S, H, self.qk_nope_head_dim + self.v_head_dim)
        kv = kv.transpose(1, 2)                                             # (B,H,S,nope+v)
        k_nope, v = kv.split([self.qk_nope_head_dim, self.v_head_dim], dim=-1)

        k = torch.cat([k_nope, k_pe_all.unsqueeze(1).expand(B, H, S, -1)], dim=-1)
        q = torch.cat([q_nope, q_pe], dim=-1)

        return F.scaled_dot_product_attention(
            q, k, v,
            attn_mask=mask,
            is_causal=(mask is None and T > 1),
            dropout_p=self.dropout_p if self.training else 0.0,
            scale=self.softmax_scale,
        )

    def _attn_absorb(self, q_nope, q_pe, c_kv_all, k_pe_all, mask):
        """矩阵吸收:全程在 kv_lora_rank 维的潜在空间里做注意力。"""
        B, H, T, _ = q_nope.shape
        W = self.wkv_b.weight.view(H, self.qk_nope_head_dim + self.v_head_dim, self.kv_lora_rank)
        W_uk = W[:, : self.qk_nope_head_dim]        # (H, nope, c)
        W_uv = W[:, self.qk_nope_head_dim:]         # (H, v, c)

        # q_nope @ W_UK  —— 把升维矩阵吸收进 Q,K 就不用还原了
        q_absorb = torch.einsum("bhtd,hdc->bhtc", q_nope, W_uk.to(q_nope.dtype))

        scores = (torch.einsum("bhtc,bsc->bhts", q_absorb, c_kv_all)
                  + torch.einsum("bhtd,bsd->bhts", q_pe, k_pe_all)) * self.softmax_scale

        S = c_kv_all.shape[1]
        if mask is not None:
            scores = scores + _to_additive(mask, scores.dtype)
        elif T > 1:
            causal = torch.ones(T, S, dtype=torch.bool, device=scores.device).tril(S - T)
            scores = scores.masked_fill(~causal, torch.finfo(scores.dtype).min)

        attn = scores.softmax(dim=-1, dtype=torch.float32).type_as(scores)
        if self.training and self.dropout_p > 0:
            attn = F.dropout(attn, self.dropout_p)

        x_lat = torch.einsum("bhts,bsc->bhtc", attn, c_kv_all)              # 仍在潜在空间
        return torch.einsum("bhtc,hdc->bhtd", x_lat, W_uv.to(x_lat.dtype))  # 输出时才用 W_UV 还原


def _to_additive(mask: torch.Tensor, dtype) -> torch.Tensor:
    if mask.dtype == torch.bool:
        return torch.zeros_like(mask, dtype=dtype).masked_fill(~mask, torch.finfo(dtype).min)
    return mask.to(dtype)