Files
FoodClassifier/SegFormer/training/dataset.py
T
2025-12-09 18:33:18 +08:00

348 lines
11 KiB
Python

"""
SegFormer训练数据集类
功能说明:
1. 加载图像和对应的分割mask
2. 数据增强(翻转、旋转、颜色变换等)
3. 预处理和标准化
4. 批量数据加载
"""
import os
import numpy as np
from PIL import Image
import torch
from torch.utils.data import Dataset
from pathlib import Path
from typing import Optional, Tuple, List
import albumentations as A
from albumentations.pytorch import ToTensorV2
class FoodSegmentationDataset(Dataset):
"""
食物分割数据集
数据格式:
- 图像:RGB图像 (.jpg, .png等)
- 标注:PNG格式mask,像素值为类别ID (0=背景, 1=食物)
"""
def __init__(
self,
data_root: str,
split: str = 'train',
image_size: int = 512,
num_classes: int = 2,
augmentation: bool = True
):
"""
初始化数据集
Args:
data_root: 数据集根目录(包含images和annotations文件夹)
split: 'train' 或 'val'
image_size: 输入图像尺寸(将resize到此尺寸)
num_classes: 类别数(包括背景)
augmentation: 是否使用数据增强(仅训练集)
"""
self.data_root = Path(data_root)
self.split = split
self.image_size = image_size
self.num_classes = num_classes
self.augmentation = augmentation and (split == 'train')
# 构建图像和mask的路径
self.images_dir = self.data_root / "images" / split
self.masks_dir = self.data_root / "annotations" / split
# 获取所有图像文件
self.image_files = sorted(list(self.images_dir.glob("*.*")))
# 过滤:只保留有对应mask的图像
self.valid_samples = []
for img_path in self.image_files:
mask_path = self.masks_dir / (img_path.stem + '.png')
if mask_path.exists():
self.valid_samples.append((img_path, mask_path))
print(f"✓ {split}集加载完成: {len(self.valid_samples)} 个样本")
# 构建数据增强pipeline
self.transform = self._build_transforms()
def _build_transforms(self):
"""
构建数据增强和预处理pipeline
使用albumentations库进行高效的数据增强
注意:对于分割任务,增强操作需要同时应用到图像和mask
"""
if self.augmentation:
# 训练集:激进的数据增强(因为数据量小)
transform = A.Compose([
# 1. 尺寸调整
A.Resize(self.image_size, self.image_size),
# 2. 几何变换(同时作用于图像和mask)
A.HorizontalFlip(p=0.5), # 50%概率水平翻转
A.VerticalFlip(p=0.3), # 30%概率垂直翻转
A.Rotate(limit=30, p=0.5), # ±30度旋转
A.ShiftScaleRotate(
shift_limit=0.1, # 平移±10%
scale_limit=0.2, # 缩放±20%
rotate_limit=20, # 旋转±20度
p=0.5
),
# 3. 颜色增强(仅作用于图像)
A.RandomBrightnessContrast(
brightness_limit=0.2,
contrast_limit=0.2,
p=0.5
),
A.HueSaturationValue(
hue_shift_limit=20,
sat_shift_limit=30,
val_shift_limit=20,
p=0.5
),
# 4. 模糊和噪声
A.OneOf([
A.GaussianBlur(blur_limit=(3, 5), p=1.0),
A.MedianBlur(blur_limit=5, p=1.0),
], p=0.3),
# 5. 标准化(使用ImageNet均值和标准差)
A.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225],
),
# 6. 转换为Tensor
ToTensorV2(),
])
else:
# 验证集:仅resize和标准化
transform = A.Compose([
A.Resize(self.image_size, self.image_size),
A.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225],
),
ToTensorV2(),
])
return transform
def __len__(self) -> int:
"""返回数据集大小"""
return len(self.valid_samples)
def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:
"""
获取一个样本
Args:
idx: 样本索引
Returns:
image: 图像Tensor (3, H, W)
mask: 分割mask Tensor (H, W),像素值为类别ID
"""
img_path, mask_path = self.valid_samples[idx]
# 1. 读取图像和mask
image = np.array(Image.open(img_path).convert('RGB'))
mask = np.array(Image.open(mask_path))
# 2. 应用数据增强
# albumentations会自动将增强同时应用到image和mask
transformed = self.transform(image=image, mask=mask)
image = transformed['image'] # Tensor (3, H, W)
mask = transformed['mask'] # 可能是ndarray或Tensor
# 3. 将mask转换为Tensor(检查类型)
if isinstance(mask, np.ndarray):
mask = torch.from_numpy(mask).long()
else:
mask = mask.long() # 已经是Tensor,直接转换类型
# 4. 检查mask的有效性
# 确保mask的值在[0, num_classes-1]范围内
if mask.max() >= self.num_classes:
print(f"警告: {mask_path.name} 包含无效类别ID {mask.max()}")
mask = torch.clamp(mask, 0, self.num_classes - 1)
return image, mask
def get_sample_info(self, idx: int) -> dict:
"""
获取样本的元信息(用于调试和可视化)
Args:
idx: 样本索引
Returns:
info: 包含文件名、路径等信息的字典
"""
img_path, mask_path = self.valid_samples[idx]
return {
'image_name': img_path.name,
'mask_name': mask_path.name,
'image_path': str(img_path),
'mask_path': str(mask_path),
}
def get_dataloaders(
data_root: str,
batch_size: int = 4,
image_size: int = 512,
num_workers: int = 0,
num_classes: int = 2
) -> Tuple[torch.utils.data.DataLoader, torch.utils.data.DataLoader]:
"""
创建训练集和验证集的DataLoader
Args:
data_root: 数据集根目录
batch_size: 批大小
image_size: 图像尺寸
num_workers: 数据加载进程数(CPU训练时设为0)
num_classes: 类别数
Returns:
train_loader: 训练集DataLoader
val_loader: 验证集DataLoader
"""
# 创建训练集
train_dataset = FoodSegmentationDataset(
data_root=data_root,
split='train',
image_size=image_size,
num_classes=num_classes,
augmentation=True # 训练集使用数据增强
)
# 创建验证集
val_dataset = FoodSegmentationDataset(
data_root=data_root,
split='val',
image_size=image_size,
num_classes=num_classes,
augmentation=False # 验证集不使用数据增强
)
# 创建DataLoader
train_loader = torch.utils.data.DataLoader(
train_dataset,
batch_size=batch_size,
shuffle=True, # 训练集打乱顺序
num_workers=num_workers,
pin_memory=False, # CPU训练时设为False
drop_last=True if len(train_dataset) > batch_size else False
)
val_loader = torch.utils.data.DataLoader(
val_dataset,
batch_size=batch_size,
shuffle=False, # 验证集不打乱
num_workers=num_workers,
pin_memory=False,
drop_last=False
)
print(f"\n✓ DataLoader创建完成")
print(f" 训练集: {len(train_dataset)} 样本, {len(train_loader)} 批次")
print(f" 验证集: {len(val_dataset)} 样本, {len(val_loader)} 批次")
return train_loader, val_loader
def test_dataset():
"""
测试数据集加载是否正常
用于开发调试,验证:
1. 数据能否正确加载
2. 数据增强是否正常工作
3. 数据的shape和类型是否正确
"""
import matplotlib.pyplot as plt
print("="*60)
print("数据集测试")
print("="*60)
# 配置
DATA_ROOT = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/segformer_format"
# 创建数据集
dataset = FoodSegmentationDataset(
data_root=DATA_ROOT,
split='train',
image_size=512,
augmentation=True
)
# 测试读取第一个样本
image, mask = dataset[0]
info = dataset.get_sample_info(0)
print(f"\n样本信息:")
print(f" 文件名: {info['image_name']}")
print(f" 图像shape: {image.shape}") # 应该是 (3, 512, 512)
print(f" Mask shape: {mask.shape}") # 应该是 (512, 512)
print(f" Mask唯一值: {torch.unique(mask).numpy()}") # 应该是 [0, 1]
print(f" 图像数据范围: [{image.min():.3f}, {image.max():.3f}]")
# 可视化前3个样本(含数据增强效果)
fig, axes = plt.subplots(3, 3, figsize=(12, 12))
for i in range(3):
image, mask = dataset[i]
# 反标准化用于显示
mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)
image_denorm = image * std + mean
image_denorm = torch.clamp(image_denorm, 0, 1)
# 转换为numpy用于显示
image_np = image_denorm.permute(1, 2, 0).numpy()
mask_np = mask.numpy()
# 显示图像
axes[i, 0].imshow(image_np)
axes[i, 0].set_title(f"样本{i+1}: 图像")
axes[i, 0].axis('off')
# 显示mask
axes[i, 1].imshow(mask_np, cmap='gray', vmin=0, vmax=1)
axes[i, 1].set_title(f"样本{i+1}: Mask")
axes[i, 1].axis('off')
# 显示叠加
overlay = image_np.copy()
red_mask = np.zeros_like(overlay)
red_mask[mask_np == 1] = [1, 0, 0]
overlay = overlay * 0.6 + red_mask * 0.4
axes[i, 2].imshow(overlay)
axes[i, 2].set_title(f"样本{i+1}: 叠加")
axes[i, 2].axis('off')
plt.tight_layout()
plt.savefig("dataset_test_result.png", dpi=120)
print(f"\n✓ 测试结果已保存: dataset_test_result.png")
plt.show()
print("\n" + "="*60)
print("✓ 数据集测试通过!")
print("="*60)
if __name__ == "__main__":
test_dataset()