Transformers
Safetensors
English
mla
deepseek-moe
mtp
custom-code
tinystories
from-scratch
Eval Results (legacy)
Instructions to use nowordsxiaomu/DeepSeek-Flash-Mini with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use nowordsxiaomu/DeepSeek-Flash-Mini with Transformers:
# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("nowordsxiaomu/DeepSeek-Flash-Mini", device_map="auto") - Notebooks
- Google Colab
- Kaggle
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)
|