矩形图像(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
相关产品推荐
相关产品推荐

