You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Transformer多头注意力为何共享KQV权重矩阵?

关于多头注意力共享KQV投影权重的疑问解答

代码示例

self.values = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.keys = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.queries = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.fc_out = nn.Linear(heads * self.head_dim, embed_size)

def forward(self, values, keys, query, mask):
    # Get number of training examples
    N = query.shape[0]

    value_len, key_len, query_len = values.shape[1], keys.shape[1], query.shape[1]

    # Split the embedding into self.heads different pieces
    values = values.reshape(N, value_len, self.heads, self.head_dim)
    keys = keys.reshape(N, key_len, self.heads, self.head_dim)
    query = query.reshape(N, query_len, self.heads, self.head_dim)

    values = self.values(values)  # (N, value_len, heads, head_dim)
    keys = self.keys(keys)  # (N, key_len, heads, head_dim)
    queries = self.queries(query)  # (N, query_len, heads, heads_dim)

    # Einsum does matrix mult. for query*keys for each training example
    # with every other training example, don't be confused by einsum
    # it's just how I like doing matrix multiplication & bmm

    energy = torch.einsum("nqhd,nkhd->nhqk", [queries, keys])

问题解答

首先明确:这种共享KQV投影权重的实现是标准多头注意力的简化版本,和Transformer原论文里的实现不同——原生版本是每个注意力头都有独立的K、Q、V投影权重。

关于你的疑问,拆解来看:

  • 反向传播信号不是核心原因:共享权重时,所有头的投影操作完全一致,反向传播时每个头的梯度会叠加到同一个权重矩阵上,相当于所有头协同更新同一个投影规则,但这是共享权重带来的结果,而非设计的初衷。
  • 真正的设计原因是减参:原生多头注意力的参数量是3 * heads * embed_size * head_dim,而共享权重的版本参数量仅为3 * head_dim * head_dim,参数量大幅降低,适合计算资源有限的场景。
  • 这种实现的局限性:代价是丢失了原生多头注意力的核心优势——每个头无法学到差异化的语义注意力模式,因为所有头的投影完全相同,后续的注意力计算本质是在重复做类似的任务,只能从注意力分数层面做区分,没法从投影层就捕捉不同的特征。

内容的提问来源于stack exchange,提问作者ZCao

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 15:48:47