RuntimeError: The size of tensor a (8) must match the size of tensor b (4) at non-singleton dimension 3

用户使用 Hugging Face Transformers 库(版本 5.10.1)中的 StaticCache.early_initialization 方法,为包含线性注意力层(如 layer_types = ["full_attention", "linear_attention"] )的混

快速结论:此报错通常在使用 StaticCache.early_initialization 时触发,原因是该方法错误地为线性注意力层(linear-attention layers,如 mamba/conv/linear_attention)预分配了形状为 (batch, num_heads, 0, head_dim)conv_states 张量,导致后续 update_conv_state 在复制数据时张量维度不匹配。优先检查模型配置中是否包含 linear_attention 层类型,并升级 transformers 到包含修复的版本。

问题场景

用户使用 Hugging Face Transformers 库(版本 5.10.1)中的 StaticCache.early_initialization 方法,为包含线性注意力层(如 layer_types = ["full_attention", "linear_attention"])的混合缓存(hybrid cache)预分配静态缓存时触发。该报错在调用 cache.layers[1].update_conv_state(...) 时出现。

报错原文

RuntimeError: The size of tensor a (8) must match the size of tensor b (4) at non-singleton dimension 3

原因分析

可能原因:Cache.early_initialization 方法中,代码假设所有缓存层都可以通过 (num_heads, head_dim) 推导出形状,并统一预分配张量。但对于线性注意力层(继承自 LinearAttentionCacheLayerMixin 的层,如 mamba/conv/linear_attention),其状态(如 conv_states)的维度由卷积核大小、SSM 状态大小等参数决定,无法通过 (num_heads, head_dim) 推导。因此,代码会错误地分配一个形状为 (batch, num_heads, 0, head_dim) 的张量(注意第3维长度为0),并将其标记为已初始化。当第一次执行 update_conv_state(...) 时,调用 conv_states.copy_(...) 会因张量形状不匹配(如大小为8 和 4)而抛出 RuntimeError

环境排查

  • Transformers 版本:确认版本是否为 5.10.1(该问题在后续提交 #46446 中修复)。
  • 模型配置:检查 config.layer_types 是否包含 "linear_attention" 或类似线性注意力层类型。
  • 缓存初始化方式:确认是否使用了 StaticCache.early_initialization 方法。

解决步骤

  1. 第一步:升级 Transformers 库到包含修复的版本(基于提交 #46446 的后续版本,如 >= 5.10.2 或 nightly build)。如果无法升级,可手动应用补丁。
  2. 第二步:early_initialization 方法中,添加防御性检查。使用 getattr(layer, "is_initialized", None) 读取初始化标志,如果返回 None(表示该层不使用键/值惰性初始化,例如线性注意力层)或 True(表示已初始化,防止重复初始化),则跳过该层的预分配。
  3. 第三步:修改 Cache.is_initialized 属性,使其忽略那些没有 is_initialized 属性的线性注意力层(否则会触发 AttributeError),仅基于注意力层(attention layers)的初始化状态判断混合缓存是否已初始化。
  4. 第四步:重新运行初始化流程,此时线性注意力层将保持延迟初始化,直到第一次 update_conv_state 时才基于真实状态创建张量。

验证方法

运行以下代码片段(基于 Issue 中的复现脚本),确认不再出现 RuntimeError,且 cache.layers[1].conv_states.shape 不再包含长度为0的维度(如应为 torch.Size([1, 2, 8, 4]) 而非 torch.Size([1, 2, 0, 8])):

import torch
from transformers import LlamaConfig, StaticCache

cfg = LlamaConfig(num_hidden_layers=2, num_attention_heads=4, num_key_value_heads=2, hidden_size=32)
cfg.layer_types = ["full_attention", "linear_attention"]
cache = StaticCache(config=cfg, max_cache_len=8)

cache.early_initialization(batch_size=1, num_heads=2, head_dim=8, dtype=torch.float32, device="cpu")
print(cache.layers[1].conv_states.shape)          # 应不再输出 torch.Size([1, 2, 0, 8])
cache.layers[1].update_conv_state(torch.zeros((1, 8, 4)))  # 应不再抛出 RuntimeError
print("问题已修复")

参考来源

huggingface/transformers #46439

GamsGo AI

AI 工具推荐

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

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

了解 GamsGo AI

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

这个方案解决了吗?

celebrityanime
celebrityanime
文章: 16283

发表回复

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