ADE20K数据集实战指南:从数据加载到语义分割模型训练

1次阅读
没有评论

共计 3484 个字符,预计需要花费 9 分钟才能阅读完成。

image.webp

背景介绍

ADE20K 是 MIT 发布的场景解析数据集,包含 2 万多张图像和 150 个精细标注的语义类别(从天空、建筑到盆栽植物等)。这个数据集在场景理解任务中具有重要地位,因为:

ADE20K 数据集实战指南:从数据加载到语义分割模型训练

  • 类别覆盖全面,包含室内外场景
  • 标注质量高,每个像素都有精确标签
  • 场景复杂度适中,适合作为研究基准

相比 PASCAL VOC(20 类)和 Cityscapes(30 类),ADE20K 提供了更丰富的语义粒度。例如 ” 椅子 ” 在 VOC 中是一个类别,而在 ADE20K 中会细分为 ” 办公椅 ”、” 餐椅 ” 等子类。

常见痛点分析

实际使用中发现三个主要挑战:

  1. 数据加载效率:解压后数据集超过 5GB,传统加载方式会导致内存溢出

  2. 标注解析复杂:同时使用 JSON 存储物体元数据和 PNG 存储像素级标签,需要特殊处理

  3. 类别不平衡:像 ” 天空 ” 这样的类别出现频率是 ” 浴缸 ” 的 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)

数据增强重点

对小样本类别使用针对性增强:

  1. 随机裁剪时确保包含稀有类别
  2. 对稀有类别图像提高重复采样率
  3. 颜色增强只应用于常见类别

完整训练示例

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()

关键避坑指南

  1. 内存优化
  2. 设置 pin_memory=True 加速 CPU 到 GPU 的数据传输
  3. 使用 torch.cuda.empty_cache() 定期清理显存

  4. 多 GPU 训练

  5. 使用 DistributedSampler 确保数据均匀分配
  6. 注意验证集的评估要在主进程进行

  7. 评估指标

  8. mIoU 计算时需要忽略 255(边界 / 无效像素)
  9. 官方提供的评估代码会处理类别重新映射
# 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 包含物体实例

后续建议:

  1. 使用在 ADE20K 预训练的模型进行迁移学习
  2. 尝试 HRNet 等新颖架构提升小物体识别
  3. 结合场景图生成等高层语义任务

完整可运行代码参考:ADE20K Colab Notebook (虚拟链接)

通过本指南,你应该能够:
– 高效加载和预处理 ADE20K 数据
– 处理复杂的类别不平衡问题
– 训练并评估语义分割模型
– 避免常见的内存和计算陷阱

正文完
 0
评论(没有评论)