创建DataSet

1. dataset组织形式:

train, val, test为二级目录,之下是以label为文件夹名的目录。常见于分类数据集

代表:imageNet数据集

 from torch.utils.data import Dataset
    from PIL import Image
    import os
    
    class MyData(Dataset):
        def __init__(self,root_dir,label_dir):  # root_dir:dataset/train  label_dir:ants
            self.root_dir = root_dir
            self.label_dir = label_dir
            self.path = os.path.join(self.root_dir,self.label_dir)  # 用//拼接两个路径
            self.img_list = os.listdir(self.path)  # 将self.path目录下的图片转化为列表,目的是给每张图片添加索引(会按照文件名排序再添加索引)
    
        def __getitem__(self, idx):
            img_name = self.img_list[idx]  # 获取到的是文件名,不是文件
            img_path=os.path.join(self.path,img_name) # 获取文件路径
            img=Image.open(img_path)
            label=self.label_dir
            return img,label
    
        def __len__(self):
            return len(self.img_list)
    
    ants_data=MyData(root_dir="dataset/train",label_dir="ants")
    img,label=ants_data[0]
    img.show()
    
    bees_data=MyData(root_dir="dataset/train",label_dir="bees")
    
    data=ants_data+bees_data  # 拼接,后一个数据集索引从len(ants_data)开始,使用还是data[i]返回img和label

img_list的表现形式:

ant_data的表现形式:

data的表现形式(注意标黄可以看到第二个数据集的索引从124开始了):

2. dataset的组织形式:

train, val, test为二级目录,images和labels为三级目录,里面存放图片和txt文件,匹配的图片和txt文件名相同。

代表:TACO数据集

class MyData(Dataset):
    def __init__(self,root_dir, split="train"):
        self.root_dir = root_dir   # dataset
        self.image_dir=os.path.join(root_dir,split,"images")
        self.label_dir=os.path.join(root_dir,split,"labels")
        self.img_list = os.listdir(self.image_dir)
        self.label_list=os.listdir(self.label_dir)

    def __getitem__(self, idx):
        img_name=self.img_list[idx]
        img_path = os.path.join(self.image_dir,img_name)
        img=Image.open(img_path)

        label_path=os.path.join(self.label_dir,img_name.replace(".jpg",".txt"))
        file=open(label_path)
        label=file.read()
        file.close()

        return img, label

    def __len__(self):
        return len(self.img_list)

data=MyData(root_dir="dataset",split="train")
img,label=data[0]
print(label)
img.show()

data的表现形式:

img_list:

label_list:

其他:

创建DataLoader

test_loader=DataLoader(dataset=test_set,batch_size=64,shuffle=True,drop_last=False)

常见的参数

1. dataset:要从哪个数据集采样

2. batch_size:一次采样多少数据

3. shuffle:一轮采样是否打乱顺序

4. drop_last:最后一个batch样本数不够batch_size时是否还要采样

5. sampler:每轮样本的采样规则。默认的有RandomSampler和SequentialSampler

注意

1. shuffle和sample不能同时生效。DataLoader的采样顺序完全由sampler决定。

如果显示传入sampler,shuffle被忽略;没有传入sampler,shuffle=True会让DataLoader自动创建RandomSampler;shuffle=False 会创建 SequentialSampler。

可以看到在dataloader内部是没有shuffle这个变量的,有的只是sampler。

shuffle=True则每轮都随机采样,因此可以看到轮之间每次采样的样本大概率是不同的。

 

shuffle=False则每轮都顺序采样,因此可以看到轮之间每次采样的样本都是相同的,和dataset的样本顺序相同。

2. sampler的其他类型

子集采样。多用于从训练数据集中划分出验证集。

indices就是你希望 DataLoader 从 dataset 中采样的样本下标列表。

这里指定了indices只能从第0,2,

5,7,10个样本中选择,可以发现第0,2,5,7,10个样本的标签分别是3 8 6 6 0

from torch.utils.data import SubsetRandomSampler

indices = [0, 2, 5, 7, 10]   # 3 8 6 6 0
sampler = SubsetRandomSampler(indices)
test_loader=DataLoader(dataset=test_set,batch_size=64,drop_last=False,sampler=sampler)

for epoch in range(3):
    for data in test_loader:
        imgs, labels=data
        print(labels)

最终输出的label:

按权重采样。多用于类别不平衡的情况。

data = list(range(9))
labels = torch.tensor([0, 0, 0, 1, 1, 1, 1, 2, 2]) 

class_count = torch.bincount(labels)
class_weight=1.0/class_count.float()
sample_weights = class_weight[labels]

sampler = WeightedRandomSampler(
    weights=sample_weights,
    num_samples=10,     # 每个 epoch 采 10 次
    replacement=True    # 允许重复采样
)

loader = DataLoader(
    data,
    batch_size=5,
    sampler=sampler
)

for epoch in range(3):
    print(f"\nEpoch {epoch}")
    for batch in loader:
        print(batch)

sampler的num_samples设置一个epoch的采样数,dataloader的batch_size设置每个batch装多少个样本。每轮就会迭代num_samples/batch_size次。因此下面每个epoch输出是两行,表示两次迭代。如果没有sampler,那么每轮就迭代len(dataset) / batch_size次。

输出:

epoch 0每个类别输出的次数分别是2 6 2 ,epoch1:3 3 4,epoch2:5 3 2。在权重差异比较小的情况下,即使按权重比例抽样,但仍然是随机的。比如第一个epoch仍然是类别1数目更多。

数据:

注:

本人也在学习中,此帖用作个人记录,如有错误或不全,欢迎指出,共同进步!

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐