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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 23:07:04