-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdataloader.py
More file actions
194 lines (160 loc) · 6.87 KB
/
Copy pathdataloader.py
File metadata and controls
194 lines (160 loc) · 6.87 KB
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
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
import os
import pickle
import sys
import time
import random
from PIL import Image
from torch.utils.data import random_split, Dataset, DataLoader
import torch
import torchvision.transforms as transforms
from gensim.models import Word2Vec
from tokenizer import Tokenizer
class PKLDataset(Dataset):
def __init__(self, img_folder = None, all_data = None, transform = None): # 两种构造方式
"""
Args:
pkl_file (str): 存储 .pkl 文件的目录路径。
transform (callable, optional): 用于处理数据的转换函数。
"""
self.img_folder = img_folder
self.transform = transform
self.all_data = all_data
def __len__(self):
return len(self.all_data)
def __getitem__(self, idx):
"""
Args:
idx (int): 数据索引。
Returns:
dict: 包含图像和标签的数据。
"""
img_id = self.all_data[idx]['ID']
img_path = os.path.join(self.img_folder, f'{img_id}.png')
img = Image.open(img_path).convert('RGB')
# Apply transformations if provided
if self.transform is not None:
img = self.transform(img)
label = self.all_data[idx]['label']
# 给 label 加标记
label = '<start>' + ' ' + label + ' ' + '<end>'
label_length = len(label.split())
return img, label, label_length
# def __load_pkl(self):
# pkl_files = os.listdir(self.pkl_file)
# all_data = []
# T0 = time.perf_counter()
# for file_name in pkl_files:
# batch_file = os.path.join(self.pkl_file, file_name)
# with open(batch_file, 'rb') as f:
# data = pickle.load(f)
# all_data.extend(data)
# T1 = time.perf_counter()
# print(f"读取时间 {T1 - T0} s")
# return all_data
def build_vocab(self):
captions = ['<start> ' + data['label'] + ' <end>' for data in self.all_data]
tokenized_captions = [caption.split() for caption in captions]
self.max_length = max([len(caption) for caption in tokenized_captions])
vocab = Word2Vec(tokenized_captions, vector_size = 100, window = 5, min_count = 1, workers = 4, seed = 42)
return vocab
# only image
class TestDataset(Dataset):
def __init__(self, img_folder, ids, transform = None):
self.img_folder = img_folder
self.ids = ids
self.transform = transform
def __len__(self):
return len(self.ids)
def __getitem__(self, index):
img_id = self.ids[index]
img_path = os.path.join(self.img_folder, f"{img_id}.png")
img = Image.open(img_path).convert('RGB')
if self.transform is not None:
img = self.transform(img)
return img_id, img
# def create_dataloader(directory, batch_size = 32, shuffle = True, num_workers = 0, transform = None):
# """
# Args:
# directory (str): 数据集目录路径。
# batch_size (int): 批量大小。
# shuffle (bool): 是否随机打乱数据。
# num_workers (int): 使用的子进程数。
# transform (callable, optional): 数据预处理函数。
# Returns:
# DataLoader: 数据加载器。
# """
# dataset = PKLDataset(directory, transform = transform)
# dataloader = DataLoader(dataset, batch_size = batch_size, shuffle = shuffle, num_workers = num_workers)
# return dataloader
def split_list(data, train_ratio = 0.9, seed = None):
if seed is not None:
random.seed(seed)
shuffled_data = data.copy()
random.shuffle(shuffled_data)
train_size = int(len(shuffled_data) * train_ratio)
train_data = shuffled_data[:train_size]
valid_data = shuffled_data[train_size:]
return train_data, valid_data
def cal_memory(all_data):
total_memory = 0
for data in all_data:
total_memory += sys.getsizeof(data['ID'])
total_memory += sys.getsizeof(data['label'])
total_memory += sys.getsizeof(data['image'])
print(f"所占字节大小:{total_memory} B")
print(f"{total_memory / 1024 / 1024} MB")
def load_pkl(pkl_dir):
all_data = []
pkl_files = os.listdir(pkl_dir)
T0 = time.perf_counter()
for file_name in pkl_files:
batch_file = os.path.join(pkl_dir, file_name)
with open(batch_file, 'rb') as f:
data = pickle.load(f)
all_data.extend(data)
T1 = time.perf_counter()
print(f"读取时间 {T1 - T0} s")
# print(f"所占字节大小:{sys.getsizeof(all_data)} B")
cal_memory(all_data)
return all_data
def create_dataloader(img_dir, pkl_dir, train_ratio = 0.9, train_batch_size = 32, valid_batch_size = 1000,
transform = None, seed = None):
# todo config
all_data = load_pkl(pkl_dir)
# dataset = PKLDataset(img_folder = img_dir, pkl_file = pkl_dir, transform = transform)
trainlist, validlist = split_list(all_data, train_ratio = train_ratio, seed = seed)
trainset = PKLDataset(img_folder = img_dir, transform = transform, all_data = trainlist)
validset = PKLDataset(img_folder = img_dir, transform = transform, all_data = validlist)
vocab = trainset.build_vocab()
# 创建数据加载器
train_loader = DataLoader(trainset, batch_size = train_batch_size, shuffle = True)
valid_loader = DataLoader(validset, batch_size = valid_batch_size, shuffle = False)
return train_loader, valid_loader, vocab
def create_testloader(img_dir, ids, test_batch_size = 1000, transform = None):
testset = TestDataset(img_dir, ids, transform)
test_loader = DataLoader(testset, batch_size = test_batch_size, shuffle = False)
return test_loader
if __name__ == '__main__':
img_folder = "datasets/train/images"
pkl_dir = "datasets/train/pkls"
test_img_folder = 'datasets/test/images'
transform = transforms.Compose([
transforms.Resize((40, 240)), # Resize image to a fixed size
transforms.ToTensor(), # Convert image to tensor
transforms.Normalize(mean = [0.485, 0.456, 0.406], std = [0.229, 0.224, 0.225]) # Normalize image
])
# # train_loader, valid_loader, vocab = create_dataloader(img_folder, pkl_dir, transform = transform)
# # tokenizer = Tokenizer(vocab)
# test_id = [104011, 105060]
# ids = [i for i in range(test_id[0], test_id[1] + 1)] # 0 ~ 1049
# testset = TestDataset(test_img_folder, ids, transform)
# (id, img) = testset[1049]
# print(id)
# print(img.shape)
dataloader, valid_dataloader, vocab = create_dataloader(img_dir = img_folder,
pkl_dir = pkl_dir,
valid_batch_size = 1000,
transform = transform
)
# for imgs, refer, length in valid_dataloader:
# print(refer)