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
相关产品推荐
相关产品推荐

