Cityscapes数据集图像分割训练:如何将类别从30类缩减至7类?
Cityscapes数据集类别缩减(30类→7类)的高效实现方法
当然可以将Cityscapes的30类缩减至7类,你之前逐图修改类别值的方法效率低下,核心原因是单图循环操作的冗余开销。以下是两种高效的实现方案:
方案一:批量预处理转换(一次性生成新标签集)
- 定义类别映射规则:先根据任务需求,把Cityscapes的30个原始类别ID映射到7个新类别ID。比如将所有机动车(car、truck、bus等)合并为一类,行人与骑行者合并为一类,建筑、道路等各归为一类,示例映射字典如下:
# 示例:根据实际需求调整原ID和新ID的对应关系 class_mapping = { 7: 0, # road → 新类别0 11: 1, # building → 新类别1 13: 1, # wall → 合并到building类(新类别1) 14: 1, # fence → 合并到building类(新类别1) 20: 2, # person → 新类别2 21: 2, # rider → 合并到person类(新类别2) 24: 3, # car → 新类别3 25: 3, # truck → 合并到car类(新类别3) 26: 3, # bus → 合并到car类(新类别3) # 其他原类别按需求映射到7类中的对应ID,未列出的可设为背景或忽略 } - 生成全局映射数组:利用numpy的向量化特性,创建一个索引为原类别ID、值为新类别ID的数组,这是实现批量转换的关键:
import numpy as np from PIL import Image import os # Cityscapes官方最大类别ID为34 max_old_id = 34 mapping_array = np.zeros(max_old_id + 1, dtype=np.uint8) # 填充映射关系,未指定的原ID默认映射为背景(0) for old_id, new_id in class_mapping.items(): mapping_array[old_id] = new_id - 批量处理所有标签图:遍历标签文件夹,用numpy加载图像后直接通过索引完成转换,速度远快于单图循环:
input_label_dir = "/path/to/gtFine/train" output_label_dir = "/path/to/gtFine_7class/train" # 创建输出目录 os.makedirs(output_label_dir, exist_ok=True) for city_dir in os.listdir(input_label_dir): city_path = os.path.join(input_label_dir, city_dir) output_city_path = os.path.join(output_label_dir, city_dir) os.makedirs(output_city_path, exist_ok=True) for filename in os.listdir(city_path): if "labelIds" in filename: label_path = os.path.join(city_path, filename) # 加载原始标签图(单通道整数图像) label_img = np.array(Image.open(label_path)) # 批量转换类别 new_label_img = mapping_array[label_img] # 保存新标签图 Image.fromarray(new_label_img).save(os.path.join(output_city_path, filename))
方案二:数据加载时动态映射(推荐,无需修改原数据集)
如果不想占用额外磁盘空间存储新标签集,可以在训练阶段的DataLoader中动态完成类别映射,省去预处理时间:
- 同样定义上述的
mapping_array - 在自定义Dataset的
__getitem__方法中,读取原始标签后直接转换:from torch.utils.data import Dataset class Cityscapes7ClassDataset(Dataset): def __init__(self, root, split, transform=None): self.root = root self.split = split self.transform = transform # 初始化类别映射数组 max_old_id = 34 self.mapping_array = np.zeros(max_old_id + 1, dtype=np.uint8) class_mapping = { # 同方案一的映射规则 7: 0, 11:1, 13:1, 14:1, 20:2, 21:2, 24:3, 25:3, 26:3, ... } for old_id, new_id in class_mapping.items(): self.mapping_array[old_id] = new_id # 加载图像路径列表(省略路径解析逻辑,可参考Cityscapes官方数据集加载代码) self.img_paths, self.label_paths = self._load_paths() def __getitem__(self, idx): img = Image.open(self.img_paths[idx]).convert("RGB") label_img = np.array(Image.open(self.label_paths[idx])) # 动态转换类别 new_label = self.mapping_array[label_img] if self.transform: img, new_label = self.transform(img, new_label) return img, new_label def __len__(self): return len(self.img_paths)
注意事项
- 确保所有原始类别ID都被映射到7个目标类别中,避免出现未定义的ID导致训练错误
- 映射规则需严格匹配任务需求,比如若任务关注交通参与者,可将行人、骑行者、机动车分别归为独立类别,非目标类别统一设为背景
- 若使用PyTorch训练,可将
mapping_array转为torch.tensor并移至GPU,进一步提升转换速度
内容的提问来源于stack exchange,提问作者user21075974
相关产品推荐
相关产品推荐

