如何在Hugging Face中复用Llama3.2 3B模型的KV缓存?
复用Llama3.2的KV缓存与Embedding以加速长文本多问答处理
我需要用Transformer处理一段固定长文本,之后让用户针对该文本提出不同的独立问题,输入格式为「固定长文本 + 动态用户问题」。为了避免每次重新计算长文本的注意力,我想知道如何存储并复用该长文本的Keys和Values(额外希望能复用Embeddings)。目标模型是Llama3.2 3B,使用Python实现,任意框架均可。
我的实验代码与问题
代码1:尝试复用缓存但未生效
import torch import time from transformers import AutoTokenizer, AutoModelForCausalLM # 加载模型 tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-1B") model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-1B") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device) tokenized = tokenizer(long_text, return_tensors="pt") with torch.no_grad(): n = 5000 t1 = time.time() # 步骤1:计算5000个token的KV缓存 outputs = model.generate( input_ids=tokenized['input_ids'].to(device)[:, :n], attention_mask=torch.ones((1, n), dtype=torch.int).to(device), temperature=0.7, max_new_tokens=1, use_cache=True, return_dict_in_generate=True ) t2 = time.time() past_kv = outputs.past_key_values # 步骤2:尝试复用之前的计算结果 outputs = model.generate( input_ids=outputs.sequences, attention_mask=torch.ones((1, outputs.sequences.shape[1]), dtype=torch.int).to(device), past_key_values=past_kv, # 这里的缓存未被使用,为什么? temperature=0.7, max_new_tokens=1, use_cache=True, return_dict_in_generate=True ) t3 = time.time() print(t2-t1) # 约3秒 print(t3-t2) # 同样约3秒,应该更快才对
问题:第二次调用generate时,缓存没有生效,耗时和第一次几乎一样。
代码2:仅传入最后一个token加缓存抛出异常
# 步骤2 new_token_id = outputs.sequences[:, -1:].to(device) # 只使用最后一个token past_kv = outputs.past_key_values outputs2 = model.generate( # 抛出异常 input_ids=new_token_id, past_key_values=past_kv, max_new_tokens=1, use_cache=True, return_dict_in_generate=True )
异常信息:
IndexError: index -1 is out of bounds for dimension 0 with size 0
问题原因分析
代码1中缓存未生效的原因:第二次调用
generate时传入了完整的outputs.sequences(包含之前的所有token),模型会默认重新处理所有输入token,忽略传入的past_key_values。只有当输入是新的增量token时,模型才会复用缓存。代码2报错的原因:
- 缺少
attention_mask参数:当使用past_key_values时,需要传入对应于当前输入token的注意力掩码,同时要确保掩码长度与缓存的上下文长度匹配。 - 输入token与缓存的维度不匹配:Llama的KV缓存包含每一层的K和V张量,其形状与上下文长度绑定,若输入的注意力掩码未正确覆盖上下文+新token的长度,会导致索引越界。
- 缺少
正确解决方案
步骤1:预计算长文本的KV缓存与Embedding
不要用generate,直接调用模型的forward方法处理长文本,得到完整的KV缓存和Embedding:
import torch import time from transformers import AutoTokenizer, AutoModelForCausalLM # 加载模型与tokenizer tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-1B") model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-1B") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device) # 预处理固定长文本 long_text = "你的固定长文本内容..." tokenized_long = tokenizer(long_text, return_tensors="pt").to(device) context_len = tokenized_long["input_ids"].shape[1] # 预计算长文本的KV缓存和Embedding with torch.no_grad(): outputs = model( **tokenized_long, use_cache=True, return_dict=True ) past_key_values = outputs.past_key_values # 长文本的KV缓存 # 如果需要复用Embedding,可以保存outputs.last_hidden_state
步骤2:复用缓存处理用户问题
对于每个用户问题,将其token化后作为增量输入,传入预计算的past_key_values,同时正确设置attention_mask:
def process_user_question(question, past_kv, context_len): # 处理用户问题 tokenized_question = tokenizer(question, return_tensors="pt").to(device) question_len = tokenized_question["input_ids"].shape[1] # 构建注意力掩码:上下文部分全1,问题部分全1 attention_mask = torch.cat([ torch.ones((1, context_len), dtype=torch.int).to(device), tokenized_question["attention_mask"] ], dim=1) # 生成回答,复用缓存 with torch.no_grad(): generate_outputs = model.generate( input_ids=tokenized_question["input_ids"], attention_mask=attention_mask, past_key_values=past_kv, max_new_tokens=50, temperature=0.7, use_cache=True, return_dict_in_generate=True ) # 解析回答 answer = tokenizer.decode(generate_outputs.sequences[0], skip_special_tokens=True) return answer, generate_outputs.past_key_values # 可选:返回更新后的缓存(如果后续要继续生成) # 示例:处理第一个用户问题 question1 = "针对固定长文本的第一个问题?" answer1, updated_past_kv = process_user_question(question1, past_key_values, context_len) print(answer1) # 处理第二个用户问题(注意:如果问题是独立的,应该重新使用原始的past_key_values,而非updated_past_kv) question2 = "针对固定长文本的第二个问题?" answer2, _ = process_user_question(question2, past_key_values, context_len) print(answer2)
关键注意事项
- 独立问题需复用原始缓存:如果用户的问题是独立的(而非连续对话),每次都要使用预计算的
past_key_values,而不是上一次生成后的更新缓存,避免上下文混淆。 - 注意力掩码的正确构建:掩码长度必须等于「上下文长度 + 当前输入token长度」,否则会导致维度不匹配错误。
- 禁用自动缓存重置:确保
use_cache=True始终开启,模型会自动在增量输入时复用缓存并更新。
内容的提问来源于stack exchange,提问作者Matt
相关产品推荐
相关产品推荐

