增加图像分割相应的代码
This commit is contained in:
@@ -0,0 +1,344 @@
|
||||
"""
|
||||
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()
|
||||
Reference in New Issue
Block a user