PyTorch中Linear操作Reshape后结果差异过大的原因排查
问题分析:Transformer QKV拆分后矩阵不匹配的原因
问题描述
尝试从Transformer大语言模型的拼接QKV矩阵中提取用于KV缓存的Q矩阵时,拆分注意力头前manual_q与unsplit_q完全匹配,但拆分头后两者差异极大,触发断言错误。代码如下:
def get_q_matrix(self, x): batch_size, seq_len, n_embd = x.size() debug_qkv = self.query_key_value(x) # shape (batch_size, seq_len, n_embd) unsplit_q, _, _ = debug_qkv.split( self.n_embd, dim=-1 ) # shape (batch_size, seq_len, n_embd // 3) debug_qkv = debug_qkv.view( batch_size, seq_len, self.n_head, 3 * self.head_size ) # shape (batch_size, seq_len, 4, 96) q, _, _ = debug_qkv.split( self.head_size, dim=-1 ) # shape (batch_size, seq_len, 4, 96 // 3) # Ensure correct weight and bias extraction weight = self.query_key_value.weight bias = self.query_key_value.bias q_weight, k_weight, v_weight = weight.chunk(3, dim=0) q_bias, k_bias, v_bias = bias.chunk(3, dim=0) manual_q = F.linear(x, q_weight, q_bias) manual_q = manual_q # shape (batch_size, seq_len, n_embd) assert torch.allclose(unsplit_q, manual_q) # Passes print(torch.max(torch.abs(unsplit_q - manual_q))) # tensor(0.) manual_q = manual_q.view(batch_size, seq_len, self.n_head, self.head_size) print(torch.max(torch.abs(q - manual_q))) # tensor(35.6218) assert torch.allclose(q, manual_q) # AssertionError return manual_q
核心原因:QKV维度拆分顺序完全错误
问题出在对拼接后的QKV矩阵的维度拆分逻辑上:
拼接QKV矩阵的结构:
query_key_value层的输出维度是3 * hidden_size(其中hidden_size = n_head * head_size),其最后一维的排布是:[Q_head1, Q_head2, ..., Q_headN, K_head1, K_head2, ..., K_headN, V_head1, ..., V_headN]也就是先完整存放所有注意力头的Q向量,再存放所有头的K向量,最后是所有头的V向量。
代码中的错误操作:
你先将整个QKV矩阵view为(batch_size, seq_len, n_head, 3*head_size),再按head_size拆分Q/K/V。这相当于把3*hidden_size的维度拆成了n_head个组,每个组包含同一个头的Q、K、V片段,完全违背了原始QKV的排布逻辑。这种拆分方式会把Q的第1头、K的第1头、V的第1头的部分数据混在一起,自然和手动计算的Q矩阵(先提取完整Q再拆头)完全不匹配。
修正方案
调整拆分顺序:先把QKV矩阵拆分为完整的Q、K、V三个部分,再对每个部分拆分注意力头:
def get_q_matrix(self, x): batch_size, seq_len, n_embd = x.size() # 定义单个Q/K/V的维度 hidden_size = self.n_head * self.head_size debug_qkv = self.query_key_value(x) # shape (batch_size, seq_len, 3*hidden_size) # 第一步:先拆分出完整的Q、K、V(每个维度为hidden_size) unsplit_q, unsplit_k, unsplit_v = debug_qkv.split(hidden_size, dim=-1) # 第二步:对Q拆分注意力头 q = unsplit_q.view(batch_size, seq_len, self.n_head, self.head_size) # 手动计算Q矩阵并拆头 weight = self.query_key_value.weight bias = self.query_key_value.bias q_weight, k_weight, v_weight = weight.chunk(3, dim=0) q_bias, k_bias, v_bias = bias.chunk(3, dim=0) manual_q = F.linear(x, q_weight, q_bias) manual_q = manual_q.view(batch_size, seq_len, self.n_head, self.head_size) # 现在断言会通过 assert torch.allclose(q, manual_q) print(torch.max(torch.abs(q - manual_q))) # tensor(0.) return manual_q
内容的提问来源于stack exchange,提问作者genjong
相关产品推荐
相关产品推荐

