345 lines
11 KiB
Python
345 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 (H, W)
|
|
|
|
# 3. 将mask转换为Tensor
|
|
mask = torch.from_numpy(mask).long()
|
|
|
|
# 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()
|