基于Cityscapes的语义分割:浮点标签下交叉熵损失使用难题
解决Cityscapes语义分割中CrossEntropyLoss的标签类型与忽略索引问题
核心问题根源
你犯了一个常见错误:把ToTensor()用在了标签数据上。Cityscapes的标签不是像素强度值,而是类别索引整数(比如道路是0,建筑是1,忽略区域是255)。ToTensor()会把[0,255]的整数归一化到[0,1]的浮点数,比如原本的类别1会变成1/255≈0.0039,转成LongTensor时会被取整为0;忽略索引255会变成1.0,转long后也是0,这就导致所有标签都失去了有效类别信息。
正确处理流程
1. 分离图像与标签的变换
图像可以正常用ToTensor()做归一化,但标签必须保留原始的整数索引,不能做归一化:
from torchvision import transforms from PIL import Image import torch import numpy as np # 图像变换:ToTensor将[0,255]转成[0,1],可选加上数据集均值方差归一化 img_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 标签变换:直接转成LongTensor,保留原始整数索引 label_transform = transforms.Compose([ transforms.Lambda(lambda x: torch.tensor(np.array(x), dtype=torch.long)) ]) # 读取数据时分别应用变换 img = Image.open(img_path).convert('RGB') label = Image.open(label_path) img_tensor = img_transform(img) label_tensor = label_transform(label)
2. 正确配置CrossEntropyLoss的忽略索引
Cityscapes的忽略区域标签是255,初始化CrossEntropyLoss时指定ignore_index=255,损失计算时会自动跳过这些区域:
import torch.nn as nn # 假设模型输出20个有效类别(Cityscapes的标准类别数) loss_fn = nn.CrossEntropyLoss(ignore_index=255) # 计算损失:模型输出为(B, C, H, W)的logits,标签为(B, H, W)的LongTensor logits = model(img_tensor.unsqueeze(0)) # 增加batch维度 loss = loss_fn(logits, label_tensor.unsqueeze(0))
额外说明
- CrossEntropyLoss要求标签必须是
LongTensor类型,且每个元素是对应类别的索引(0到C-1,C为类别数),不能是浮点数。 - 如果使用自定义数据集加载器,务必确保标签数据在加载过程中没有被做归一化处理,始终保留原始的整数类别ID。
内容的提问来源于stack exchange,提问作者Jacob Kang
相关产品推荐
相关产品推荐

