Llama模型自定义计算past_key_values与模型输出不一致问题排查
问题:手动计算的Llama past_key_values与官方输出不匹配
我尝试通过attention_layers和hidden_state提取特定层的past key、value对,实现代码如下:
import torch import torch.nn.functional as F from transformers import LlamaConfig from transformers import LlamaModel, LlamaTokenizer, LlamaForCausalLM tokenizer = LlamaTokenizer.from_pretrained(path_to_llama2) # Load the configuration and enable required outputs config = LlamaConfig.from_pretrained(path_to_llama2) config.output_hidden_states = True config.output_attentions = True # To get self_attn_weights and biases if needed config.use_cache = True # To get past_key_values model = LlamaForCausalLM.from_pretrained(path_to_llama2, config=config) model.eval() input_text = "Once upon a time" inputs = tokenizer(input_text, return_tensors='pt') outputs = model(**inputs) hidden_states = outputs.hidden_states # List of hidden states from each layer state_dict = model.state_dict() # Function to compute past_key_values for a single layer def compute_past_key_values_for_layer(layer_idx, hidden_state): attention_layers = [layer.self_attn for layer in model.model.layers] W_q = state_dict[f'model.layers.{layer_idx}.self_attn.q_proj.weight'] W_k = state_dict[f'model.layers.{layer_idx}.self_attn.k_proj.weight'] W_v = state_dict[f'model.layers.{layer_idx}.self_attn.v_proj.weight'] queries = torch.matmul(hidden_state, W_q.T) keys = torch.matmul(hidden_state, W_k.T) values = torch.matmul(hidden_state, W_v.T) batch_size, seq_length, hidden_dim = hidden_state.size() num_attention_heads = attention_layers[layer_idx].num_heads head_dim = hidden_dim // num_attention_heads keys = keys.view(batch_size, seq_length, num_attention_heads, head_dim) keys = keys.permute(0, 2, 1, 3) values = values.view(batch_size, seq_length, num_attention_heads, head_dim) values = values.permute(0, 2, 1, 3) return keys, values past_key_values = [] for i, hidden_state in enumerate(hidden_states[1:]): # Skip the embedding layer keys, values = compute_past_key_values_for_layer(i, hidden_state) past_key_values.append((keys, values)) past_key_values = tuple(past_key_values)
但通过该方式得到的past_key_values与outputs.past_key_values中对应层的数值不匹配,请问出现该问题的原因是什么?有哪些解决建议?
问题原因分析
- 忽略了投影层的偏置参数:Llama的q/k/v投影层(
q_proj/k_proj/v_proj)不仅包含权重矩阵,还有对应的偏置参数。你只加载了权重,没有将偏置加入计算,导致输出结果与官方实现存在偏差。 - hidden_state的输入对应关系错误:官方的
past_key_values是用当前注意力层的输入hidden state(即上一层的输出)计算的,但你的代码中使用了hidden_states[1:]——这是当前层输出的hidden state,输入输出的顺序不匹配,导致计算的k/v基于错误的特征。 - 未应用旋转位置编码(RoPE):Llama的注意力机制会对queries和keys应用RoPE编码,这是实现相对位置编码的核心步骤。你直接计算完k/v就返回,跳过了这一步,这是数值不匹配的核心原因。
- 张量维度处理的细节差异:官方实现中对k/v张量进行维度变换后会调用
.contiguous()确保内存连续,虽然不改变数值,但可能导致张量的存储顺序与你的手动实现不同,在某些场景下会被判定为不匹配。
解决建议
包含投影层偏置参数:在加载权重时同时获取偏置,并在计算时加入:
def compute_past_key_values_for_layer(layer_idx, hidden_state): attention_layers = [layer.self_attn for layer in model.model.layers] # 同时加载权重和偏置 W_q = state_dict[f'model.layers.{layer_idx}.self_attn.q_proj.weight'] b_q = state_dict[f'model.layers.{layer_idx}.self_attn.q_proj.bias'] W_k = state_dict[f'model.layers.{layer_idx}.self_attn.k_proj.weight'] b_k = state_dict[f'model.layers.{layer_idx}.self_attn.k_proj.bias'] W_v = state_dict[f'model.layers.{layer_idx}.self_attn.v_proj.weight'] b_v = state_dict[f'model.layers.{layer_idx}.self_attn.v_proj.bias'] # 计算时加入偏置 queries = torch.matmul(hidden_state, W_q.T) + b_q keys = torch.matmul(hidden_state, W_k.T) + b_k values = torch.matmul(hidden_state, W_v.T) + b_v # 后续维度处理保持不变 batch_size, seq_length, hidden_dim = hidden_state.size() num_attention_heads = attention_layers[layer_idx].num_heads head_dim = hidden_dim // num_attention_heads keys = keys.view(batch_size, seq_length, num_attention_heads, head_dim) keys = keys.permute(0, 2, 1, 3).contiguous() values = values.view(batch_size, seq_length, num_attention_heads, head_dim) values = values.permute(0, 2, 1, 3).contiguous() return keys, values修正hidden_state的输入来源:将循环中的
hidden_states[1:]改为hidden_states[:-1],因为hidden_states[0]是embedding层输出,作为第0层注意力的输入;hidden_states[1]是第0层输出,作为第1层注意力的输入,以此类推:past_key_values = [] for i, hidden_state in enumerate(hidden_states[:-1]): # 取每一层的输入hidden state keys, values = compute_past_key_values_for_layer(i, hidden_state) past_key_values.append((keys, values)) past_key_values = tuple(past_key_values)应用RoPE编码:复用模型内置的RoPE实现,对queries和keys进行编码:
def compute_past_key_values_for_layer(layer_idx, hidden_state): attention_layers = [layer.self_attn for layer in model.model.layers] # 加载权重和偏置(同上) W_q = state_dict[f'model.layers.{layer_idx}.self_attn.q_proj.weight'] b_q = state_dict[f'model.layers.{layer_idx}.self_attn.q_proj.bias'] W_k = state_dict[f'model.layers.{layer_idx}.self_attn.k_proj.weight'] b_k = state_dict[f'model.layers.{layer_idx}.self_attn.k_proj.bias'] W_v = state_dict[f'model.layers.{layer_idx}.self_attn.v_proj.weight'] b_v = state_dict[f'model.layers.{layer_idx}.self_attn.v_proj.bias'] queries = torch.matmul(hidden_state, W_q.T) + b_q keys = torch.matmul(hidden_state, W_k.T) + b_k values = torch.matmul(hidden_state, W_v.T) + b_v batch_size, seq_length, hidden_dim = hidden_state.size() num_attention_heads = attention_layers[layer_idx].num_heads head_dim = hidden_dim // num_attention_heads # 拆分多头 queries = queries.view(batch_size, seq_length, num_attention_heads, head_dim) keys = keys.view(batch_size, seq_length, num_attention_heads, head_dim) values = values.view(batch_size, seq_length, num_attention_heads, head_dim) # 应用RoPE编码 rotary_emb = model.model.layers[layer_idx].self_attn.rotary_emb position_ids = torch.arange(seq_length, dtype=torch.long, device=hidden_state.device).unsqueeze(0) cos, sin = rotary_emb(queries, position_ids) queries = rotary_emb.apply_rotary_pos_emb(queries, cos, sin) keys = rotary_emb.apply_rotary_pos_emb(keys, cos, sin) # 调整维度并确保连续 keys = keys.permute(0, 2, 1, 3).contiguous() values = values.permute(0, 2, 1, 3).contiguous() return keys, values直接复用模型内部逻辑(推荐):避免手动实现的误差,直接调用对应层的注意力前向过程获取past_key_values:
past_key_values = [] attention_mask = inputs.get('attention_mask', None) current_hidden_state = hidden_states[0] # 从embedding层输出开始 for layer_idx in range(len(model.model.layers)): layer = model.model.layers[layer_idx] # 调用注意力层前向过程,获取past_key_value attn_outputs = layer.self_attn( current_hidden_state, attention_mask=attention_mask, use_cache=True ) past_key_values.append(attn_outputs.past_key_value) # 更新current_hidden_state为当前层输出(用于下一层输入) current_hidden_state = layer(current_hidden_state, attention_mask=attention_mask)[0] past_key_values = tuple(past_key_values)
内容的提问来源于stack exchange,提问作者lazytux
相关产品推荐
相关产品推荐

