File size: 7,472 Bytes
5ea0d01 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | 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
|