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

如何在PyTorch框架下的计算机视觉任务中使用软标签?

基于PyTorch的软标签图像数据集训练实现方案

第一步:自定义数据集类读取数据

你可以通过继承PyTorch的Dataset类,同时读取图像文件和csv里存储的软标签,以下是可直接复用的实现:

import os
import pandas as pd
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from PIL import Image

class SoftLabelImageDataset(Dataset):
    def __init__(self, csv_path, img_root_dir, transform=None):
        # 读取存储软标签的csv文件,默认第一列为图像文件名,后续列依次对应各类别的概率
        self.label_df = pd.read_csv(csv_path)
        self.img_root_dir = img_root_dir
        self.transform = transform
        # 自动计算类别数:总列数减1(排除第一列的图像文件名)
        self.num_classes = len(self.label_df.columns) - 1

    def __len__(self):
        return len(self.label_df)

    def __getitem__(self, idx):
        # 拼接图像完整路径并读取
        img_path = os.path.join(self.img_root_dir, self.label_df.iloc[idx, 0])
        image = Image.open(img_path).convert("RGB")
        # 读取软标签并转换为浮点型张量
        soft_label = torch.tensor(self.label_df.iloc[idx, 1:].values, dtype=torch.float32)
        # 可选:如果csv里的概率和不为1,可加下面这行做归一化
        # soft_label = soft_label / soft_label.sum()
        
        if self.transform:
            image = self.transform(image)
        
        return image, soft_label

完成数据集类定义后,按如下方式构建训练所需的DataLoader:

# 图像预处理逻辑可根据自身需求调整
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# 替换为你自己的csv路径、图像存储根目录
train_dataset = SoftLabelImageDataset(
    csv_path="./train_labels.csv",
    img_root_dir="./train_images",
    transform=transform
)

train_dataloader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)

第二步:软标签适配的损失函数选择

PyTorch原生的CrossEntropyLoss默认支持硬标签输入,针对软标签场景可选择以下两种方案:

# 方案1:PyTorch 1.10及以上版本可直接使用CrossEntropyLoss
# 当输入的target形状和模型输出logits形状一致时,会自动计算两个概率分布的交叉熵
criterion = torch.nn.CrossEntropyLoss()

# 方案2:兼容旧版本的实现,使用KLDivLoss
# 本质和交叉熵优化等价,训练时需要先对模型输出做log_softmax处理
criterion = torch.nn.KLDivLoss(reduction="batchmean")

第三步:完整训练流程示例

以下是基于ResNet18的训练样例,你可以替换为自己的模型结构:

import torchvision.models as models

# 加载模型,修改输出头匹配你的类别数,这里假设类别数为10
model = models.resnet18(pretrained=False)
num_ftrs = model.fc.in_features
model.fc = torch.nn.Linear(num_ftrs, 10)

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

# 定义优化器
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

# 训练循环
num_epochs = 10
for epoch in range(num_epochs):
    model.train()
    total_loss = 0.0
    for images, soft_labels in train_dataloader:
        images = images.to(device)
        soft_labels = soft_labels.to(device)
        
        optimizer.zero_grad()
        outputs = model(images)
        
        # 使用CrossEntropyLoss时直接计算损失
        loss = criterion(outputs, soft_labels)
        
        # 若使用KLDivLoss,替换为下面两行代码
        # log_prob = torch.nn.functional.log_softmax(outputs, dim=1)
        # loss = criterion(log_prob, soft_labels)
        
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item() * images.size(0)
    
    avg_loss = total_loss / len(train_dataset)
    print(f"Epoch {epoch+1}/{num_epochs},平均损失:{avg_loss:.4f}")

注意事项

  • csv文件的列顺序要固定,列顺序和类别的对应关系需要在训练、验证、测试阶段保持一致
  • 软标签必须为浮点型张量,不能使用长整型(long),否则会触发损失函数报错
  • 推理阶段不需要特殊处理,直接对模型输出做softmax即可得到预测的类别概率分布

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 21:36:05