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
相关产品推荐
相关产品推荐

