Allow static cache to be larger than sequence length / batch size for encoder-decoder models

这个报错通常发生在手动创建 StaticCache / EncoderDecoderCache 并在 generate() 外部复用、但 cache 的 batch size 或 encoder 序列长度与当前输入不一致的场景。优先排查是否把 cache 初始化成了“精确尺寸”,而不是“大于等于实际

快速结论:这个报错通常发生在手动创建 StaticCache/EncoderDecoderCache 并在 generate() 外部复用、但 cache 的 batch size 或 encoder 序列长度与当前输入不一致的场景。优先排查是否把 cache 初始化成了“精确尺寸”,而不是“大于等于实际使用的最大尺寸”。

适用环境:Transformers 库;T5 等 encoder-decoder 模型(示例使用 google-t5/t5-small);torch.float16;手动构造 StaticCacheEncoderDecoderCache。Issue 中未确认具体操作系统、Python、CUDA、PyTorch 版本和显卡型号。

最快修复方案:如果手动初始化 cache 而不通过 generate(cache_implementation=...),需要自己把 batch size 设为当前输入的 batch size、把 max cache length 设为与 encoder 序列长度完全匹配。Issue 确认的变 batch size 复用支持已在 #37394 加入;变 encoder 长度的那半部分仍未完成,暂无确认的一步修复方案。

注意事项:generate(cache_implementation=...)generate() 内部会处理 cache 尺寸,问题多出现在手写 cache 的场景。直接通过 past_key_value.key_cache[self.layer_idx] 访问 cache 不是稳定公开用法,T5/Whisper 及其派生模型会为 cross-attention 直接访问该结构,改动时容易踩坑。

问题场景

用户在运行 Transformers 的 encoder-decoder 模型(示例为 google-t5/t5-small)做生成时,手动创建了 StaticCacheEncoderDecoderCache,然后把 past_key_values=cache 传给 model.generate(**input_ids, past_key_values=cache)。该场景来自 executorch 导出需求,希望所有内存预先分配、cache 可复用,因此试图让 static cache 的容量大于实际生成所需的 batch size 和 encoder 序列长度。

报错原文

Allow static cache to be larger than sequence length / batch size for encoder-decoder models

原因分析

主要原因在于 StaticCache 当前按“精确尺寸”约束初始化与使用,而不是按“容量上限”工作:

  • cross-attention cache 的尺寸必须严格等于 encoder 序列长度;
  • self-attention 和 cross-attention cache 的 batch size 必须严格等于生成时的 batch size。

因此,当传入的输入 batch size 或 encoder 序列长度小于初始化时的值,cache 无法被复用;多次调用 generate 且长度不同时,cross-attention static cache 会针对每个 encoder 长度重新构建。Issue 讨论中确认,decoder-only 模型已支持这种复用,encoder cache 部分需要对齐同样的行为;而变 batch size 的复用已由后续 PR 支持。

环境排查

  • 确认使用的 Transformers 版本是否已包含 #37394 带来的“小于 max_batch_size 的 batch 复用”支持。
  • 确认是手动构造 StaticCache/EncoderDecoderCache,还是通过 generate(cache_implementation=...) 内部创建。
  • 确认 cache 初始化时的 max_batch_size 与当前输入 batch size 的关系。
  • 确认 cross-attention cache 的 max_cache_len 是否严格等于 encoder 序列长度。
  • 确认模型是否为 T5/Whisper 及其派生模型,因为这些模型会为 cross-attention 直接访问 cache。
  • Issue 未确认操作系统、Python、CUDA、PyTorch、显卡型号,排查时需自行记录。

解决步骤

  1. 优先在 generate() 中使用 cache_implementation 让内部管理 cache,而不是手动构造并复用 cache。
  2. 若必须手动构造 cache,则把 batch size 设置为当前输入的实际 batch size,并把 cross-attention cache 长度设置为与 encoder 序列长度完全匹配。
  3. 若使用较旧的 Transformers,升级到已包含 #37394 的版本,以获得 batch size 小于 max_batch_size 时的复用能力。
  4. 对于变 encoder 序列长度这一半问题,目前没有已确认的修复方案;讨论中表示相关 PR 仍欢迎,可跟踪后续进展。
  5. 不要依赖 past_key_value.key_cache[self.layer_idx] 直接访问 cache 作为稳定接口,尤其在裁剪或复用 cache 时。

验证方法

用小于 max_batch_size 的 batch 和不超过 max_cache_len 的 encoder 序列长度重复调用生成,确认不再为每个 encoder 长度重建 cross-attention static cache,且生成结果与不传外部 cache 时一致。若仍出现尺寸不匹配,说明当前版本尚未覆盖变 encoder 长度这一部分。

参考来源

huggingface/transformers #35444

GamsGo AI

AI 工具推荐

想把多个 AI 模型放在一个入口?

GamsGo AI 集成 ChatGPT、DeepSeek、Gemini、Claude、Midjourney、Veo 等常用模型,适合写作、绘图、视频和日常 AI 工作流。

了解 GamsGo AI

推广链接:通过此链接购买,我可能获得佣金,不影响你的价格。

这个方案解决了吗?

celebrityanime
celebrityanime
文章: 23905

发表回复

您的邮箱地址不会被公开。 必填项已用 * 标注