用于DataLoader的pytorch数据集
暂时介绍 image-mask型数据集, 以人手分割数据集 EGTEA Gaze+ 为例.
准备数据文件夹
需要将Image和Mask分开存放, 对应文件的文件名必须保持一致. 提醒: Mask 图像一般为 png 单通道
EGTEA Gaze+ 数据集下载解压后即得到如下的目录, 无需处理
hand14k
┣━ Images
┃ ┣━ OP01-R01-PastaSalad_000014.jpg
┃ ┣━ OP01-R01-PastaSalad_000015.jpg
┃ ┣━ OP01-R01-PastaSalad_000016.jpg
┃ ┗━ ···
┗━ Masks
┣━ OP01-R01-PastaSalad_000014.png
┣━ OP01-R01-PastaSalad_000015.png
┣━ OP01-R01-PastaSalad_000016.png
┗━ ···
生成路径文件, 划分数据集
脚本如下:import cv2 as cv
import numpy as np
import PIL.Image as Image
import os
np.random.seed(42)
def split_dataset():
# 读取图像文件
images_path = "./Images/"
images_list = os.listdir(images_path) # 每次返回文件列表顺序不一致
images_list.sort() # 需要排序处理
# 读取标签/Mask图像
labels_path = "./Masks/"
labels_list = os.listdir(labels_path)
labels_list.sort()
# 创建路径文件 (使用二进制编码, 避免操作系统不匹配)
train_file = "./train.data"
test_file = "./test.data"
if os.path.isfile(train_file) and os.path.isfile(test_file):
return
train_file = open(train_file, "wb")
test_file = open(test_file, "wb")
# 外汇返佣
split_ratio = 0.8
for image, label in zip(images_list, labels_list):
image = os.path.join(images_path, image)
label = os.path.join(labels_path, label)
if os.path.basename(image).split('.')[0] != os.path.basename(label).split('.')[0]:
continue
file = train_file if np.random.rand() < split_ratio else test_file
file.write((image + "\t" + label + "\n").encode("utf-8"))
train_file.close()
test_file.close()
print("成功划分数据集!")
def read_image(path):
img = np.array(Image.open(path))
if img.ndim == 2:
img = cv.merge([img, img, img])
return img
def test_read():
train_file = "./test.data"
with open(train_file, 'rb') as f:
datalist = f.readlines()
datalist = [(k, v) for k, v in map(lambda x: x.decode('utf-8').strip('\n').split('\t'), datalist)]
item = datalist[np.random.randint(42)]
image = read_image(item[0])
mask = read_image(item[1])
cv.imshow("image", image)
cv.imshow("mask", mask)
cv.waitKey(0)
cv.destroyAllWindows()
if __name__ == '__main__':
split_dataset()
test_read()
派生 Dataset 类
class MyDataset(Dataset):
def __init__(
self, data_file, data_dir, transform_trn=None, transform_val=None
):
"""
Args:
data_file (string): Path to the data file with annotations.
data_dir (string): Directory with all the images.
transform_{trn, val} (callable, optional): Optional transform to be applied
on a sample.
"""
with open(data_file, 'rb') as f:
datalist = f.readlines()
self.datalist = [(k, v) for k, v in map(lambda x: x.decode('utf-8').strip('\n').split('\t'), datalist)]
self.root_dir = data_dir
self.transform_trn = transform_trn
self.transform_val = transform_val
self.stage = 'train'
def set_stage(self, stage):
self.stage = stage
def __len__(self):
return len(self.datalist)
def __getitem__(self, idx):
img_name = os.path.join(self.root_dir, self.datalist[idx][0])
msk_name = os.path.join(self.root_dir, self.datalist[idx][1])
def read_image(x):
img_arr = np.array(Image.open(x))
if len(img_arr.shape) == 2: # grayscale
img_arr = np.tile(img_arr, [3, 1, 1]).transpose(1, 2, 0)
return img_arr
image = read_image(img_name)
mask = np.array(Image.open(msk_name))
if img_name != msk_name:
assert len(mask.shape) == 2, 'Masks must be encoded without colourmap'
sample = {'image': image, 'mask': mask}
if self.stage == 'train':
if self.transform_trn:
sample = self.transform_trn(sample)
elif self.stage == 'val':
if self.transform_val:
sample = self.transform_val(sample)
return sample
构造DataLoader
# 定义Transform
composed_trn = transforms.Compose([ResizeShorterScale(shorter_side, low_scale, high_scale),
Pad(crop_size, [123.675, 116.28, 103.53], ignore_label),
RandomMirror(),
RandomCrop(crop_size),
Normalise(*normalise_params),
ToTensor()])
composed_val = transforms.Compose([Normalise(*normalise_params),
ToTensor()])
# 导入数据集
trainset = MyDataset(data_file=train_list,
data_dir=train_dir,
transform_trn=composed_trn,
transform_val=composed_val)
valset = MyDataset(data_file=val_list,
data_dir=val_dir,
transform_trn=None,
transform_val=composed_val)
# 构建生成器
train_loader = DataLoader(trainset,
batch_size=batch_size,
shuffle=True,
num_workers=num_workers,
pin_memory=True,
drop_last=True)
val_loader = DataLoader(valset,
batch_size=1,
shuffle=False,
num_workers=num_workers,
pin_memory=True)
训练
for i, sample in enumerate(train_loader):
image = sample['image'].cuda()
target = sample['mask'].cuda()
image_var = torch.autograd.Variable(image).float()
target_var = torch.autograd.Variable(target).long()
# Compute output
output = net(image_var)
...
原文链接:https://blog.csdn.net/Augurlee/article/details/103652444
用于DataLoader的pytorch数据集的更多相关文章
- PyTorch 数据集类 和 数据加载类 的一些尝试
最近在学习PyTorch, 但是对里面的数据类和数据加载类比较迷糊,可能是封装的太好大部分情况下是不需要有什么自己的操作的,不过偶然遇到一些自己导入的数据时就会遇到一些问题,因此自己对此做了一些小实 ...
- Pytorch数据集读取
Pytorch中数据集读取 在机器学习中,有很多形式的数据,我们就以最常用的几种来看: 在Pytorch中,他自带了很多数据集,比如MNIST.CIFAR10等,这些自带的数据集获得和读取十分简便: ...
- Pytorch数据集读入——Dataset类,实现数据集打乱Shuffle
在进行相关平台的练习过程中,由于要自己导入数据集,而导入方法在市面上五花八门,各种库都可以应用,在这个过程中我准备尝试torchvision的库dataset torchvision.datasets ...
- [Pytorch数据集下载] 下载MNIST数据缓慢的方案
步骤一 首先访问下面的网站,手工下载数据集.http://yann.lecun.com/exdb/mnist/ 把四个压缩包下载到任意文件夹,以便之后使用. 步骤二 把自己电脑上已经下载好的数据集的文 ...
- PyTorch 之 DataLoader
DataLoader DataLoader 是 PyTorch 中读取数据的一个重要接口,该接口定义在 dataloader.py 文件中,该接口的目的: 将自定义的 Dataset 根据 batch ...
- 什么是pytorch(4.数据集加载和处理)(翻译)
数据集加载和处理 这里主要涉及两个包:torchvision.datasets 和torch.utils.data.Dataset 和DataLoader torchvision.datasets是一 ...
- 【pytorch】torch.utils.data.DataLoader
简介 DataLoader是PyTorch中的一种数据类型.用于训练/验证/测试时的数据按批读取. torch.utils.data.DataLoader(dataset, batch_size=1, ...
- pytorch加载语音类自定义数据集
pytorch对一下常用的公开数据集有很方便的API接口,但是当我们需要使用自己的数据集训练神经网络时,就需要自定义数据集,在pytorch中,提供了一些类,方便我们定义自己的数据集合 torch.u ...
- [实现] 利用 Seq2Seq 预测句子后续字词 (Pytorch)2
最近有个任务:利用 RNN 进行句子补全,即给定一个不完整的句子,预测其后续的字词.本文使用了 Seq2Seq 模型,输入为 5 个中文字词,输出为 1 个中文字词.目录 关于RNN 语料预处理 搭建 ...
随机推荐
- Retrofit RestAdapter 配置说明
RestAdapter.Builder builder = new RestAdapter.Builder(); builder.setEndpoint(ip地址 ...
- Chrome谷歌页面翻译增强插件开发
最近想做一个Chrome的插件(看别的博客说其实叫插件不准确,应该叫拓展,大家叫习惯了就按习惯的来吧).一开始咱先直接看了Chrome开发(360翻译)和chrome extensions(这个官方的 ...
- (转)linux nc命令使用详解
linux nc命令使用详解 原文:https://www.2cto.com/os/201306/220971.html 功能说明:功能强大的网络工具 语 法:nc [-hlnruz][-g<网 ...
- GUI_FlowLayout
void setBounds(x, y, width, height) 设置窗体坐标,窗体大小 import java.awt.Frame; public class IntegerDemo { pu ...
- Centos 下更改MySQL源数据存放目录(datadir)
MySQL在安装完成之后,其源数据默认存放在 /var/lib/mysql/ 目录下,一般情况下,该目录在根目录下,由于Linux系统默认 根目录所在挂载的磁盘容量有限,随着生产数据的不断产生,该目 ...
- Jmeter+ SeureCRT + Pinpoint
1.环境配置 [相关操作] 下载jdk http://www.oracle.com/technetwork/java/javase/downloads/jdk8-downloads-2133151.h ...
- 实验报告一&第三周学习总结
一.实验报告 1.打印输出所有的"水仙花数",所谓"水仙花数"是指一个3位数,其中各位数字立方和等于该数本身.例如,153是一个"水仙花数" ...
- 关于时间API
如何正确处理时间 现实生活的世界里,时间是不断向前的,如果向前追溯时间的起点,可能是宇宙出生时,又或是宇宙出现之前, 但肯定是我们目前无法找到的,我们不知道现在距离时间原点的精确距离.所以我们要表示时 ...
- 最全的 Java 知识总结- Github 日增 10 star
项目地址: 如果觉得有帮助,希望大家给个 star 鼓励以下:同时也希望大家多多 fork,一起加入进来. 为什么选择做这个开源项目 首先,希望提高自己:因为选择做这个,自己肯定就会花时间去提高自己的 ...
- java SSM框架 代码生成器 快速开发平台 websocket即时通讯 shiro redis
A代码编辑器,在线模版编辑,仿开发工具编辑器,pdf在线预览,文件转换编码 B 集成代码生成器 [正反双向](单表.主表.明细表.树形表,快速开发利器)+快速表单构建器 freemaker模版技术 , ...