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

PyTorch中past_key_values缓存用法及拼接结果不一致问题

关于GPT2中past_key_values缓存的两个问题解答

1. 如何用past_key_values实现缓存功能?

past_key_values本质是Transformer模型注意力层的历史key和value张量缓存,核心作用是避免重复计算已处理序列的注意力,大幅提升长文本生成的效率。在Hugging Face Transformers库中的使用步骤非常清晰:

  • 第一步:处理前缀序列时,开启use_cache=True参数,模型会在返回结果中附带past_key_values,里面存储了每一层注意力模块的历史key、value张量。
  • 第二步:后续生成新token时,无需输入完整历史序列,只传入当前待处理的新token,同时将past_key_values传入模型。模型会自动复用缓存的历史数据,仅计算新token的注意力和输出。

对应你代码里的核心逻辑:

# 处理前缀序列,获取缓存
uncomplete_ids = ids[:, :-1]
output = model(input_ids=uncomplete_ids, use_cache=True)
past_key_values = output.past_key_values

# 传入新token+缓存,生成下一个词
last_id = ids[:, -1:]
output = model(input_ids=last_id, past_key_values=past_key_values)

这种方式在长序列生成时能显著减少计算量,避免重复计算历史序列的注意力权重。

2. 为什么两次生成结果不一致?

问题出在你直接输入整句时的logits处理逻辑错误,和past_key_values本身无关:

当你直接输入完整序列"Hello, my dog is cute"时,模型输出的logits形状是(1, 5, vocab_size)(5是序列的token数量)。你直接对squeeze(0)后的(5, vocab_size)执行torch.multinomial,得到的是序列中每个token对应的下一个词采样结果,而你取的next_word_index.tolist()[0]是第一个token"Hello"的下一个词,并非最后一个token"cute"的下一个词。

而用past_key_values的方式,你输入的是最后一个token"cute",模型输出的logits是(1,1,vocab_size),采样的是"cute"对应的下一个词——两者采样的是完全不同位置的预测结果,自然不一致。

修正直接输入的代码,只取最后一个位置的logits即可对齐结果:

output = model(input_ids=ids)
logits = output.logits[:, -1, :]  # 仅取最后一个token的logits
probabilities = F.softmax(logits, dim=-1)
next_word_index = torch.multinomial(probabilities, 1)
next_word = tokenizer.decode(next_word_index.tolist()[0])

修正后重新运行,两种方式的采样结果就会一致。

内容的提问来源于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.09 05:27:38