如何训练每个标签组可输出多类别的Multi-label classification模型
多标签分组分类模型训练实现方案
1. 修正模型定义错误
你当前的初始化方法存在参数笔误,需先修正类别数声明:
class MyModel(nn.Module): # 传入两个分组的类别数 def __init__(self, n_classes_layout=5, n_classes_style=4): super().__init__() self.base_model = models.resnet50(pretrained=True).to(device) last_channel = self.base_model.fc.in_features self.base_model.fc = nn.Sequential() self.layout = nn.Sequential( nn.Dropout(0.2), nn.Linear(last_channel, n_classes_layout), nn.Sigmoid() ) self.style = nn.Sequential( nn.Dropout(0.2), nn.Linear(last_channel, n_classes_style), nn.Sigmoid() ) def forward(self, x): base = self.base_model(x) return self.layout(base), self.style(base)
2. 补全损失函数
两个分组均为独立的多标签分类任务,采用二元交叉熵损失加权求和即可:
def loss_fn(outputs, targets, layout_loss_weight=0.5): o1, o2 = outputs t1, t2 = targets # 分别计算两个分组的BCELoss loss_layout = nn.BCELoss()(o1, t1.float()) loss_style = nn.BCELoss()(o2, t2.float()) # 加权合并损失,可根据任务优先级调整权重 return layout_loss_weight * loss_layout + (1-layout_loss_weight) * loss_style
3. 数据集预处理要求
- 每个样本的标签按(layout_label, style_label)格式组织:
- layout_label为长度为5的0/1张量,其中3个位置为1
- style_label为长度为4的0/1张量,其中2个位置为1
- 图像预处理对齐ResNet预训练逻辑:resize到224/256尺寸、随机水平翻转、随机裁剪,归一化使用ImageNet数据集的均值
[0.485, 0.456, 0.406]、方差[0.229, 0.224, 0.225]
4. 训练流程配置
4.1 优化器设置
推荐分两阶段微调,收敛更快效果更好:
- 第一阶段冻结ResNet50主干权重,仅训练两个分类头:使用Adam优化器,学习率设置为
1e-3,训练3-5轮 - 第二阶段解冻全部权重,整体微调:学习率调整为
1e-4,训练5-10轮,可加入学习率衰减策略
4.2 训练循环示例
model = MyModel().to(device) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) for epoch in range(total_epochs): model.train() for img, (layout_target, style_target) in dataloader: img = img.to(device) layout_target = layout_target.to(device) style_target = style_target.to(device) optimizer.zero_grad() layout_out, style_out = model(img) loss = loss_fn((layout_out, style_out), (layout_target, style_target)) loss.backward() optimizer.step()
5. 推理后处理适配输出要求
不要使用固定0.5阈值筛选正类,直接按概率取TopK即可匹配输出要求:
def post_process(layout_out, style_out): # layout取Top3的位置设为1 _, layout_topk = torch.topk(layout_out, k=3, dim=-1) layout_res = torch.zeros_like(layout_out).scatter_(1, layout_topk, 1).int() # style取Top2的位置设为1 _, style_topk = torch.topk(style_out, k=2, dim=-1) style_res = torch.zeros_like(style_out).scatter_(1, style_topk, 1).int() return layout_res, style_res
内容的提问来源于stack exchange,提问作者Tùng Dương Quang
相关产品推荐
相关产品推荐

