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

多标签分类任务输入输出批量大小不一致问题求助

问题排查与解决方案

一、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],经过三次卷积+池化后,特征图尺寸变化为:

  1. 第一次卷积+池化:224 → 112
  2. 第二次卷积+池化:112 → 56
  3. 第三次卷积+池化: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

额外验证建议

  1. 在forward方法中打印每一步张量形状,确认特征尺寸变化:
print("conv1+pool后形状:", x.shape)
print("conv2+pool后形状:", x.shape)
print("conv3+pool后形状:", x.shape)
print("展平后形状:", x.shape)
  1. 用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 10:02:10