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

矩形图像(128,256)适配Vision Transformer的patch_size设置问题

解决Vision Transformer适配矩形图像(128,256)的问题

核心问题分析

  • 用patch_size=16时,128/16=8、256/16=16,总patch数为128,但原ViT代码大概率硬编码了224×224正方形对应的14×14=196个patch,导致后续层维度不匹配,触发RuntimeError。
  • 设patch_size=(16,32)时报错,是因为原实现仅支持等宽高的int型patch_size,未处理tuple类型参数,导致取模、张量拼接等操作失败。

正确的patch_size配置与代码修正

1. 重写支持矩形Patch的嵌入层

修改Patch Embedding模块,允许传入tuple型patch尺寸,并重新计算patch数量:

import torch
import torch.nn as nn

class PatchEmbed(nn.Module):
    def __init__(self, img_size=(128,256), patch_size=(16,32), in_chans=3, embed_dim=768):
        super().__init__()
        self.img_size = img_size
        self.patch_size = patch_size
        # 确保图像尺寸能被patch尺寸整除,避免残差区域
        assert img_size[0] % patch_size[0] == 0 and img_size[1] % patch_size[1] == 0, \
            "图像宽高必须分别能被patch的宽高整除"
        self.num_patches = (img_size[0] // patch_size[0]) * (img_size[1] // patch_size[1])
        
        # 用Conv2d实现分块嵌入,卷积核尺寸对应patch_size
        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)

    def forward(self, x):
        # 输入x shape: (batch_size, channels, height, width)
        x = self.proj(x)  # 输出shape: (B, embed_dim, H/patch_h, W/patch_w)
        x = x.flatten(2)  # 展平patch维度: (B, embed_dim, num_patches)
        x = x.transpose(1, 2)  # 调整维度顺序: (B, num_patches, embed_dim)
        return x

2. 修改ViT主模型适配矩形输入

替换原模型的PatchEmbed模块,传入矩形参数并调整位置嵌入维度:

class ViT(nn.Module):
    def __init__(self, img_size=(128,256), patch_size=(16,32), in_chans=3, embed_dim=768, 
                 depth=12, num_heads=12, mlp_ratio=4., num_classes=10):
        super().__init__()
        self.patch_embed = PatchEmbed(
            img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim
        )
        num_patches = self.patch_embed.num_patches
        
        # 类别嵌入与位置嵌入
        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
        # 位置嵌入长度为patch数+1(包含cls_token)
        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
        nn.init.trunc_normal_(self.pos_embed, std=.02)
        
        # Transformer编码器层(简化实现,实际可复用torch官方TransformerEncoder)
        self.blocks = nn.ModuleList([
            nn.TransformerEncoderLayer(
                d_model=embed_dim, nhead=num_heads, dim_feedforward=int(embed_dim*mlp_ratio),
                batch_first=True, norm_first=True
            ) for _ in range(depth)
        ])
        self.norm = nn.LayerNorm(embed_dim)
        self.head = nn.Linear(embed_dim, num_classes)

    def forward(self, x):
        batch_size = x.shape[0]
        x = self.patch_embed(x)
        
        # 扩展cls_token到当前batch尺寸
        cls_tokens = self.cls_token.expand(batch_size, -1, -1)
        # 拼接cls_token与patch嵌入
        x = torch.cat((cls_tokens, x), dim=1)
        
        # 添加位置嵌入
        x = x + self.pos_embed
        
        # 编码器前向传播
        for blk in self.blocks:
            x = blk(x)
        x = self.norm(x)
        
        # 用cls_token的输出做分类
        return self.head(x[:, 0])

3. 关键注意事项

  • 本次配置中(128,256)与(16,32)刚好整除,得到8×8=64个patch,位置嵌入维度完全匹配,不会出现尺寸错误。
  • 如果无法找到整除的patch尺寸,可先对图像做比例缩放,再裁剪到能被patch整除的尺寸,避免拉伸失真。
  • 训练时输入图像必须严格保持(128,256)尺寸,建议用中心裁剪或随机裁剪保留原始矩形比例,不要强制拉伸为正方形。

4. 测试代码验证

# 测试模型前向传播是否正常
model = ViT(img_size=(128,256), patch_size=(16,32), num_classes=10)
test_input = torch.randn(2, 3, 128, 256)
output = model(test_input)
print(output.shape)  # 应输出torch.Size([2, 10])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 18:43:20