多标签分类任务输入输出批量大小不一致问题求助
问题排查与解决方案
一、Hugging Face ResNet18新增多分类分支失败的原因及修正
直接给model2新增classifier_artist等属性不会被模型的前向传播逻辑识别,因为Hugging Face的ResNetForImageClassification默认仅调用classifier层输出结果。要实现多任务分类,需重写模型的前向传播方法,或替换原分类器为多分支结构。
修正方案:基于Hugging Face ResNet18改造多任务模型
from transformers import ResNetForImageClassification import torch.nn as nn # 加载预训练模型 model2 = ResNetForImageClassification.from_pretrained("microsoft/resnet-18") # 定义各类别数量 num_classes_artist = 129 num_classes_style = 27 num_classes_genre = 11 # 移除原有分类器,新增三个任务分支 model2.classifier = nn.Identity() model2.classifier_artist = nn.Sequential( nn.Dropout(p=0.2, inplace=True), nn.Linear(in_features=512, out_features=num_classes_artist, bias=True) ).to(device) model2.classifier_style = nn.Sequential( nn.Dropout(p=0.2, inplace=True), nn.Linear(in_features=512, out_features=num_classes_style, bias=True) ).to(device) model2.classifier_genre = nn.Sequential( nn.Dropout(p=0.2, inplace=True), nn.Linear(in_features=512, out_features=num_classes_genre, bias=True) ).to(device) # 重写前向传播方法 def forward(self, pixel_values, labels=None): outputs = self.resnet(pixel_values) pooled_output = outputs.pooler_output # 获取ResNet池化特征,形状为[batch_size, 512] # 生成三个任务的输出 artist_logits = self.classifier_artist(pooled_output) style_logits = self.classifier_style(pooled_output) genre_logits = self.classifier_genre(pooled_output) # 损失计算逻辑(按需调整) loss = None if labels is not None: criterion_artist = nn.CrossEntropyLoss() criterion_style = nn.CrossEntropyLoss() criterion_genre = nn.CrossEntropyLoss() loss_artist = criterion_artist(artist_logits, labels["artist"]) loss_style = criterion_style(style_logits, labels["style"]) loss_genre = criterion_genre(genre_logits, labels["genre"]) loss = loss_artist + loss_style + loss_genre return {"loss": loss, "artist_logits": artist_logits, "style_logits": style_logits, "genre_logits": genre_logits} # 替换模型的forward方法 ResNetForImageClassification.forward = forward
改造后,模型前向传播会同时输出三个任务的结果,torchinfo也会显示新增的分类分支。
二、自定义WikiartModel批量大小不匹配的问题根源及修复
输出batch size变为98的核心原因是特征展平的维度计算错误:
输入为[32, 3, 224, 224],经过三次卷积+池化后,特征图尺寸变化为:
- 第一次卷积+池化:224 → 112
- 第二次卷积+池化:112 → 56
- 第三次卷积+池化:56 → 28
因此展平后的特征维度应为256 * 28 * 28,而非你写的256 * 16 * 16。错误的维度导致x.view(-1, 256*16*16)将32个样本的特征强行拼接为98个样本,直接引发batch size不匹配。
修复方案:修正展平维度计算
import torch import torch.nn as nn import torch.nn.functional as F class WikiartModel(nn.Module): def __init__(self, num_artists, num_genres, num_styles): super(WikiartModel, self).__init__() # 共享卷积层 self.conv1 = nn.Conv2d(3, 64, kernel_size=3, padding =1) self.conv2 = nn.Conv2d(64, 128, kernel_size=3, padding=1) self.conv3 = nn.Conv2d(128, 256, kernel_size=3, padding=1) self.pool = nn.MaxPool2d(2, 2) # 计算正确的展平维度:224/(2^3)=28 flattened_dim = 256 * 28 * 28 # 艺术家分类分支 self.fc_artist1 = nn.Linear(flattened_dim, 512) self.fc_artist2 = nn.Linear(512, num_artists) # 流派分类分支 self.fc_genre1 = nn.Linear(flattened_dim, 512) self.fc_genre2 = nn.Linear(512, num_genres) # 风格分类分支 self.fc_style1 = nn.Linear(flattened_dim, 512) self.fc_style2 = nn.Linear(512, num_styles) def forward(self, x): # 共享卷积层前向传播 x = self.pool(F.relu(self.conv1(x))) # 形状:[32,64,112,112] x = self.pool(F.relu(self.conv2(x))) # 形状:[32,128,56,56] x = self.pool(F.relu(self.conv3(x))) # 形状:[32,256,28,28] x = x.view(-1, 256 * 28 * 28) # 形状:[32, 256*28*28] # 艺术家分支输出 artists_out = F.relu(self.fc_artist1(x)) artists_out = self.fc_artist2(artists_out) # 形状:[32,129] # 流派分支输出 genre_out = F.relu(self.fc_genre1(x)) genre_out = self.fc_genre2(genre_out) # 形状:[32,11] # 风格分支输出 style_out = F.relu(self.fc_style1(x)) style_out = self.fc_style2(style_out) # 形状:[32,27] return artists_out, genre_out, style_out # 定义各类别数量 num_artists = 129 # 包含"未知艺术家" num_genres = 11 # 包含"未知流派" num_styles = 27
额外验证建议
- 在forward方法中打印每一步张量形状,确认特征尺寸变化:
print("conv1+pool后形状:", x.shape) print("conv2+pool后形状:", x.shape) print("conv3+pool后形状:", x.shape) print("展平后形状:", x.shape)
- 用torchinfo验证模型输入输出:
from torchinfo import summary model = WikiartModel(num_artists, num_genres, num_styles) summary(model, input_size=(32, 3, 224, 224))
此时三个分支的batch size均为32,与输入一致,损失计算不会再报错。
内容的提问来源于stack exchange,提问作者Sinestro101
相关产品推荐
相关产品推荐

