共计 3484 个字符,预计需要花费 9 分钟才能阅读完成。
背景介绍
ADE20K 是 MIT 发布的场景解析数据集,包含 2 万多张图像和 150 个精细标注的语义类别(从天空、建筑到盆栽植物等)。这个数据集在场景理解任务中具有重要地位,因为:

- 类别覆盖全面,包含室内外场景
- 标注质量高,每个像素都有精确标签
- 场景复杂度适中,适合作为研究基准
相比 PASCAL VOC(20 类)和 Cityscapes(30 类),ADE20K 提供了更丰富的语义粒度。例如 ” 椅子 ” 在 VOC 中是一个类别,而在 ADE20K 中会细分为 ” 办公椅 ”、” 餐椅 ” 等子类。
常见痛点分析
实际使用中发现三个主要挑战:
-
数据加载效率:解压后数据集超过 5GB,传统加载方式会导致内存溢出
-
标注解析复杂:同时使用 JSON 存储物体元数据和 PNG 存储像素级标签,需要特殊处理
-
类别不平衡:像 ” 天空 ” 这样的类别出现频率是 ” 浴缸 ” 的 300 倍以上
技术实现方案
1. 高效数据加载
使用 PyTorch 的 Dataset 类实现按需加载,避免一次性读取所有数据:
from torch.utils.data import Dataset
from PIL import Image
import json
import os
class ADE20KDataset(Dataset):
def __init__(self, root, transform=None):
self.root = root
self.transform = transform
self.img_dir = os.path.join(root, 'images')
self.ann_dir = os.path.join(root, 'annotations')
# 获取所有文件名(不带后缀)
self.filenames = [f.split('.')[0] for f in os.listdir(self.img_dir)
if f.endswith('.jpg')]
def __len__(self):
return len(self.filenames)
def __getitem__(self, idx):
base_name = self.filenames[idx]
# 加载图像
img_path = os.path.join(self.img_dir, f"{base_name}.jpg")
image = Image.open(img_path).convert('RGB')
# 加载标注
ann_path = os.path.join(self.ann_dir, f"{base_name}.png")
annotation = Image.open(ann_path)
if self.transform:
image, annotation = self.transform(image, annotation)
return image, annotation
关键优化点:
- 只存储文件名而非完整路径
- 使用 PIL 的惰性加载特性
- 支持 transform 组合操作
2. 标注解析技巧
ADE20K 的标注 PNG 使用特殊的颜色编码:
import numpy as np
# 示例:将标注 PNG 转换为类别 ID 矩阵
def decode_annotation(annotation):
"""
参数:
annotation: PIL.Image 对象
返回:
np.ndarray 形状(H,W) 值为 0 -149 的类别 ID
"""
arr = np.array(annotation)
# R 通道存储物体 ID,G 通道存储部件 ID
return arr[..., 0].astype(np.int64) - 1 # 减 1 使 ID 从 0 开始
3. 处理类别不平衡
样本加权策略
from sklearn.utils.class_weight import compute_class_weight
# 计算类别权重
def calculate_weights(dataset, n_classes=150):
"""遍历数据集统计类别分布"""
pixel_counts = np.zeros(n_classes)
for _, ann in dataset:
classes = np.unique(decode_annotation(ann))
for cls in classes:
if 0 <= cls < n_classes: # 过滤无效 ID
pixel_counts[cls] += 1
# 计算权重(出现频率越低的类别权重越高)
weights = compute_class_weight(
'balanced',
classes=np.arange(n_classes),
y=pixel_counts
)
return torch.tensor(weights, dtype=torch.float32)
数据增强重点
对小样本类别使用针对性增强:
- 随机裁剪时确保包含稀有类别
- 对稀有类别图像提高重复采样率
- 颜色增强只应用于常见类别
完整训练示例
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torch.optim import Adam
# 1. 初始化
model = UNet(num_classes=150) # 示例模型
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 2. 数据准备
dataset = ADE20KDataset('path/to/ADE20K', transform=augmentations)
weights = calculate_weights(dataset)
train_loader = DataLoader(
dataset,
batch_size=8,
shuffle=True,
pin_memory=True, # 加速 GPU 传输
num_workers=4
)
# 3. 损失函数(带类别权重)
criterion = nn.CrossEntropyLoss(weight=weights.to(device))
# 4. 训练循环
for epoch in range(100):
model.train()
for images, annotations in train_loader:
images = images.to(device)
labels = decode_annotation(annotations).to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
关键避坑指南
- 内存优化:
- 设置
pin_memory=True加速 CPU 到 GPU 的数据传输 -
使用
torch.cuda.empty_cache()定期清理显存 -
多 GPU 训练:
- 使用
DistributedSampler确保数据均匀分配 -
注意验证集的评估要在主进程进行
-
评估指标:
- mIoU 计算时需要忽略 255(边界 / 无效像素)
- 官方提供的评估代码会处理类别重新映射
# mIoU 计算示例
def compute_miou(preds, labels, n_classes=150):
"""
preds: (B, C, H, W)
labels: (B, H, W)
"""
# 转换预测结果为类别 ID
preds = torch.argmax(preds, dim=1)
# 初始化混淆矩阵
cm = torch.zeros((n_classes, n_classes), dtype=torch.int64)
# 统计每个类别的预测情况
for p, l in zip(preds.flatten(), labels.flatten()):
if 0 <= l < n_classes: # 忽略 255
cm[l, p] += 1
# 计算各类 IoU
intersection = torch.diag(cm)
union = cm.sum(0) + cm.sum(1) - intersection
iou = intersection.float() / union.float()
return iou.mean().item() # 返回 mIoU
总结与延伸
相比其他数据集:
| 数据集 | 类别数 | 图像数 | 特点 |
|---|---|---|---|
| ADE20K | 150 | 20k+ | 室内外均衡 |
| Cityscapes | 30 | 5k | 街景专用 |
| COCO-Stuff | 172 | 164k | 包含物体实例 |
后续建议:
- 使用在 ADE20K 预训练的模型进行迁移学习
- 尝试 HRNet 等新颖架构提升小物体识别
- 结合场景图生成等高层语义任务
完整可运行代码参考:ADE20K Colab Notebook (虚拟链接)
通过本指南,你应该能够:
– 高效加载和预处理 ADE20K 数据
– 处理复杂的类别不平衡问题
– 训练并评估语义分割模型
– 避免常见的内存和计算陷阱
正文完
