在深度学习模型日益庞大、复杂的当下,推理空间占用的优化显得尤为关键,而MLA 凭借其独特的设计,在这方面表现得极为出色。其通过先进的压缩算法和巧妙的数据处理方式,将原本需要大量空间存储的键( K )和值( V )信息进行高效压缩,从而大幅降低了模型在推理过程中对内存和显存的占用。
节省推理空间占用只是MLA 带来的基础益处,由此还引发了一系列积极的连锁反应。由于占用空间的减少,模型在推理时的数据读取和传输速度得到了显著提升。这就好比在一条原本拥堵的道路上,车辆(数据)能够更加顺畅地行驶,减少了等待和延误,从而加快了整个推理流程。推理速度的提升意味着模型能够更快地给出预测结果,这对于一些对实时性要求极高的应用场景,如自动驾驶、在线实时翻译等,具有至关重要的意义。
此外,MLA 节省推理空间占用的特性还为模型的部署提供了更大的灵活性。在资源有限的设备上,如移动设备、嵌入式系统等,传统的注意力模型可能因空间占用过大而无法顺利运行。而MLA 则打破了这一限制,使得这些设备也能够轻松承载和运行复杂的深度学习模型,极大地拓展了模型的应用 范围。
在接下来的内容中,我们将深入探究MLA 是如何实现节省推理空间占用的具体机制,包括其背后的算法原理、关键技术的运用等。同时,我们还将通过实际案例和实验数据,直观地展示MLA 在节省推理空间占用以及提升推理效率方面的卓越表现,让读者对这一创新技术有更全面、更深入的认识。
9.2.1 优化的MLA 模型实现1 :压缩低秩空间
MLA 的创新之一就是低秩空间压缩,我们通过将输入的向量压缩到低秩维度从而完成了维度变换。而在具体实现上,我们则可以通过定义和组成对应的可训练参数完成模型的结构。
1. 变换参数的定义
#首先我们需要定义MLA中的参数,在这里我们定义参数如下: self.q_proj_up = d_model * 2 #先把Query的维度升高 self.qk_proj_down = d_model #再降低Query的维度进行计算 self.query_head_dim = self.qk_proj_down // n_heads #这里是进行压缩kv的维度,可以设置为原有d_model大小的2/3为好 self.kv_lora_rank = int((d_model * 3)//2) #变为hidden压缩了维度
在上面代码中,我们分别设置了MLA 中的参数,q_proj_up 与qk_proj_down 是对输入的查询( Q ) 进行计算,我们采用双参数期望能够在输入的参数较少(在推理时只有1 个token )的情况下,也能获得一个较为合理的特征向量。
kv_lora_rank 对整体的Key 和Value 整合计算维度,其作用是建立一个过度的低维处理向量,用于后期的整合。
2. Query 的维度变换处理
下面我们首先看Query 的维度变换,代码如下所示:
self.W_dq = torch.nn.Parameter(0.01 * torch.randn(d_model, self.q_proj_dim)) self.q_norm = torch.nn.LayerNorm(self.q_proj_dim) self.W_Uq = torch.nn.Parameter(0.01 * torch.randn(self.q_proj_dim, self.qk_proj_dim))
在这里,我们首先依据公式定义了多个映射参数,旨在对输入的特征向量 x 进行变换。其中,qk_proj_dim 与qk_head_num 的设定是为了调整维度的尺寸,以便我们人为地修正输入和输出的维度。在进行qk 计算时,我们仅需确保最后一个维度保持一致即可。
而Query 的计算代码如下所示:
compressed_q = x @ self.W_dq compressed_q = self.q_norm(compressed_q) query = compressed_q @ self.W_Uq q_nope = self._split_heads(query, self.num_heads, self.query_head_dim) # 将query分割为多头 shape = [-1,6,48,64]
3. Key 与Value 的维度变换处理
对于Key 和Value 的维度变换与处理,我们同样按推理公式中首先对其进行变换,之后将其变换为输入维度,代码如下所示:
q_absorb = torch.einsum('bqd,hdc->bhqc', compressed_kv, self.kv_b_proj[:,:,
:self.query_head_dim])
这里我们首先完成了Key 与Value 的总体维度变换,将其拆分后我们所需要的Key 和Value 维度矩阵,计算过程如下:
v_nope = torch.einsum('bqc,hcd->bqhd', compressed_kv, self.kv_b_proj[:,:,self.query_head_dim:])
attn_output = torch.einsum('bhql,blhd->bhqd', attn_weights,v_nope)
顺便说一下,MLA 中的低秩压缩的核心思想是:我们可以用一个较小的矩阵来近似表示一个大矩阵。而这个小矩阵正常来说是通过矩阵分解得到的,但在MLA 中是直接设计小矩阵的维度,让它们通过训练自动学习到一个好的低秩表示。
秩就是信息压缩的维度,我们通过一个具体的例子来理解这个过程:想象一个5120 维的向量,它可能代表了一幅图像的所有像素值。当我们用一个秩为512 的变换去处理它时,本质上我们是在说:这5120 个数字中,实际上可以用512 个独立的特征来表达主要信息。
而这种压缩之所以有效,是因为在实际应用中:
l 真实数据通常存在大量冗余。
l 并非所有维度都同等重要。
因此在实际中保留最重要的维度往往足以表达数据的主要特征。
9.2.2 优化的MLA 模型实现2 :核心注意力矩阵计算
注意力模型MLA 核心优化之一就是对核心计算使用矩阵计算的方法,在这里我们按公式完成了新的注意力计算优化,代码如下所示:
def _attn(self, q_nope, compressed_kv, attention_mask=None):
q_cope = q_nope.clone()
# 为了与compressed_kv结合计算attention_score
if True:
q_absorb = torch.einsum('bqd,hdc->bhqc', compressed_kv, self.kv_b_proj[:,:,:self.query_head_dim])
attn_weights = torch.einsum('bhqc,bhlc->bhlq', q_absorb, q_nope)
else:
pass
# 缩放注意力权重
attn_weights = attn_weights / torch.full([], q_nope.size(-1) ** 0.5, dtype=attn_weights.dtype,device=attn_weights.device)
query_length, key_length = q_nope.size(-2), compressed_kv.size(-2)
causal_mask = self.bias[:, :, key_length - query_length: key_length, :key_length].to(q_nope.device)
attn_weights = torch.where(causal_mask, attn_weights.to(attn_weights.dtype), self.mask_value)
if attention_mask is not None:
# 如果有额外的注意力掩码,应用它
attn_weights = attn_weights + attention_mask
attn_weights += self.cope(q_cope, attn_weights)
attn_weights = torch.nn.functional.softmax(attn_weights, dim=-1) # 计算Softmax
attn_weights = attn_weights.type(compressed_kv.dtype)
attn_weights = self.attn_dropout(attn_weights) # 应用dropout
if True:
v_nope = torch.einsum('bqc,hcd->bqhd', compressed_kv, self.kv_b_proj[:,:,self.query_head_dim:])
attn_output = torch.einsum('bhql,blhd->bhqd', attn_weights,v_nope)
else:
pass
return attn_output, attn_weights
在上面代码中,我们分别对两个部分进行优化,首先是attn_weights 的计算,而同样的attn_output 部分我们也进行优化,将维度变换过程整合到一个完整过程中实现。另外,读者在具体操作时,可能注意到代码的最后部分使用了if …else 条件语句进行判断。这是由于在这个位置进行维度变换时可以将if 条件下的两个变换过程整合成一个。这样的计算过程虽然会节省时间,但是这种变换对于硬件资源消耗较大,部分GPU 并不适合这种直接维度变换的操作,因此我们建议有兴趣的读者可以自行尝试。
9.2.3 优化的MLA 模型实现3 :对显存KV Cache 部分的压缩
与传统的我们分别对生成的Key 和Value 值进行缓存不同,MLA 中首先对Key 和Value 进行一个整体变换,将其压缩后通过维度计算重新获取对应的Key 和Value 值。在这个过程中我们可以通过缓存中间过程,也就是压缩的整体值从而完成对值的缓存。代码如下所示:
if layer_past is not None: current_kv = x @ self.W_dkv compressed_kv = torch.cat([layer_past, current_kv], dim=1) else: compressed_kv = x @ self.W_dkv present = compressed_kv #在这里进行了键值对的压缩 compressed_kv = compressed_kv@self.W_duv compressed_kv = self.kv_norm(compressed_kv)
在这里我们完成了对键值对的联合压缩过程,这也是MLA 的创新,我们可以将其理解为:
l x :是输入的原始向量。
l W_dkv :是一个特殊的下投影矩阵,它同时服务于键和值的压缩。
l compressed_kv :是压缩后的潜在向量,它包含了键和值共享的信息。
compressed_kv 作为一个存储关键信息供模型在计算时快速访问它包含了键和值的压缩信息,这个向量的维度一般认为比原始的键值对小得多,因此存储这个压缩向量比存储完整的键和值更节省空间。
9.2.4 带有缓存的MLA 注意力模型完整实现
我们已经对MLA 模型的各个模块进行了详细的讲解,而MLA 注意力机制的一个显著优点,就是其在推理阶段能够出色地完成缓存的优化与利用。这一特性不仅提升了模型的运行效率,还为实际应用带来了更多的便利性。
接下来,本小节将着手构建一个带有缓存功能的MLA 完整模型。通过整合缓存机制,我们将进一步展示MLA 模型在实际应用中的优势。以下是实现这一完整模型的代码示例:
class CoPEMLA(torch.nn.Module):
def __init__(self, config, layer_idx=None):
super().__init__()
self.config = config
self.max_position_embeddings = max_positions = config.max_position_embeddings
# 创建一个下三角矩阵,用于因果掩码(causal mask)
self.bias = torch.tril(torch.ones((max_positions, max_positions), dtype=torch.bool)).view(1, 1, max_positions,max_positions)
self.mask_value = torch.tensor(-1E+9) # 用于掩码的极大负值
# 维度参数
d_model = config.hidden_size
self.d_model = torch.tensor(d_model)
self.num_heads = n_heads = config.num_attention_heads
# 投影维度
self.q_proj_up = d_model * 2 # 先把Query的维度升高
self.qk_proj_down = d_model # 再降低Query的维度进行计算
self.query_head_dim = self.qk_proj_down // n_heads
# 这里是进行压缩kv的维度,可以设置为原有d_model大小的2/3为好
self.kv_lora_rank = int((d_model * 3)//2) # 变为hidden压缩了维度
self.qk_nope_head_dim = self.query_head_dim * 2
self.kv_attn_dim = (self.query_head_dim + self.qk_nope_head_dim)
# Q投影
self.W_dq = torch.nn.Parameter(0.01 * torch.randn(d_model, self.q_proj_up))
self.q_norm = torch.nn.RMSNorm(self.q_proj_up)
self.W_Uq = torch.nn.Parameter(0.01 * torch.randn(self.q_proj_up, self.qk_proj_down))
# KV投影
self.W_dkv = torch.nn.Parameter(0.01 * torch.randn((d_model), (self.kv_lora_rank)))
self.W_duv = torch.nn.Parameter(0.01 * torch.randn((self.kv_lora_rank), (self.kv_lora_rank)))
self.kv_norm = torch.nn.RMSNorm((self.kv_lora_rank))
self.kv_b_proj = torch.nn.Parameter(0.01 * torch.randn(size=(self.num_heads, (self.kv_attn_dim), self.kv_lora_rank)))
# 输出
self.W_o = torch.nn.Parameter(0.01 * torch.randn(d_model, d_model))
self.attn_dropout = torch.nn.Dropout(config.dropout)
#kq_weight要参与多头后的query与key计算
self.uk_weight = torch.torch.nn.Parameter(0.01 * torch.randn(self.num_heads,self.kv_lora_rank, self.query_head_dim))
#v_proj_weight要参与value计算,将压缩的内容重新映射回value
self.uv_weight = torch.torch.nn.Parameter(0.01 * torch.randn(self.num_heads,self.kv_lora_rank, self.query_head_dim))
self.cope = updat_moudle.CoPE(config.max_position_embeddings, self.query_head_dim)
def forward(self, x, layer_past=None, attention_mask=None):
# Q投影
compressed_q = x @ self.W_dq
compressed_q = self.q_norm(compressed_q)
query = compressed_q @ self.W_Uq
q_nope = self._split_heads(query, self.num_heads, self.query_head_dim)
# KV投影
if layer_past is not None:
current_kv = x @ self.W_dkv
compressed_kv = torch.cat([layer_past, current_kv], dim=1)
else:
compressed_kv = x @ self.W_dkv
present = compressed_kv #在这里进行了键值对的压缩
compressed_kv = compressed_kv@self.W_duv
compressed_kv = self.kv_norm(compressed_kv)
# 计算注意力输出和注意力权重
attn_output, attn_weights = self._attn(q_nope, compressed_kv,attention_mask)
attn_output = self._merge_heads(attn_output, self.num_heads, self.query_head_dim) # 合并多头
attn_output = attn_output @ self.W_o
attn_output = self.attn_dropout(attn_output) # 应用dropout
outputs = (attn_output, present) # 返回注意力输出和当前的key、value
return outputs
def _split_heads(self, tensor, num_heads, attn_head_size):
"""
将隐藏层维度分割为多头注意力的头和头的大小。
"""
new_shape = tensor.size()[:-1] + (num_heads, attn_head_size)
tensor = tensor.view(new_shape)
return tensor.permute(0, 2, 1, 3) # (batch, head, seq_length, head_features)
def _attn(self, q_nope, compressed_kv, attention_mask=None):
q_cope = q_nope.clone()
#为了与compressed_kv结合计算attention_score
if True:
q_absorb = torch.einsum('bqd,hdc->bhqc', compressed_kv, self.uk_weight)
attn_weights = torch.einsum('bhqc,bhlc->bhlq', q_absorb, q_nope)
else:
pass
# 缩放注意力权重
attn_weights = attn_weights / torch.full([], q_nope.size(-1) ** 0.5, dtype=attn_weights.dtype,device=attn_weights.device)
query_length, key_length = q_nope.size(-2), compressed_kv.size(-2)
causal_mask = self.bias[:, :, key_length - query_length: key_length, :key_length].to(q_nope.device)
attn_weights = torch.where(causal_mask, attn_weights.to(attn_weights.dtype), self.mask_value)
if attention_mask is not None:
# 如果有额外的注意力掩码,应用它
attn_weights = attn_weights + attention_mask
attn_weights += self.cope(q_cope, attn_weights)
attn_weights = torch.nn.functional.softmax(attn_weights, dim=-1) # 计算softmax
attn_weights = attn_weights.type(compressed_kv.dtype)
attn_weights = self.attn_dropout(attn_weights) # 应用dropout
if True:
v_nope = torch.einsum('bqc,hcd->bqhd', compressed_kv, self.uv_weight)
attn_output = torch.einsum('bhql,blhd->bhqd', attn_weights,v_nope)
else:
#下面这个压缩代码在训练时会报错
pass
return attn_output, attn_weights
def _merge_heads(self, tensor, num_heads, attn_head_size):
"""
将多头注意力的头和头的大小合并回隐藏层维度。
"""
tensor = tensor.permute(0, 2, 1, 3).contiguous()
new_shape = tensor.size()[:-2] + (num_heads * attn_head_size,)
return tensor.view(new_shape)
上面代码完整实现了标准的MLA 注意力模型。其中加粗的部分代码为注意力计算模块,在具体使用上读者可以将MLA 替代原有我们的上一章对比计算中的经典多头注意力模型,读者可以自行尝试。
本文节选自《DeepSeek原生应用与智能体开发实践》一书,获出版社和作者授权发布,仅供读者个人学习使用。
