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

如何训练每个标签组可输出多类别的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 优化器设置

推荐分两阶段微调,收敛更快效果更好:

  1. 第一阶段冻结ResNet50主干权重,仅训练两个分类头:使用Adam优化器,学习率设置为1e-3,训练3-5轮
  2. 第二阶段解冻全部权重,整体微调:学习率调整为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 12:39:04