单头注意力、MHA和MLA详解

目录

  • Self-Attention单头注意力
    • [Single Head Attention代码](#Single Head Attention代码)
    • [Single Head Attention流程图](#Single Head Attention流程图)
    • 掩码矩阵的形式
  • [多头注意力MHA(Multi-Head Attention)](#多头注意力MHA(Multi-Head Attention))
    • [Multi-Head Attention代码](#Multi-Head Attention代码)
    • MHA流程图
  • [线性注意力(Linear Attention)](#线性注意力(Linear Attention))
    • [Linear Attention与普通 Self-Attention 的对比](#Linear Attention与普通 Self-Attention 的对比)
  • [MLA(Multi-head Latent Attention)](#MLA(Multi-head Latent Attention))
    • 背景知识
      • [1. 低秩分解](#1. 低秩分解)
      • [2. ColumnParallelLinear和RowParallelLinear函数](#2. ColumnParallelLinear和RowParallelLinear函数)
        • [2.1 ColumnParallelLinear](#2.1 ColumnParallelLinear)
        • [2.2 RowParallelLinear](#2.2 RowParallelLinear)
        • [2.3 通过一个例子来解释](#2.3 通过一个例子来解释)
          • [2.3.1 先看"不并行"时,正常是怎么算的](#2.3.1 先看“不并行”时,正常是怎么算的)
          • [2.3.2 现在切开(RowParallelLinear 登场)](#2.3.2 现在切开(RowParallelLinear 登场))
          • [2.3.3 为什么叫"部分和(Partial Sum)"?](#2.3.3 为什么叫“部分和(Partial Sum)”?)
          • [2.3.4 救星登场:All-Reduce 通信](#2.3.4 救星登场:All-Reduce 通信)
          • [2.3.5 最终结果:](#2.3.5 最终结果:)
        • [2.4 为什么ColumnParallelLinear和RowParallelLinear是"天生一对"?](#2.4 为什么ColumnParallelLinear和RowParallelLinear是“天生一对”?)
    • [Multi-head Latent Attention代码](#Multi-head Latent Attention代码)

Self-Attention单头注意力

Single Head Attention代码

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

class SelfAttention(nn.Module):
    def __init__(self, embed_size):
        super(SelfAttention, self).__init__()
        
        # embed_size:输入嵌入向量的维度,即每个 token 的特征向量长度。
        self.embed_size = embed_size
        
        # 三个线性变换层
        # 创建Value投影矩阵W^V。输入输出维度都是 embed_size,没有偏置 bias=False。
        # 将输入 x 映射为 "值" 向量,代表每个token实际携带的信息内容。
        self.values = nn.Linear(embed_size, embed_size, bias=False)

        # 创建Key投影矩阵W^K。将输入 x 映射为 "键" 向量,代表每个 token "是什么",用于被查询匹配。
        self.keys = nn.Linear(embed_size, embed_size, bias=False)

        # 创建Query投影矩阵W^Q。将输入 x 映射为 "查询" 向量,代表每个token"在找什么",用于去匹配其他token的key。
        self.queries = nn.Linear(embed_size, embed_size, bias=False)

        # 直观理解:
        # 想象你在图书馆找书 ------ Query 是你的需求("我想看科幻"),Key 是每本书的标签("科幻/文学/历史"),
        # Value 是书的实际内容。Query 和 Key 匹配度越高,对应的 Value 就越被关注。

    # 前向传播函数。参数:
    # x:输入张量,形状 (N, seq_length, embed_size)
    # mask:可选的掩码张量,用于屏蔽某些位置(如 padding 或因果掩码)
    def forward(self, x, mask):
    
        # 解构输入 x 的形状:
        # N:batch size(一个批次有多少个样本)
        # seq_length:序列长度(每个样本有多少个token)
        # _: 丢弃 embed_size(不需要显式使用)
        N, seq_length, _ = x.size() 

        # Q、K、V的投影计算
        # x的shape原本为(N, seq_length, embed_size)
        # 经过后计算得到的queries的shape仍然是(N, seq_length, embed_size)
        queries = self.queries(x) 
        
        # keys和values的计算与queries类似
        keys = self.keys(x)        # shape仍然是(N, seq_length, embed_size)
        values = self.values(x)    # shape仍然是(N, seq_length, embed_size)

        # 计算注意力分数
        # 计算每个 query 与所有 key 的相似度
        # keys.transpose(1, 2):将keys从(N, seq_length, embed_size)转置为(N, embed_size, seq_length),实现行列互换。
        # torch.bmm(batch matrix multiply):批量矩阵乘法。对于每个batch的矩阵进行运算:
        # queries形状(seq_length, embed_size)点乘keys^T形状(embed_size, seq_length)
        # 结果形状(seq_length, seq_length) ---- 即每个 token 对每个 token 的注意力原始分数
        # 最终attention_scores 形状为 (N, seq_length, seq_length)
        attention_scores = torch.bmm(queries, keys.transpose(1, 2))

        # 掩码处理
        # mask == 0:找出 mask 中值为 0 的位置(表示这些位置需要被屏蔽),
        # 将这些位置的注意力分数设为负无穷大(用 -1e20 近似)
        # 为什么这样做? 
        # 因为后续的 softmax 会将 e^(-∞) ≈ 0,使得被屏蔽位置的注意力权重趋近于 0,从而实现"忽略这些位置"的效果。
        if mask is not None:
            attention_scores = attention_scores.masked_fill(mask == 0, float('-1e20'))
        
        # (self.embed_size ** 0.5)也就是√d_k
        # 当d_k很大时,点积结果方差变大,softmax会趋向于极端的one-hot分布,梯度很小。除以√d_k可以缓解这个问题。
        # F.softmax(..., dim=-1)沿最后一个维度(即每个query对所有key的那一行)做softmax
        # 将缩放后的分数转化为概率分布(每个token对其他token的关注权重,和为1)
        # 结果attention形状为 (N, seq_length, seq_length),即注意力权重矩阵。
        attention = F.softmax(attention_scores / (self.embed_size ** 0.5), dim=-1)

        # 加权求和(输出)
        # 用注意力权重对 values 做加权求和:
        # attention 形状 (N, seq_length, seq_length)----权重
        # values 形状 (N, seq_length, embed_size)----被加权的内容
        # torch.bmm 结果形状 (N, seq_length, embed_size)
        output = torch.bmm(attention, values)
        
        # 在每个batch里,attention与values的点乘,也就是:
        # token i = 所有token的value向量按注意力权重加起来的和。
        return output

Single Head Attention流程图

掩码矩阵的形式

因果掩码是一个上三角矩阵,对角线及以下为1(可见),对角线以上为0(屏蔽):

多头注意力MHA(Multi-Head Attention)

Multi-Head Attention代码

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

class MultiHeadAttention(nn.Module):
    def __init__(self, embed_size, num_heads):
        super(MultiHeadAttention, self).__init__()
        # 嵌入维度大小
        self.embed_size = embed_size
        # 注意力头的数量
        self.num_heads = num_heads
        # 每个注意力头的维度。例如 embed_size=512, num_heads=8 → head_dim=64。
        # 核心思想:把一个大维度拆成多个小维度,每个头在小空间里独立做注意力。
        self.head_dim = embed_size // num_heads
        
        # 确保嵌入尺寸可以被头数整除
        assert (
            self.head_dim * num_heads == embed_size
        ), "嵌入尺寸需要能够被头数整除"

        # 为多头注意力机制定义线性层
        self.values = nn.Linear(embed_size, embed_size, bias=False)
        self.keys = nn.Linear(embed_size, embed_size, bias=False)
        self.queries = nn.Linear(embed_size, embed_size, bias=False)
        
        # 定义一个线性层用于合并多个头的输出,这个有bias(默认bias=True)
        self.fc_out = nn.Linear(embed_size, embed_size)

    def forward(self, x, mask):
        N = x.shape[0]  # 批量大小
        seq_length = x.shape[1]  # 序列长度
        
        # x的shape是(N, seq_length, embed_size)
        # 将输入x通过线性变换后分割成num_heads份,每份是head_dim维。
        # x的shape (N, seq_length, embed_size)最后一个维度embed_size被拆分成了 (num_heads, head_dim)
        # 变换后的shape是(N, seq_length, num_heads, head_dim)
        values = self.values(x).view(N, seq_length, self.num_heads, self.head_dim)
        keys = self.keys(x).view(N, seq_length, self.num_heads, self.head_dim)
        queries = self.queries(x).view(N, seq_length, self.num_heads, self.head_dim)
        
        # 把num_heads提前到第二维,这样每个头就是独立的 (seq_len, head_dim) 矩阵
        # 后续可以并行计算所有头的注意力,互不干扰。
        # 从(N, seq_length, num_heads, head_dim) 变为 (N, num_heads, seq_length, head_dim)
        values = values.permute(0, 2, 1, 3)  
        keys = keys.permute(0, 2, 1, 3)
        queries = queries.permute(0, 2, 1, 3)
        
        # 计算缩放点积注意力得分(爱因斯坦求和)
        # n: batch维度N  q:query length  k:key length  h:head  d:head_dim
        # 在 d 上做内积 → 结果 nhqk:(N, heads, Q_seq, K_seq)  相对于前面的单头的注意力,多了一个heads维度
        attention_scores = torch.einsum("nhqd,nhkd->nhqk", [queries, keys])
        if mask is not None:
            # 如果有掩码,则应用掩码到注意力得分上
            attention_scores = attention_scores.masked_fill(mask.unsqueeze(1) == 0, float('-1e20'))
            
        # 应用softmax函数并根据head_dim开平方根进行缩放
        attention = F.softmax(attention_scores / (self.head_dim ** 0.5), dim=-1)
        
        # 将注意力权重与值向量相乘
        # attention的shape (N, num_heads, Q_seq_length, K_seq_length)  
        # values的shape (N, num_heads, K_seq_length, head_dim)
        #  在K维度上求和 → nqhd:(N, Q, heads, D)
        out = torch.einsum("nhqk,nhkd->nqhd", [attention, values]).reshape(
            N, seq_length, self.embed_size
        )
        
        # 通过线性层组合来自所有头的信息
        out = self.fc_out(out)
        return out

MHA流程图

线性注意力(Linear Attention)

Linear Attention与普通 Self-Attention 的对比

MLA(Multi-head Latent Attention)

近些年,Attention的模型出现了很多改进版本,但主要思想是,减少KV的数量,达到减少KV Cache显存的使用量,MLA的思想是对KV的维度降维,实现减少对显存的占用率。

背景知识

1. 低秩分解

在讲解MLA之前,先了解一下什么是低秩分解

举一个例子:

复制代码
直接用一个矩阵做投影:
输入 (5120维) × W_big (5120×24576)──▶  输出 (24576维)
                         ↑
                    W_big有1.26亿个参数!
                    
这太重了。

低秩分解的想法:这个大矩阵可以近似为两个小矩阵的乘积:
W_big ≈ W2 × W1
(5120×24576) ≈  (5120×1536) × (1536×24576) 

由标准模式的1步转变为低秩模式的2步,参数量由1.26亿降为4560万。

复制代码
标准模式(1 步)
x (5120)  ────────── wq ──────────▶  Q (24576)
               5120 × 24576
               参数量 ≈ 1.26 亿
               
低秩模式(2 步)
x (5120) ── wq_a ──▶  q_compressed (1536) ── RMSNorm ── wq_b ──▶  Q (24576)
         5120×1536                                       1536×24576
         ≈ 787 万                                             ≈ 3775 万
                         └────── 合计 ≈ 4560 万 ──────┘

为什么叫"低秩"?

任何一个大矩阵 W (5120×24576),如果它的秩不高(即信息有大量冗余),就可以用两个小矩阵的乘积来逼近。

中间维度 r 就叫秩(rank),这里 r = q_lora_rank = 1536。秩越小,压缩越狠,参数越少,但逼近精度越差。1536 是 DeepSeek 找到的平衡点。

直觉理解:

  • wq_a 把信息"收拢"到一个低维瓶颈 1536
  • wq_b 再把信息"展开"到各头 24576
  • 瓶颈层迫使模型只保留最关键的信息,丢弃冗余
  • 类似于自动编码器的 Encoder-Decoder 结构

2. ColumnParallelLinear和RowParallelLinear函数

ColumnParallelLinear 和 RowParallelLinear 是分布式大模型训练和推理中(特别是基于 Megatron-LM 提出的张量并行 Tensor Parallelism, TP)极其重要的两个核心算子。

它们的作用是将一个巨大的线性层(全连接层)的参数矩阵切分到多张 GPU 上,使得每张 GPU 只需要计算和存储一部分权重,从而突破单张显存的瓶颈。这两个函数通常是成对使用的。

2.1 ColumnParallelLinear

将权重矩阵按 列(Column) 进行切分,分配给多张 GPU。

假设我们有一个线性运算 Y = W ⋅ X Y = W \cdot X Y=W⋅X。

如果 W W W 非常大,我们把 W W W 沿着列方向切分为多块(假设有 2 张 GPU): W = W 1 , W 2 W = W_1, W_2 W=W1,W2

  • GPU 1 负责计算 Y 1 = W 1 ⋅ X Y_1 = W_1 \cdot X Y1=W1⋅X
  • GPU 2 负责计算 Y 2 = W 2 ⋅ X Y_2 = W_2 \cdot X Y2=W2⋅X
输入与输出特征
  • 输入 X X X:在所有 GPU 上都是完全相同的拷贝(前向传播时,每张卡都会拿到完整输入)。
  • 输出 Y Y Y:每张 GPU 算出的结果是最终输出的"一部分"。GPU 1 拿到了特征向量的左半部分,GPU 2 拿到了右半部分。
通信代价

Tip:

零通信! 在前向传播(Forward)阶段,不需要任何 GPU 间的通信。每张显卡独立计算自己分得的列。

代码中的体现

self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.qk_head_dim)

self.wkv_b = ColumnParallelLinear(self.kv_lora_rank, self.n_heads * (self.qk_nope_head_dim + self.v_head_dim))

在 MLA(多头隐向量注意力)中,生成 Q、K、V 的投影矩阵(如 self.wq, self.wkv_b )用的就是列并行。在注意力机制中,切分列本质上就是把"不同的注意力头(Attention Heads)"分配给了不同的 GPU

2.2 RowParallelLinear

将权重矩阵按 行(Row) 进行切分,分配给多张 GPU。

假设我们要进行下一步计算 Z = Y ⋅ W o u t Z = Y \cdot W_{out} Z=Y⋅Wout,其中 Y Y Y 是刚才列并行的输出(已经切分为 Y 1 Y_1 Y1 和 Y 2 Y_2 Y2 散落在两张卡上)。

Y = Y 1 , Y 2 Y = Y_1, Y_2 Y=Y1,Y2

此时,我们将 W o u t W_{out} Wout 沿着行方向切分:

W o u t = W o u t 1 W o u t 2 W_{out} = \begin{bmatrix} W_{out1} \\ W_{out2} \end{bmatrix} Wout=Wout1Wout2

计算结果则为: Z = Y 1 ⋅ W o u t 1 + Y 2 ⋅ W o u t 2 Z = Y_1 \cdot W_{out1} + Y_2 \cdot W_{out2} Z=Y1⋅Wout1+Y2⋅Wout2。

输入与输出特征
  • 输入 :GPU 1 使用它刚才计算出的 Y 1 Y_1 Y1 乘以 W o u t 1 W_{out1} Wout1;GPU 2 使用 Y 2 Y_2 Y2 乘以 W o u t 2 W_{out2} Wout2。
  • 输出 :每张显卡各自算出来的只是一部分 部分和(Partial Sum)
通信代价

IMPORTANT

需要 All-Reduce! 为了得到最终完整的 Z Z Z,必须在这个算子计算完成后,执行一次跨 GPU 的 All-Reduce 操作,把所有 GPU 上算出来的部分和累加起来。这样所有 GPU 才能拿到完整且一致的 Z Z Z 结果。

代码中的体现

self.wo = RowParallelLinear(self.n_heads * self.v_head_dim, self.dim)

注意力机制最后一步的输出投影矩阵 self.wo 使用的就是行并行。前面各张 GPU 分别算完了属于自己的注意力头的输出,通过按行切分的 wo 相乘后,执行 All-Reduce 将结果累加,输出最终完整的隐状态向量。

2.3 通过一个例子来解释
2.3.1 先看"不并行"时,正常是怎么算的

假设我们现在要做一个极其简单的矩阵乘法(向量乘以矩阵): Z = X ⋅ W Z = X \cdot W Z=X⋅W。

为了方便理解,假设 Y Y Y 是一个有 4 个数字的向量, W W W 是一列有 4 个数字的权重:

  • X = x 1 , x 2 , x 3 , x 4 X = x_1, x_2, x_3, x_4 X=x1,x2,x3,x4
  • W = w 1 w 2 w 3 w 4 W = \begin{bmatrix} w_1 \\ w_2 \\ w_3 \\ w_4 \end{bmatrix} W= w1w2w3w4
    按照正常的矩阵乘法(点乘)规则,最终结果 Z Z Z 是把对应位置相乘,然后全部加起来
    Z = x 1 w 1 + x 2 w 2 + x 3 w 3 + x 4 w 4 Z = x_1 w_1 + x_2 w_2 + x_3 w_3 + x_4 w_4 Z=x1w1+x2w2+x3w3+x4w4
    这就是最终完整、正确的结果
2.3.2 现在切开(RowParallelLinear 登场)

在上一层(列并行)结束时,输入特征 Y Y Y 已经被切成了两半散落在两张显卡上了(这是列并行的特性造成的):

  • GPU 1 手里只拿着 Y Y Y 的左半边: Y l e f t = y 1 , y 2 Y_{left} = y_1, y_2 Yleft=y1,y2
  • GPU 2 手里只拿着 Y Y Y 的右半边: Y r i g h t = y 3 , y 4 Y_{right} = y_3, y_4 Yright=y3,y4
    这时候遇到了我们要解释的 RowParallelLinear(行并行) 。它的规则是把权重 W W W 横向(按行)切断
  • GPU 1 分配到 W W W 的上半部: W t o p = w 1 w 2 W_{top} = \begin{bmatrix} w_1 \\ w_2 \end{bmatrix} Wtop=w1w2
  • GPU 2 分配到 W W W 的下半部: W b o t = w 3 w 4 W_{bot} = \begin{bmatrix} w_3 \\ w_4 \end{bmatrix} Wbot=w3w4

2.3.3 为什么叫"部分和(Partial Sum)"?

现在,两张显卡开始自己算自己的(不互相通信)。

按照矩阵乘法规则,它们各自手里的向量和权重正好能乘起来:

  • GPU 1 自己算的输出(设为 Z 1 Z_1 Z1)
    Z 1 = Y l e f t ⋅ W t o p = y 1 , y 2 w 1 w 2 = y 1 w 1 + y 2 w 2 Z_1 = Y_{left} \cdot W_{top} = y_1, y_2 \cdot \begin{bmatrix} w_1 \\ w_2 \end{bmatrix} = \mathbf{y_1 w_1 + y_2 w_2} Z1=Yleft⋅Wtop=y1,y2w1w2=y1w1+y2w2
  • GPU 2 自己算的输出(设为 Z 2 Z_2 Z2)
    Z 2 = Y r i g h t ⋅ W b o t = y 3 , y 4 w 3 w 4 = y 3 w 3 + y 4 w 4 Z_2 = Y_{right} \cdot W_{bot} = y_3, y_4 \cdot \begin{bmatrix} w_3 \\ w_4 \end{bmatrix} = \mathbf{y_3 w_3 + y_4 w_4} Z2=Yright⋅Wbot=y3,y4w3w4=y3w3+y4w4
    你对比一下正常的最终结果 Z Z Z 和 GPU 各自算出来的 Z 1 , Z 2 Z_1, Z_2 Z1,Z2:
    正常的 Z = ( y 1 w 1 + y 2 w 2 ) + ( y 3 w 3 + y 4 w 4 ) Z = (y_1 w_1 + y_2 w_2) + (y_3 w_3 + y_4 w_4) Z=(y1w1+y2w2)+(y3w3+y4w4)
    所以:Z = Z 1 + Z 2 Z = Z_1 + Z_2 Z=Z1+Z2
    看出来了吗?
    GPU 1 算出来的 Z 1 Z_1 Z1 不是最终结果,它只是最终结果"前半部分的加和"。
    GPU 2 算出来的 Z 2 Z_2 Z2 也不是最终结果,它只是最终结果"后半部分的加和"。
    这就是为什么我们把它们叫做 部分和(Partial Sum)。任何一张单卡拿着自己的 Partial Sum 是没法往下做后续模型计算的,因为数据不完整。

2.3.4 救星登场:All-Reduce 通信

为了让模型继续往下算,两张卡必须把完整的 Z Z Z 拼凑出来。这时候就需要调用底层分布式通信协议(比如 NCCL)里的 All-Reduce 操作。

!NOTE

All-Reduce(全规约) 的动作可以用大白话翻译为:

"兄弟们,把你们手里的结果都拿出来,我们统一加在一起,然后把最终的总和,抄送给每一个人。"

  • 步骤 A (Reduce) :GPU 1 拿出 Z 1 Z_1 Z1,GPU 2 拿出 Z 2 Z_2 Z2,通过显卡间的网线(如 NVLink)传数据,在某处相加得到完整的 Z = Z 1 + Z 2 Z = Z_1 + Z_2 Z=Z1+Z2。
  • 步骤 B (All) :把相加后完整的 Z Z Z,广播发回给 GPU 1 和 GPU 2。
2.3.5 最终结果:

执行完 All-Reduce 之后,GPU 1 手里的内存从毫无意义的 Z 1 Z_1 Z1 变成了完整的 Z Z Z;GPU 2 手里的内存也从 Z 2 Z_2 Z2 变成了完整的 Z Z Z。

两张卡重新拥有了完全相同且正确的数据,又可以愉快地进行下一层网络(如新一轮的列并行)的计算了!

通过这种数学上的"拆解分配 -> 独立相乘 -> 汇聚相加",模型成功绕过了单卡算力/显存的瓶颈,这正是张量并行(Tensor Parallelism)的魔法所在。

2.4 为什么ColumnParallelLinear和RowParallelLinear是"天生一对"?

在 Transformer 的注意力层(Attention)或前馈网络层(FFN/MLP)中,这两个算子总是采用 列并行 → \rightarrow → 行并行 的组合策略。

这么设计的精妙之处在于将分布式训练的通信开销降到了最低。其数据流转过程如下:

步骤 阶段 使用算子 / 操作 是否需要 GPU 通信
1 投影阶段 经过 ColumnParallelLinear 生成 Q/K/V ❌ 零通信
2 计算阶段 执行 Attention 打分或激活函数计算 ❌ 零通信
3 聚合阶段 经过 RowParallelLinear 计算部分和 ❌ 零通信
4 同步阶段 模块出口处执行一次 All-Reduce 汇聚结果 一次通信

如果随意切分矩阵,可能导致每做一步矩阵乘法都需要在 GPU 之间传输大量数据,挤爆网络带宽。列-行组合保证了一个庞大的计算模块(如整个 Attention Layer)内部是完全独立的"零通信"状态,只有在离开该模块时才进行一次高效的数据汇聚。

Multi-head Latent Attention代码

以下代码摘自DeepSeek V3:

传统的 Transformer(如 LLaMA 使用的 MHA/GQA)在推理时需要缓存每个 Token 的完整 K 和 V 矩阵(即 KV Cache),这会占用巨大的显存。MLA 的核心思想是通过"低秩投影"将 K 和 V 压缩到一个共享的低维潜在空间(Latent Space)中,从而将 KV Cache 的显存占用降低 10 倍以上,同时由于独特的推导,它在数学上完全等价于标准注意力,不损失性能。

python 复制代码
class MLA(nn.Module):
    """
    Multi-Headed Attention Layer (MLA).

    Attributes:
        dim (int): Dimensionality of the input features. 模型隐藏维度
        n_heads (int): Number of attention heads.
        n_local_heads (int): Number of local attention heads for distributed systems.
        q_lora_rank (int): Rank for low-rank query projection. 
        kv_lora_rank (int): Rank for low-rank key/value projection. 
        qk_nope_head_dim (int): Dimensionality of non-positional query/key projections.
        qk_rope_head_dim (int): Dimensionality of rotary-positional query/key projections.
        qk_head_dim (int): Total dimensionality of query/key projections.
        v_head_dim (int): Dimensionality of value projections.
        softmax_scale (float): Scaling factor for softmax in attention computation.
    """
    # 核心思想:标准多头注意力每个头存完整的 K、V,总 KV Cache = 2 × n_heads × head_dim。
    # MLA 把 KV 压缩到共享的低秩潜在空间(kv_lora_rank = 512),再按需展开到各头,KV Cache 只存压缩后的向量和位置编码,显存减少 10 倍以上。
    def __init__(self, args: ModelArgs):
        super().__init__()
        self.dim = args.dim
        self.n_heads = args.n_heads
        
        # 张量并行(Tensor Parallelism)设置:n_local_heads代表当前GPU负责计算的注意力头数量
        # world_size:分布式训练的GPU数量,每个 GPU 只负责部分头
        self.n_local_heads = args.n_heads // world_size 

        # Q压缩后的潜在维度
        self.q_lora_rank = args.q_lora_rank
        # KV压缩后的潜在维度
        self.kv_lora_rank = args.kv_lora_rank
        # 无位置编码的QK维度
        self.qk_nope_head_dim = args.qk_nope_head_dim
        # 旋转位置编码的维度
        self.qk_rope_head_dim = args.qk_rope_head_dim

        # qk_head_dim:每个头完整 Q/K 的维度 = 无位置部分 + RoPE 部分。例如 128 + 64 = 192
        # Q 和 K 被拆成两部分处理:不施加位置编码的部分(nope)+ 施加 RoPE 的部分(rope)
        self.qk_head_dim = args.qk_nope_head_dim + args.qk_rope_head_dim
        self.v_head_dim = args.v_head_dim

        
        if self.q_lora_rank == 0:
            # 如果 q_lora_rank 为 0,退化为标准的多头注意力Q投影
            self.wq = ColumnParallelLinear(self.dim, self.n_heads * self.qk_head_dim)
        else:
            # 否则,对 Q 也进行低秩压缩:先压缩到 q_lora_rank,做归一化,再展开到各个头
            # 类似于前面提到的 LoRA 思想:
            # x (5120) ── wq_a ──▶  q_compressed (1536) ── RMSNorm ── wq_b ──▶  Q (24576)
            self.wq_a = Linear(self.dim, self.q_lora_rank)
            self.q_norm = RMSNorm(self.q_lora_rank)
            self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.qk_head_dim)
        
        # wkv_a 负责将输入 x 压缩成:Latent KV (维度 kv_lora_rank) + 用于 RoPE 的 K (维度 qk_rope_head_dim)
        # 注意:这里只有一个矩阵,所有注意力头共享这个压缩后的 Latent 向量!    
        self.wkv_a = Linear(self.dim, self.kv_lora_rank + self.qk_rope_head_dim)
        self.kv_norm = RMSNorm(self.kv_lora_rank)

        # wkv_b 负责将压缩后的 Latent KV 展开(解压)成各个头独享的 K_nope 和 V
        self.wkv_b = ColumnParallelLinear(self.kv_lora_rank, self.n_heads * (self.qk_nope_head_dim + self.v_head_dim))
        
        # 最终的输出投影矩阵
        self.wo = RowParallelLinear(self.n_heads * self.v_head_dim, self.dim)

        # Attention 里的 softmax 缩放因子 (1 / sqrt(d))
        self.softmax_scale = self.qk_head_dim ** -0.5
        
        # 针对长文本推断(Context Extension)的缩放调整(类似 Yarn 算法的乘子)
        if args.max_seq_len > args.original_seq_len:
            mscale = 0.1 * args.mscale * math.log(args.rope_factor) + 1.0
            self.softmax_scale = self.softmax_scale * mscale * mscale

        # kv cache显存分配
        if attn_impl == "naive":
            # 朴素实现:开辟标准的 K 和 V Cache。维度极大 (seq_len, heads, head_dim)
            self.register_buffer("k_cache", 
                torch.zeros(args.max_batch_size, args.max_seq_len, self.n_local_heads, self.qk_head_dim), persistent=False)
            self.register_buffer("v_cache", 
                torch.zeros(args.max_batch_size, args.max_seq_len, self.n_local_heads, self.v_head_dim), persistent=False)
        else:
            # 优化实现(实际部署用):KV Cache只需要存压缩后的 kv_lora_rank 向量,以及一小段共享的pe(位置编码)向量。
            # 这就是 MLA 能省十几倍显存的根本原因!
            self.register_buffer("kv_cache", 
                torch.zeros(args.max_batch_size, args.max_seq_len, self.kv_lora_rank), persistent=False)
            self.register_buffer("pe_cache", 
                torch.zeros(args.max_batch_size, args.max_seq_len, self.qk_rope_head_dim), persistent=False)

    def forward(self, x: torch.Tensor, start_pos: int, freqs_cis: torch.Tensor, mask: Optional[torch.Tensor]):
        """
        Forward pass for the Multi-Headed Attention Layer (MLA).

        Args:
            x (torch.Tensor): Input tensor of shape (batch_size, seq_len, dim).
            start_pos (int): Starting position in the sequence for caching.
            freqs_cis (torch.Tensor): Precomputed complex exponential values for rotary embeddings.
            mask (Optional[torch.Tensor]): Mask tensor to exclude certain positions from attention.

        Returns:
            torch.Tensor: Output tensor with the same shape as the input.
        """
        bsz, seqlen, _ = x.size() # (batch_size, seq_length, embedding_size)
        end_pos = start_pos + seqlen

        # ================= 1. 计算 Query =================
        if self.q_lora_rank == 0:  # 退化为标准的多头注意力
            q = self.wq(x)
        else:  # 低秩压缩
            q = self.wq_b(self.q_norm(self.wq_a(x)))

        # q的shape (batch_size, seq_length, embed_size)最后一个维度embed_size被拆分成了 (num_heads, head_dim)
        # 变换后的shape是(N, seq_length, num_heads, head_dim)  这一步与前面的MHA的forward中的shape转换是一致的
        q = q.view(bsz, seqlen, self.n_local_heads, self.qk_head_dim)
        
        # 将 Q 拆分为两半:不带位置编码的部分 (q_nope) 和 需要应用位置编码的部分 (q_pe)
        q_nope, q_pe = torch.split(q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
        # 仅对 q_pe 施加旋转位置编码 (RoPE)
        q_pe = apply_rotary_emb(q_pe, freqs_cis)

        # ================= 2. 计算并压缩 KV =================
        kv = self.wkv_a(x)    # [batch, seq, kv_lora_rank + qk_rope_head_dim]
        # 拆分为:Latent KV (内容向量) 和 k_pe (需要位置编码的 K)
        kv, k_pe = torch.split(kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
        # 仅对 k_pe 施加旋转位置编码 (注意这里 k_pe 在各个头上是共享的,所以 unsqueeze 扩展了一维)
        k_pe = apply_rotary_emb(k_pe.unsqueeze(2), freqs_cis)
        
        if attn_impl == "naive":
            q = torch.cat([q_nope, q_pe], dim=-1)
            kv = self.wkv_b(self.kv_norm(kv))
            kv = kv.view(bsz, seqlen, self.n_local_heads, self.qk_nope_head_dim + self.v_head_dim)
            k_nope, v = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)
            k = torch.cat([k_nope, k_pe.expand(-1, -1, self.n_local_heads, -1)], dim=-1)
            self.k_cache[:bsz, start_pos:end_pos] = k
            self.v_cache[:bsz, start_pos:end_pos] = v
            scores = torch.einsum("bshd,bthd->bsht", q, self.k_cache[:bsz, :end_pos]) * self.softmax_scale
        else:
            wkv_b = self.wkv_b.weight if self.wkv_b.scale is None 
                                          else weight_dequant(self.wkv_b.weight, self.wkv_b.scale, block_size) 
            wkv_b = wkv_b.view(self.n_local_heads, -1, self.kv_lora_rank)
            q_nope = torch.einsum("bshd,hdc->bshc", q_nope, wkv_b[:, :self.qk_nope_head_dim])
            self.kv_cache[:bsz, start_pos:end_pos] = self.kv_norm(kv)
            self.pe_cache[:bsz, start_pos:end_pos] = k_pe.squeeze(2)
            scores = (torch.einsum("bshc,btc->bsht", q_nope, self.kv_cache[:bsz, :end_pos]) +
                      torch.einsum("bshr,btr->bsht", q_pe, self.pe_cache[:bsz, :end_pos])) * self.softmax_scale
        if mask is not None:
            scores += mask.unsqueeze(1)
        scores = scores.softmax(dim=-1, dtype=torch.float32).type_as(x)
        if attn_impl == "naive":
            x = torch.einsum("bsht,bthd->bshd", scores, self.v_cache[:bsz, :end_pos])
        else:
            x = torch.einsum("bsht,btc->bshc", scores, self.kv_cache[:bsz, :end_pos])
            x = torch.einsum("bshc,hdc->bshd", x, wkv_b[:, -self.v_head_dim:])
        x = self.wo(x.flatten(2))
        return x
相关推荐
海天一色y2 小时前
强化学习工具函数详解:从经验回放到优势函数计算
人工智能·python·强化学习
yuhulkjv3352 小时前
告别复制粘贴式降级:纳米AI鸿蒙版导出word格式为何绕不开“AI 导出鸭”
人工智能·ai·word·harmonyos·ai导出鸭
故七月2 小时前
生成式引擎优化(GEO)的底层逻辑与产业实践
大数据·人工智能·机器学习
MacroZheng2 小时前
又一个神级画图Skill开源,再见draw.io!
java·人工智能·后端
船厂电气自动化ai大模型2 小时前
AI大模型与数学·第56课 快速傅里叶变换FFT:DFT高效优化算法,图像、音频、扩散模型工程加速核心工具
数据结构·人工智能·深度学习·算法·机器学习
TWT1212 小时前
用 TRAE Work 5 分钟搞定技术周报,再也不用周五下午憋字了
人工智能
AI导出鸭2 小时前
腾讯ima的LaTeX生成PDF文件复制后数学公式乱码,怎样修改?AI导出鸭苹果版硬核拆解
人工智能·pdf·ai导出鸭
TechEdu2026062 小时前
[人工智能]Kimi(月之暗面 Moonshot AI):长上下文、智能体与工程实践
人工智能·ai
9i编程2 小时前
借助 Trae Work 学透 Multi-Agent 代码:从「抄出来了但没懂」到完整调通 v1.0.8
人工智能·openai·ai编程
jay神2 小时前
YOLO还能不能作为模型baseline?
人工智能·深度学习·yolo·cnn·毕业设计