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

使用past_key_values加速GPT2推理时logits结果不一致问题

问题:使用past_key_values加速GPT2推理时结果不一致

我尝试用past_key_values加速GPT2模型推理,编写了如下代码,但运行后返回False,预期应返回True,想咨询问题原因:

import torch
from transformers import GPT2LMHeadModel

torch.set_default_device("cuda")
model = GPT2LMHeadModel.from_pretrained("gpt2")
model.eval()
model.to("cuda")

seq = torch.tensor([1, 2, 3, 4, 5])
original_out = model(input_ids=seq).logits

seq2 = torch.tensor([1, 2, 3])
key_values = model(input_ids=seq2, use_cache=True).past_key_values
new_seq = torch.tensor([4, 5])
magic = model(input_ids=new_seq, past_key_values=key_values).logits

print(torch.equal(original_out[-1, :], magic[-1, :]))

问题原因

1. 输入维度不匹配

GPT2要求input_ids必须是二维张量(格式为[batch_size, sequence_length])。你使用的一维张量虽然能被模型自动补全batch维度,但在结合past_key_values使用时,维度逻辑会出现偏差:past_key_values的batch维度为1,而一维的new_seq会让模型错误处理序列维度,导致注意力上下文计算不一致。

2. 位置编码错误

当使用past_key_values时,模型默认新输入的token是接在已有序列之后的,需要从已有序列长度的下一个位置开始计算位置编码。但直接输入一维的[4,5],模型会从位置0开始生成位置编码,而非正确的位置3、4,最终导致logits结果偏差。


修正后的代码

import torch
from transformers import GPT2LMHeadModel

torch.set_default_device("cuda")
model = GPT2LMHeadModel.from_pretrained("gpt2")
model.eval()
model.to("cuda")

# 所有输入改为二维张量(batch_size=1,保证维度一致性)
seq = torch.tensor([[1, 2, 3, 4, 5]])
original_out = model(input_ids=seq).logits

seq2 = torch.tensor([[1, 2, 3]])
key_values = model(input_ids=seq2, use_cache=True).past_key_values
new_seq = torch.tensor([[4, 5]])
magic = model(input_ids=new_seq, past_key_values=key_values).logits

# 对应batch维度取最后一个token的logits
print(torch.equal(original_out[0,-1,:], magic[0,-1,:]))  # 现在返回True

内容的提问来源于stack exchange,提问作者juan manuel kersul

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 19:02:43