
MHA的设计动机为什么要有MHA用单头注意力不行吗我们先放一下注意力的公式其中Softmax算子作用于每一行输出一个归一化的概率分布分布阵 在矩阵A中任意一行的和1也就是说强制将所有上下文关系归一化到这一个概率分布当中。但是自然语言中词与词的关系是多维重叠的。例如句子“苹果 昨天 发布了 新手机”对token “苹果来说”语法维度需要关注动词发布了主谓关系语义修饰维度需要关注新手机实体与产品的语义关联位置维度需要关注紧随其后的昨天。如果只有单头注意力这一概率向量就必须同时兼顾上述所有关系结果就是不同维度的信号相互抵消、平滑导致 模型无法精确捕捉特定维度的上下文关系。MHA的做法多头注意力通过将高维特征空间切分成H个低维子空间实现了特征维度的解耦在数学上每个头拥有独立的投影矩阵。这意味着模型可以在 H 个不同的低维子空间中并行计算 H组完全独立的注意力得分矩阵head1专门提取主谓关系head2专门提取代词指代head3专门提取短语局部位置。......GQA的设计动机解决KV Cache占用显存过大的问题什么是kv cache可以到我的另一篇博客回顾一下kv cache和因果掩码kv cache在多头注意力中需要把每个头的k v都缓存下来占用了大量的显存并且在计算时需要把k v从显存中搬运到计算单元中首先于GPU 带宽问题带来了吞吐量瓶颈的问题。那么减少 K V头的数量是不是就可以了确实是这样我们可以让多个Q共享同一个 KV就可以这就是分组注意力例如原先是32个多头现在把KV分成8组也就是每4个Q共享同一个KV此时KV Cache的显存消耗是MHA的 8/3225%降低了75%的显存占用大大缓解了带宽瓶颈问题。其实MQA更加激进所有的Q共享一个K V但是会导致性能下降为什么Q要保持多头呢我们先来回顾一下Q、K、V的角色1.Qquery提问者含义“我在找什么信息”作用$Q$ 代表当前 Token 站在自己的角度向全篇文章发出的“查询请求”。例子在句子苹果 发布了 新手机中Token发布了的 $Q$ 可能会问“谁发布的”找主语以及“发布了什么”找宾语。2.Kkey/标签、检索索引含义“我是什么特征凭什么吸引你”作用$K$ 是当前 Token 挂在门口的“身份标签”用来和别人的 $Q$ 做匹配。匹配过程QK点积就是在算“你提的问题和我的标签匹配度有多高”3.Vvalue、内容信息载体含义“如果我的标签被选中了我能提供什么具体语义”用V是这个 Token 真正包含的“语义内容”。一旦 Q和K 匹配成功模型就会把对应的V 按比例“抽走”并融合到输出中。因为Q承担着一个向外“探索”的角色这个视角必须是丰富的多维度的。对于同一个 Token比如苹果它在不同的“头”里需要同时提出很多不同的问题Q1头“后文有没有动词我要找我的动作。”Q2头“前文有没有修饰词我是水果还是公司”Q3头“句子里有没有代词指代我”Q4头“我是不是句子开头的主语”如果Q头太少模型就无法同时在多个不同的抽象维度语法、语义、上下文上寻找信息因此Q需要保持多头。为什么KV可以不需要那么多头呢可以从两个角度来分析从数学上由于Q是多头的可以保证后续的注意力计算的结果是不同的概率分布可以保证提取语义不同维度信息在对训练好的MHA模型进行奇异值分解SVD和相关性分析后发现可能第1~4个KV头在表达语法属性第5~8个KV头在表达“实体语义属性”而剩下的9~32个头学习到的内容和前8个头高度重合大量冗余。因此适量的减少KV头不会造成性能下降一下是一份GQA的一份实现import torch import torch.nn as nn import os import cv2 import numpy as np from collections import defaultdict import math import torch import torch.nn as nn import torch.nn.functional as F def repeat_kv(hidden_states: torch.Tensor, n_rep: int) - torch.Tensor: 按组扩展 KV 的 Head 数量使其与 Q 的 Head 数量匹配。 输入 shape: [batch_size, num_kv_heads, seq_len, head_dim] 输出 shape: [batch_size, num_q_heads, seq_len, head_dim] (其中 num_q_heads num_kv_heads * n_rep) batch_size, num_kv_heads, seq_len, head_dim hidden_states.shape if n_rep 1: return hidden_states # 通过内存视图与广播机制扩展维度避免额外复制开销 # 1. 扩展维度 - [batch_size, num_kv_heads, 1, seq_len, head_dim] # 2. 广播复制 - [batch_size, num_kv_heads, n_rep, seq_len, head_dim] # 3. 合并 Head 维度 - [batch_size, num_kv_heads * n_rep, seq_len, head_dim] hidden_states hidden_states[:, :, None, :, :].expand( batch_size, num_kv_heads, n_rep, seq_len, head_dim ) return hidden_states.reshape(batch_size, num_kv_heads * n_rep, seq_len, head_dim) class GroupedQueryAttention(nn.Module): def __init__( self, hidden_size: int 4096, num_heads: int 32, num_key_value_heads: int 8, ): super().__init__() self.hidden_size hidden_size self.num_heads num_heads # Q 的 Head 数量 (H_Q) self.num_key_value_heads num_key_value_heads # KV 的 Head 数量 (H_KV) # 每个 Head 的特征维度 self.head_dim hidden_size // num_heads # 每个 KV Head 对应的 Q Head 重复组数 (n_rep H_Q / H_KV) self.num_key_value_groups num_heads // num_key_value_heads assert num_heads % num_key_value_heads 0, num_heads 必须能被 num_key_value_heads 整除 # 线性投影层 # Q: 映射到 num_heads * head_dim # K/V: 压缩映射到 num_key_value_heads * head_dim (降低参数量和计算量) self.q_proj nn.Linear(hidden_size, num_heads * self.head_dim, biasFalse) self.k_proj nn.Linear(hidden_size, num_key_value_heads * self.head_dim, biasFalse) self.v_proj nn.Linear(hidden_size, num_key_value_heads * self.head_dim, biasFalse) self.o_proj nn.Linear(num_heads * self.head_dim, hidden_size, biasFalse) def forward( self, hidden_states: torch.Tensor, attention_mask: torch.Tensor None, ) - torch.Tensor: 前向传播计算 输入: hidden_states: [batch_size, seq_len, hidden_size] attention_mask: [batch_size, 1, seq_len, seq_len] (可选) 输出: output: [batch_size, seq_len, hidden_size] batch_size, seq_len, _ hidden_states.shape # ---------------------------------------------------------------------- # 1. 线性投影 (Linear Projections) # ---------------------------------------------------------------------- # query_states: [B, S, num_heads * head_dim] eg: [2, 16, 32 * 128] - [2, 16, 4096] # key_states: [B, S, num_kv_heads * head_dim] eg: [2, 16, 8 * 128] - [2, 16, 1024] # value_states: [B, S, num_kv_heads * head_dim] eg: [2, 16, 8 * 128] - [2, 16, 1024] query_states self.q_proj(hidden_states) key_states self.k_proj(hidden_states) value_states self.v_proj(hidden_states) # ---------------------------------------------------------------------- # 2. Reshape 与转置分离开多头维度 (Head Dim Separation) # ---------------------------------------------------------------------- # query_states: [B, num_heads, S, head_dim] eg: [2, 32, 16, 128] query_states query_states.view( batch_size, seq_len, self.num_heads, self.head_dim ).transpose(1, 2) # key_states: [B, num_kv_heads, S, head_dim] eg: [2, 8, 16, 128] key_states key_states.view( batch_size, seq_len, self.num_key_value_heads, self.head_dim ).transpose(1, 2) # value_states: [B, num_kv_heads, S, head_dim] eg: [2, 8, 16, 128] value_states value_states.view( batch_size, seq_len, self.num_key_value_heads, self.head_dim ).transpose(1, 2) # ---------------------------------------------------------------------- # 3. 按组广播复制 KV (Repeat KV for Group Matching) # ---------------------------------------------------------------------- # key_states: [B, num_heads, S, head_dim] eg: [2, 32, 16, 128] # value_states: [B, num_heads, S, head_dim] eg: [2, 32, 16, 128] key_states repeat_kv(key_states, self.num_key_value_groups) value_states repeat_kv(value_states, self.num_key_value_groups) # ---------------------------------------------------------------------- # 4. 计算注意力得分 (Scaled Dot-Product Attention) # ---------------------------------------------------------------------- # Q * K^T - [B, num_heads, S, S] eg: [2, 32, 16, 16] attn_weights torch.matmul( query_states, key_states.transpose(2, 3) ) / math.sqrt(self.head_dim) if attention_mask is not None: attn_weights attn_weights attention_mask # 在最后一个序列维度 S 上应用 Softmax 归一化 attn_weights F.softmax(attn_weights, dim-1, dtypetorch.float32).to(query_states.dtype) # ---------------------------------------------------------------------- # 5. 加权求和输出 (Weighted Sum) # ---------------------------------------------------------------------- # attn_weights * V - [B, num_heads, S, head_dim] eg: [2, 32, 16, 128] attn_output torch.matmul(attn_weights, value_states) # ---------------------------------------------------------------------- # 6. 转置并合并多头 (Concatenate Heads) # ---------------------------------------------------------------------- # attn_output: [B, S, num_heads, head_dim] - [B, S, num_heads * head_dim] # eg: [2, 16, 32, 128] - [2, 16, 4096] attn_output attn_output.transpose(1, 2).contiguous().view( batch_size, seq_len, self.num_heads * self.head_dim ) # ---------------------------------------------------------------------- # 7. 最终线性输出投影 (Output Projection) # ---------------------------------------------------------------------- # output: [B, S, hidden_size] eg: [2, 16, 4096] output self.o_proj(attn_output) return output if __name__ __main__: # 测试参数设置模拟 Qwen3-7B 参数架构 batch_size 2 seq_len 16 hidden_size 4096 num_heads 32 # Query 头数 num_kv_heads 8 # KV 头数 (4个 Q 头共享 1个 KV 头) # 实例化 GQA 模块 gqa_layer GroupedQueryAttention( hidden_sizehidden_size, num_headsnum_heads, num_key_value_headsnum_kv_heads ) # 构造假数据输入 x torch.randn(batch_size, seq_len, hidden_size) print(f输入隐藏层张量 shape: {x.shape}) # 前向计算 out gqa_layer(x) print(f输出隐藏层张量 shape: {out.shape})MLA的提出背景GQA、MQA都是为了解决KV Cache中KV占用大量缓存的问题GQA通过多头Q共享分组KV的方式而MQA更为激进的所有Q头共享一个KV这或多或少都会影响性能。那有没有一种方式即能完整保留多头注意力又能解决KV Cache的问题呢为此deepseek提出了MLA。MLA的计算过程简单的说MLA就是把输入的高维序列矩阵X先压缩成一个低维潜隐矩阵,后续再进行多头注意力计算时再通过一个解压矩阵把解压到每个头的维度KV矩阵但是在这过程中可以利用矩阵的结合律提前把解压矩阵吸收掉达到不用恢复的效果进而达到KV Cache只需要缓存一个低维矩阵也能保留多头注意力的目的。接下来我们来看一下详细的计算过程以输入维度 D 5120为例设定 MLA 的典型超参数配置如下阶段一第一步压缩生成低秩潜隐向量输入X的形状为 [B, S, 5120]1.Q的压缩输入投影压缩矩阵对Q少量压缩计算输出2.KV的压缩K的压缩分语义为位置压缩后续的注意力计算分为语义注意力和位置编码注意力输入投影压缩矩阵:[5120, 512 64] [5120, 576]计算:Compressed_KV 拆分语义KV:KV共享一个潜隐矩阵位置K:阶段二A训练阶段显式解压QKV矩阵解压成多头1.Q拆分语义和位置部分语义部分乘上升维矩阵得到reshape成[B,128,S,128]。位置部分乘上升维矩阵得到reshape成[B,128,S,64]。这里也可以看出来K的位置部分是从X来的而Q的位置部分是从来的因为已经被压缩的够狠了再用来投影出位置部分就不合适了。拼接Q语义和位置部分拼接2. KV展开成多头K展开语义部分乘上升维矩阵,得到reshape成[B, 128, S, 128]位置部分将上述的施加ROPE展开广播为[B, 128, S, 64]拼接K语义和位置部分拼接V展开乘上升维矩阵,得到reshape成[B, 128, S, 128]3.计算得分, [B,128,S,192] * [B, 128, 192, S][B, 128, S, S]阶段二B推理阶段不解压矩阵吸收1.语义部分得分其中,而即即语义得分2.位置部分得分和相乘得到计算注意力得分[B,128,S,S]3.相加语义得分位置得分Score[B,128,S,S]4.输出加权和融合矩阵吸收Score *其中,可以提前融合成[512, 5120]最后融合矩阵再和矩阵相乘得到最终的结果[B, S, 5120]可以看出来通过步骤1和步骤4的两次融合使得推理的时候不需要把解压成[B, 128,S,128]的多头KV矩阵那么每一次前向计算KV Cache就只需缓存一个512维的向量即可相比MHA的KV Cache需要缓存128*128*232768节省了98%的显存。MLA的实现import math import torch import torch.nn as nn import torch.nn.functional as F class RoPE(nn.Module): 旋转位置编码 (Rotary Position Embedding) def __init__(self, dim: int, max_len: int 4096, theta: float 10000.0): super().__init__() # inv_freq 形状: [dim / 2] inv_freq 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) t torch.arange(max_len).float() freqs torch.outer(t, inv_freq) # [max_len, dim / 2] emb torch.cat((freqs, freqs), dim-1) # [max_len, dim] self.register_buffer(cos_cached, emb.cos(), persistentFalse) self.register_buffer(sin_cached, emb.sin(), persistentFalse) def _rotate_half(self, x: torch.Tensor) - torch.Tensor: x1 x[..., : x.shape[-1] // 2] x2 x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim-1) def forward(self, x: torch.Tensor, seq_len: int) - torch.Tensor: # x 形状可能是 [B, N_h, S, D_R] 或 [B, S, D_R] cos self.cos_cached[:seq_len].to(x.dtype) sin self.sin_cached[:seq_len].to(x.dtype) if x.ndim 4: cos cos.unsqueeze(0).unsqueeze(1) # [1, 1, S, D_R] sin sin.unsqueeze(0).unsqueeze(1) elif x.ndim 3: cos cos.unsqueeze(0) # [1, S, D_R] sin sin.unsqueeze(0) return (x * cos) (self._rotate_half(x) * sin) class MultiHeadLatentAttention(nn.Module): DeepSeek Multi-Head Latent Attention (MLA) 模块 包含训练模式显式解压 K/V和推理模式矩阵吸收 Matrix Absorption def __init__( self, d_model: int 4096, # 隐层维度 D n_heads: int 128, # 注意力头数 N_h d_head: int 128, # 每个头的语义维度 D_h d_c_kv: int 512, # KV 潜隐压缩维度 D_c d_c_q: int 1536, # Q 潜隐压缩维度 D_c d_rope: int 64, # RoPE 解耦位置维度 D_R ): super().__init__() self.d_model d_model self.n_heads n_heads self.d_head d_head self.d_c_kv d_c_kv self.d_c_q d_c_q self.d_rope d_rope # 点积缩放因子 (语义维度 位置维度) self.scale (d_head d_rope) ** -0.5 # 1. Query 投影权重 self.w_dq nn.Linear(d_model, d_c_q, biasFalse) # Q 降维 [D - D_c] self.w_uq nn.Linear( d_c_q, n_heads * d_head, biasFalse ) # Q 语义升维 [D_c - N_h * D_h] self.w_qr nn.Linear( d_c_q, n_heads * d_rope, biasFalse ) # Q 位置升维 [D_c - N_h * D_R] # 2. Key / Value 投影权重 # 联合降维同时生成 KV 潜隐向量 (D_c) 与 共享位置 Key (D_R) self.w_dkv nn.Linear( d_model, d_c_kv d_rope, biasFalse ) # [D - D_c D_R] self.w_uk nn.Linear( d_c_kv, n_heads * d_head, biasFalse ) # K 语义升维 [D_c - N_h * D_h] self.w_uv nn.Linear( d_c_kv, n_heads * d_head, biasFalse ) # V 语义升维 [D_c - N_h * D_h] # 3. 输出投影权重 self.w_o nn.Linear( n_heads * d_head, d_model, biasFalse ) # [N_h * D_h - D] # 位置编码器 self.rope RoPE(dimd_rope) def forward_training( self, x: torch.Tensor, mask: torch.Tensor None ) - torch.Tensor: 训练前向传播显式解压形态 输入: x: [B, S, D_model] 输出: out: [B, S, D_model] B, S, _ x.shape # --- 1. Query 处理 --- c_q self.w_dq(x) # [B, S, D_c_q] q_c self.w_uq(c_q).view( B, S, self.n_heads, self.d_head ) # [B, S, N_h, D_h] q_c q_c.transpose(1, 2) # [B, N_h, S, D_h] q_r self.w_qr(c_q).view( B, S, self.n_heads, self.d_rope ) # [B, S, N_h, D_R] q_r q_r.transpose(1, 2) # [B, N_h, S, D_R] q_r self.rope(q_r, seq_lenS) # [B, N_h, S, D_R] (应用 RoPE) # 拼接语义与位置 Query q torch.cat([q_c, q_r], dim-1) # [B, N_h, S, D_h D_R] # --- 2. Key Value 处理 --- compressed_kv self.w_dkv(x) # [B, S, D_c_kv D_R] c_kv, k_r torch.split( compressed_kv, [self.d_c_kv, self.d_rope], dim-1 ) # c_kv: [B, S, D_c_kv] # k_r: [B, S, D_R] # 处理共享位置 Key k_r self.rope(k_r, seq_lenS) # [B, S, D_R] k_r k_r.unsqueeze(1).expand( B, self.n_heads, S, self.d_rope ) # [B, N_h, S, D_R] # 还原语义 Key 与 Value k_c self.w_uk(c_kv).view( B, S, self.n_heads, self.d_head ) # [B, S, N_h, D_h] k_c k_c.transpose(1, 2) # [B, N_h, S, D_h] v self.w_uv(c_kv).view( B, S, self.n_heads, self.d_head ) # [B, S, N_h, D_h] v v.transpose(1, 2) # [B, N_h, S, D_h] # 拼接语义与位置 Key k torch.cat([k_c, k_r], dim-1) # [B, N_h, S, D_h D_R] # --- 3. Attention 计算 --- scores ( torch.matmul(q, k.transpose(-2, -1)) * self.scale ) # [B, N_h, S, S] if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) # [B, N_h, S, S] # --- 4. 输出投影 --- out torch.matmul(attn_weights, v) # [B, N_h, S, D_h] out ( out.transpose(1, 2) .contiguous() .view(B, S, self.n_heads * self.d_head) ) # [B, S, N_h * D_h] out self.w_o(out) # [B, S, D_model] return out def forward_inference_absorbed( self, x: torch.Tensor, mask: torch.Tensor None ) - torch.Tensor: 推理前向传播矩阵吸收形态 KV Cache 中仅需保存: 1. c_kv: [B, S, D_c_kv] 2. k_r: [B, S, D_R] 全程不解压产生高维的 K 和 V 向量 B, S, _ x.shape # --- 1. Query 生成 --- c_q self.w_dq(x) # [B, S, D_c_q] q_r self.w_qr(c_q).view( B, S, self.n_heads, self.d_rope ) # [B, S, N_h, D_R] q_r q_r.transpose(1, 2) # [B, N_h, S, D_R] q_r self.rope(q_r, seq_lenS) # [B, N_h, S, D_R] # --- 2. 存入 KV Cache 的两项低维数据 --- compressed_kv self.w_dkv(x) # [B, S, D_c_kv D_R] c_kv, k_r torch.split( compressed_kv, [self.d_c_kv, self.d_rope], dim-1 ) # c_kv: [B, S, D_c_kv] - 显存只需存这个 # k_r: [B, S, D_R] - 显存只需存这个 k_r_pe self.rope(k_r, seq_lenS) # [B, S, D_R] k_r_pe k_r_pe.unsqueeze(1).expand( B, self.n_heads, S, self.d_rope ) # [B, N_h, S, D_R] # --- 3. 矩阵吸收求 Query: 将 W_UQ 与 W_UK 融合成 W_q_absorbed --- # w_uq 权重 reshape: [N_h, D_h, D_c_q] w_uq self.w_uq.weight.view(self.n_heads, self.d_head, self.d_c_q) # w_uk 权重 reshape: [N_h, D_h, D_c_kv] w_uk self.w_uk.weight.view(self.n_heads, self.d_head, self.d_c_kv) # W_q_absorbed W_UQ^T W_UK - 形状 [N_h, D_c_q, D_c_kv] w_q_absorbed torch.einsum(hdq, hdk - hqk, w_uq, w_uk) # q_absorbed: 用 c_q 直接乘以吸收矩阵生成低维 Query形状 [B, N_h, S, D_c_kv] q_absorbed torch.einsum(bsq, hqk - bhsk, c_q, w_q_absorbed) # 语义得分: 拿低维 q_absorbed 直接与压缩包 c_kv 点积没有解压 K scores_content torch.matmul( q_absorbed, c_kv.unsqueeze(1).transpose(-2, -1) ) # [B, N_h, S, S] # 位置得分 scores_position torch.matmul( q_r, k_r_pe.transpose(-2, -1) ) # [B, N_h, S, S] scores (scores_content scores_position) * self.scale if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) # [B, N_h, S, S] # --- 4. 矩阵吸收求 Output: 将 W_UV 与 W_O 融合成 W_o_absorbed --- # w_uv 权重 reshape: [N_h, D_h, D_c_kv] w_uv self.w_uv.weight.view(self.n_heads, self.d_head, self.d_c_kv) # w_o 权重 reshape: [D_model, N_h, D_h] w_o self.w_o.weight.view(self.d_model, self.n_heads, self.d_head) # W_o_absorbed W_UV^T W_O - 形状 [N_h, D_c_kv, D_model] w_o_absorbed torch.einsum(hdk, mhd - hkm, w_uv, w_o) # 对 KV Cache 中的潜隐向量 c_kv 进行注意力加权形状 [B, N_h, S, D_c_kv] attn_c_kv torch.matmul(attn_weights, c_kv.unsqueeze(1)) # 一步到位投影回 D_model 维度并对多头求和没有解压 V out torch.einsum(bhsk, hkm - bsm, attn_c_kv, w_o_absorbed) # [B, S, D_model] return out # 测试与结果验证 if __name__ __main__: # 参数设置 (缩减版配置方便本地验证) B, S, D 2, 16, 512 N_h, D_h 8, 64 D_c_kv 128 D_c_q 256 D_R 32 x torch.randn(B, S, D) mla MultiHeadLatentAttention( d_modelD, n_headsN_h, d_headD_h, d_c_kvD_c_kv, d_c_qD_c_q, d_ropeD_R, ) # 1. 训练模式前向传播 out_train mla.forward_training(x) # 2. 推理矩阵吸收模式前向传播 out_infer mla.forward_inference_absorbed(x) print(f输入 Tensor 形状: {x.shape}) print(f训练模式输出 形状: {out_train.shape}) print(f推理模式输出 形状: {out_infer.shape}) # 数学等价性校验在浮点数误差范围内对比 diff torch.max(torch.abs(out_train - out_infer)).item() print(f训练模式与推理模式输出最大绝对误差: {diff:.6f}) assert torch.allclose(out_train, out_infer, atol1e-4), 两套前向传播计算结果不一致 print(验证通过矩阵吸收逻辑与显式解压完全等价)MLA的常见问题1.既然将 128 个头的 KV 压缩成了 512 维的潜隐向量为什么不会丢失多头表达能力多头的表达力被转移到了 Query 侧和静态权重矩阵中而非在显存缓存中硬存 128 份KV.在推理阶段通过离线计算将升维矩阵融合吸收为.Query 依然能映射出 128 个独立且相互不干涉的子空间得分矩阵仍然保持 [B, 128, S, S]的 128 头完整粒度2.为什么 RoPE 位置编码不能直接作用在 512 维的潜隐向量MLA 是如何解决的破坏矩阵结合律RoPE 是随位置 m, n 动态变化的旋转矩阵如果作用在动态矩阵会加载和中间导致和无法相乘融合污染 VKV共享一个如果直接对施加旋转那么V也会带上位置旋转破坏语义特征MLA解决方案建立 64 维的独立“位置专线”,存入 Cache 专门算相对位置得分与 512 维潜隐空间算出的语义得分直接相加。3.MLA 的显存优化属于一种“空间换空间”的工程权衡它的 Trade-off 代价是什么代价静态权重显存Weight Memory增加。因为吸收后的融合矩阵维度为 [128, 512, 5120]单层静态模型权重增加了约37%以 60 层 FP16 模型为例全卡多占用 约 15 GB 静态显存。收益动态 KV Cache 显存暴降98%静态权重的增加是常数级 O(1)而 KV Cache 的节省是线性/平方级 O(B *S)在长文本与大并发推理中这个交换极其划算4.MLA 在结合 FlashAttention / PagedAttention 算子编写时有哪些难点与优化点非对称维度计算语义点积维度是 512位置点积维度是 64无法像传统 FA2 那样直接按照标准的 Tile Size如 128x128划分 Block需要专门设计 Block 布局