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

能否实现子类(类中类)预测?求三级层级图像分类方案

三级层级物体分类解决方案(PyTorch实现)

核心思路

三级分类本质是多任务并行分类:同时预测主类、一级子类、二级子类,最后将结果拼接为Class-主类ID-子类1ID-子类2ID格式。这种方式比直接把每个三级类别当作独立分类任务更高效,尤其适合子类存在特征共享的场景。

实现步骤

1. 数据准备

假设数据集标签为三元组格式(主类ID, 一级子类ID, 二级子类ID),比如(1,1,1)对应Class-1-1-1。自定义数据集类加载数据时需同时读取三个层级的标签:

from torch.utils.data import Dataset
from PIL import Image
import torchvision.transforms as transforms

class MultiLevelDataset(Dataset):
    def __init__(self, img_paths, labels, transform=None):
        self.img_paths = img_paths
        self.labels = labels  # 格式示例:[(1,1,1), (1,2,3), ...]
        self.transform = transform or transforms.Compose([
            transforms.Resize((224,224)),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
        ])
    
    def __getitem__(self, idx):
        img = Image.open(self.img_paths[idx]).convert('RGB')
        img = self.transform(img)
        main_label, sub1_label, sub2_label = self.labels[idx]
        return img, (main_label-1, sub1_label-1, sub2_label-1)  # 转0-based标签适配CrossEntropyLoss
    
    def __len__(self):
        return len(self.img_paths)

2. 模型构建

采用预训练模型做特征提取,分三个独立分支分别预测三个层级的类别:

import torch
import torch.nn as nn
from torchvision import models

class MultiLevelClassifier(nn.Module):
    def __init__(self, num_main, num_sub1, num_sub2):
        super().__init__()
        # 预训练特征提取 backbone
        self.backbone = models.resnet50(pretrained=True)
        self.feature_dim = self.backbone.fc.in_features
        self.backbone.fc = nn.Identity()  # 移除原全连接层
        
        # 三个分类头
        self.main_head = nn.Linear(self.feature_dim, num_main)
        self.sub1_head = nn.Linear(self.feature_dim, num_sub1)
        self.sub2_head = nn.Linear(self.feature_dim, num_sub2)
    
    def forward(self, x):
        features = self.backbone(x)
        main_logits = self.main_head(features)
        sub1_logits = self.sub1_head(features)
        sub2_logits = self.sub2_head(features)
        return main_logits, sub1_logits, sub2_logits

3. 训练流程

使用交叉熵损失分别计算三个层级的损失,加权求和后反向传播:

# 初始化模型、损失、优化器
model = MultiLevelClassifier(num_main=3, num_sub1=5, num_sub2=4)  # 根据实际类别数调整
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)

# 训练循环示例
for epoch in range(15):
    model.train()
    total_loss = 0.0
    for imgs, (main_labels, sub1_labels, sub2_labels) in train_dataloader:
        imgs = imgs.to(device)
        main_labels = main_labels.to(device)
        sub1_labels = sub1_labels.to(device)
        sub2_labels = sub2_labels.to(device)
        
        optimizer.zero_grad()
        main_logits, sub1_logits, sub2_logits = model(imgs)
        
        # 计算三个层级损失,可根据需求调整权重
        loss_main = criterion(main_logits, main_labels)
        loss_sub1 = criterion(sub1_logits, sub1_labels)
        loss_sub2 = criterion(sub2_logits, sub2_labels)
        total_batch_loss = loss_main + loss_sub1 + loss_sub2
        
        total_batch_loss.backward()
        optimizer.step()
        total_loss += total_batch_loss.item()
    
    print(f"Epoch {epoch+1}, Avg Loss: {total_loss/len(train_dataloader):.4f}")

4. 推理与结果生成

推理时对三个分支的输出取argmax,再拼接为要求的格式:

model.eval()
with torch.no_grad():
    for imgs, _ in test_dataloader:
        imgs = imgs.to(device)
        main_logits, sub1_logits, sub2_logits = model(imgs)
        
        # 预测结果转1-based ID
        main_pred = torch.argmax(main_logits, dim=1).cpu().numpy() + 1
        sub1_pred = torch.argmax(sub1_logits, dim=1).cpu().numpy() + 1
        sub2_pred = torch.argmax(sub2_logits, dim=1).cpu().numpy() + 1
        
        # 生成最终分类结果
        for m, s1, s2 in zip(main_pred, sub1_pred, sub2_pred):
            print(f"Class-{m}-{s1}-{s2}")

可选优化方向

  • 层级约束损失:如果子类依赖主类(如主类1的子类与主类2的子类无重叠),可添加约束逻辑,仅当主类预测正确时才计算子类损失,减少无效预测。
  • 特征分支拆分:若不同层级需要差异化特征,可在backbone后为每个层级添加独立的特征提取模块,而非共享同一特征。

内容的提问来源于stack exchange,提问作者VittalAB

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 07:25:21