split / utils /dataloader.py
wangwenguang's picture
Upload 23 files
5ea0d01 verified
Raw
History Blame
7.47 kB
import torch
from utils.utils import cvtColor, preprocess_input
import os
from PIL import Image
import numpy as np
from torch.utils.data import Dataset, DataLoader
import cv2
class UnetDataset(Dataset):
def __init__(self, data_path, input_shape, num_classes, augmentation=True ,txt_name: str = "train.txt"):
# 读取train.txt和val.txt,test.txt文件,获取训练和验证集的图像ID
with open(os.path.join(data_path, "VOC2012/ImageSets/Segmentation", txt_name), "r") as f:
self.annotation_lines = f.readlines()
# 初始化其他参数
self.length = len(self.annotation_lines) # 数据集的长度
self.input_shape = input_shape # 输入图像的形状(宽和高)
self.num_classes = num_classes # 类别数目
self.augmentation = augmentation # 是否在训练阶段,用来控制是否使用数据增强
self.data_path = data_path # 数据集路径
def __len__(self):
# 返回数据集的大小
return self.length
def __getitem__(self, index):
# 读取单个样本
annotation_line = self.annotation_lines[index] # 获取对应的annotation
name = annotation_line.split()[0] # 获取文件名,通常是图像文件的名称
# 读取JPEG图像
jpg = Image.open(os.path.join(self.data_path, "VOC2012/JPEGImages", name + ".png"))
# 读取PNG标签图像
png = Image.open(os.path.join(self.data_path, "VOC2012/SegmentationClass", name + ".png"))
# 如果是训练阶段,进行随机数据增强
jpg, png = self.get_random_data(jpg, png, self.input_shape, random=self.augmentation)
# 图像预处理
jpg = np.transpose(preprocess_input(np.array(jpg, np.float64)), [2, 0, 1])
# 标签转换成numpy数组
png = np.array(png)
# 将标签值大于类别数的部分设置为类别数(忽略这些区域)
png[png >= self.num_classes] = self.num_classes
# 将标签转换为one-hot编码
seg_labels = np.eye(self.num_classes + 1)[png.reshape([-1])]
# 重塑标签为目标形状
seg_labels = seg_labels.reshape((int(self.input_shape[0]), int(self.input_shape[1]), self.num_classes + 1))
# 返回图像、标签以及one-hot编码的标签
return jpg, png, seg_labels
# 生成随机数的函数
def rand(self, a=0, b=1):
return np.random.rand() * (b - a) + a
# 对图像和标签进行随机数据增强的函数
def get_random_data(self, image, label, input_shape, jitter=.3, hue=.1, sat=0.7, val=0.3, random=True):
# 将图像转为RGB格式
image = cvtColor(image)
label = Image.fromarray(np.array(label))
iw, ih = image.size # 获取图像的宽和高
h, w = input_shape # 获取目标图像的高和宽
if not random:
# 如果不进行随机增强(例如在验证阶段)
iw, ih = image.size
scale = min(w / iw, h / ih) # 计算缩放比例
nw = int(iw * scale) # 根据比例计算缩放后的宽
nh = int(ih * scale) # 根据比例计算缩放后的高
# 缩放图像并进行中心裁剪
image = image.resize((nw, nh), Image.BICUBIC)
new_image = Image.new('RGB', [w, h], (128, 128, 128)) # 创建一个灰色背景的图像
new_image.paste(image, ((w - nw) // 2, (h - nh) // 2)) # 将缩放后的图像粘贴到目标图像中
# 缩放标签图像并进行中心裁剪
label = label.resize((nw, nh), Image.NEAREST) # 标签使用最近邻插值
new_label = Image.new('L', [w, h], (0)) # 创建一个空白标签图像
new_label.paste(label, ((w - nw) // 2, (h - nh) // 2)) # 将缩放后的标签图像粘贴到目标图像中
return new_image, new_label
# 获取一个新的宽高比(通过调整宽和高的比例)
new_ar = iw / ih * self.rand(1 - jitter, 1 + jitter) / self.rand(1 - jitter, 1 + jitter)
scale = self.rand(0.25, 2) # 随机缩放比例
if new_ar < 1:
nh = int(scale * h)
nw = int(nh * new_ar)
else:
nw = int(scale * w)
nh = int(nw / new_ar)
image = image.resize((nw, nh), Image.BICUBIC)
label = label.resize((nw, nh), Image.NEAREST)
# 随机翻转图像
flip = self.rand() < .5
if flip:
image = image.transpose(Image.FLIP_LEFT_RIGHT)
label = label.transpose(Image.FLIP_LEFT_RIGHT)
# 在图像周围随机添加灰色边框
dx = int(self.rand(0, w - nw))
dy = int(self.rand(0, h - nh))
new_image = Image.new('RGB', (w, h), (128, 128, 128))
new_label = Image.new('L', (w, h), (0))
new_image.paste(image, (dx, dy)) # 将图像粘贴到新图像中
new_label.paste(label, (dx, dy)) # 将标签粘贴到新标签中
image = new_image
label = new_label
# 转换图像为数组
image_data = np.array(image, np.uint8)
r = np.random.uniform(-1, 1, 3) * [hue, sat, val] + 1 # 随机调整色调、饱和度和亮度
hue, sat, val = cv2.split(cv2.cvtColor(image_data, cv2.COLOR_RGB2HSV)) # 转为HSV色域
dtype = image_data.dtype # 获取数据类型
x = np.arange(0, 256, dtype=r.dtype) # 获取颜色值范围
lut_hue = ((x * r[0]) % 180).astype(dtype) # 应用色调变换
lut_sat = np.clip(x * r[1], 0, 255).astype(dtype) # 应用饱和度变换
lut_val = np.clip(x * r[2], 0, 255).astype(dtype) # 应用亮度变换
# 使用查找表(LUT)应用变换
image_data = cv2.merge((cv2.LUT(hue, lut_hue), cv2.LUT(sat, lut_sat), cv2.LUT(val, lut_val)))
image_data = cv2.cvtColor(image_data, cv2.COLOR_HSV2RGB) # 转换回RGB色域
return image_data, label # 返回经过增强的图像和标签
# DataLoader中collate_fn使用
def unet_dataset_collate(batch):
# 初始化三个列表,用于存储每个批次中的图像、标签和one-hot编码标签
images = [] # 用来存储图像数据
pngs = [] # 用来存储原始标签(通常是类别标签)
seg_labels = [] # 用来存储one-hot编码的标签
# 遍历当前批次中的每个样本(img, png, labels)
for img, png, labels in batch:
images.append(img) # 将图像添加到images列表中
pngs.append(png) # 将原始标签添加到pngs列表中
seg_labels.append(labels) # 将one-hot标签添加到seg_labels列表中
# 将列表转换为NumPy数组,然后转换为torch张量
# images的张量需要是float类型,通常用于输入图像
images = torch.from_numpy(np.array(images)).type(torch.FloatTensor)
# pngs的张量需要是long类型,通常用于标签索引
pngs = torch.from_numpy(np.array(pngs)).long()
# seg_labels的张量需要是float类型,通常用于标签的one-hot编码
seg_labels = torch.from_numpy(np.array(seg_labels)).type(torch.FloatTensor)
# 返回三个张量,分别对应图像、原始标签和one-hot标签
return images, pngs, seg_labels