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

视频分类Transformer模型中Patch Embedding维度错误的解决方法

视频分类Transformer模型Patch Embedding形状不匹配问题解决

问题详情

输入张量形状为 [batch=16, channels=3, frames=16, H=224, W=224](维度顺序:b, c, t, h, w),Patch Embedding实现代码如下:

patch_dim = in_channels * patch_size ** 2
self.to_patch_embedding = nn.Sequential(
        Rearrange('b t c (h p1) (w p2) -> b t (h w) (p1 p2 c)', p1 = patch_size, p2 = patch_size),
        nn.Linear(patch_dim, dim),     # 报错行
    )

使用参数:

patch_size = 16 
dim = 192
in_channels = 3

报错信息

运行时触发RuntimeError:

Exception has occured: RuntimeError
mat1 and mat2 shapes cannot be multiplied (9408x4096 and 768x192)

问题根源

Rearrange层维度匹配错误:输入张量的维度顺序是b, c, t, h, w,但代码里的Rearrange pattern写的是b t c (h p1) (w p2),相当于把输入的第2个维度(channels=3)当成了frames,第3个维度(frames=16)当成了channels。

错误变换后,每个patch的特征维度变成了16*16*16=4096(本该是16*16*3=768),而你计算的patch_dim=768,导致Linear层的输入维度(4096)和权重维度(768×192)不匹配,触发形状错误。

解决方法

修正Rearrange的输入维度匹配,让它对应实际的输入张量维度顺序:

方案1:直接修改Rearrange pattern

patch_dim = in_channels * patch_size ** 2
self.to_patch_embedding = nn.Sequential(
        Rearrange('b c t (h p1) (w p2) -> b t (h w) (p1 p2 c)', p1 = patch_size, p2 = patch_size),
        nn.Linear(patch_dim, dim),
    )

方案2:先调整维度顺序再做Patch分割

如果习惯保持t在c前面的pattern,可以先加一步维度重排:

patch_dim = in_channels * patch_size ** 2
self.to_patch_embedding = nn.Sequential(
        Rearrange('b c t h w -> b t c h w'),  # 将channels和frames维度交换位置
        Rearrange('b t c (h p1) (w p2) -> b t (h w) (p1 p2 c)', p1 = patch_size, p2 = patch_size),
        nn.Linear(patch_dim, dim),
    )

修正后,每个patch的特征维度会正确计算为16*16*3=768,和patch_dim一致,Linear层就能正常完成矩阵运算。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 19:50:26