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

如何解决Florence2模型Torch JIT推理无法利用past_key_values的性能问题?

问题解决思路与替代方案

一、修复JIT追踪后无法利用past_key_values的性能问题

torch.jit.trace导致性能暴跌的核心原因是:它基于固定输入形状生成计算图,而past_key_values(KV缓存)的形状会随生成token数量动态变化(每多生成一个token,缓存就追加一组新的KV张量)。trace会把缓存形状硬编码到计算图中,导致每次生成新token时无法复用已有缓存,只能重新计算整个序列的所有token,这才出现了慢10倍的情况。

解决这个问题的关键是让JIT支持动态形状与缓存复用逻辑,首选方案如下:

1. 替换为torch.jit.script编译

script是基于代码逻辑而非具体输入进行追踪,能处理动态形状、条件分支等trace无法覆盖的逻辑,完美适配带KV缓存的生成式模型。具体操作:

  • 将模型封装为自定义类,在前向方法中明确处理past_key_values的传入、更新和返回逻辑,确保逻辑可被script解析。
  • 示例代码框架:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

class ScriptableFlorenceModel(torch.nn.Module):
    def __init__(self, model):
        super().__init__()
        self.model = model

    def forward(self, input_ids, attention_mask, past_key_values=None):
        outputs = self.model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            past_key_values=past_key_values,
            use_cache=True
        )
        return outputs.logits, outputs.past_key_values

# 加载原模型
tokenizer = AutoTokenizer.from_pretrained("microsoft/Florence-2-large")
model = AutoModelForCausalLM.from_pretrained("microsoft/Florence-2-large")

# 编译为script模型
script_model = torch.jit.script(ScriptableFlorenceModel(model))
# 保存模型
script_model.save("florence2_scripted.pt")
  • 推理时,每次生成新token都传入上一步返回的past_key_values,即可正常复用缓存,恢复原模型的推理速度。

2. (不推荐)手动适配trace的动态形状

如果必须使用trace,需要强制让计算图支持动态维度:

  • 追踪前将past_key_values初始化为动态形状的张量(比如用torch.randn创建带动态维度的张量,或用torch.Tensor.size()动态获取维度)。但这种方法容易出现形状不匹配的报错,维护成本高,远不如script可靠。

二、LLM场景下无法利用KV缓存时的推理方案

如果受限于部署环境或模型特性,确实无法使用KV缓存,只能做全序列推理,可以通过以下手段尽可能提升性能:

  • 批量对齐短序列:将多个生成请求的输入序列长度对齐(给短序列补padding),利用GPU的批量并行计算能力抵消单序列全量计算的开销,降低单请求的平均耗时。
  • 限制生成长度:全序列推理的耗时与序列长度正相关,根据业务需求严格设置max_new_tokens,避免生成过长的序列。
  • 模型量化:用INT8/INT4量化压缩模型参数,减少计算量和内存占用。比如用torch.ao.quantization做后训练量化(PTQ),或用bitsandbytes做量化,量化后的模型即使全序列推理,速度也会有明显提升。
  • 预计算共享前缀:如果多个请求有相同的输入前缀(比如固定prompt),可以提前预计算该前缀的所有中间结果,将这些结果作为初始输入分发给各个请求,避免重复计算前缀部分的开销。

内容的提问来源于stack exchange,提问作者Vivek Kalyanarangan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 17:25:06