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

如何在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. 代码1中缓存未生效的原因:第二次调用generate时传入了完整的outputs.sequences(包含之前的所有token),模型会默认重新处理所有输入token,忽略传入的past_key_values。只有当输入是新的增量token时,模型才会复用缓存。

  2. 代码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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 08:29:51