MCPcopy Create free account
hub / github.com/TrustAIResearch/MLHospital / GetDataLoaderPoison

Class GetDataLoaderPoison

mlh/data_preprocessing/data_loader.py:224–333  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

222
223
224class GetDataLoaderPoison(object):
225 def __init__(self, args):
226 self.args = args
227 self.data_path = args.data_path
228 self.input_shape = args.input_shape
229
230 def parse_dataset(self, dataset):
231
232 if dataset in configs.SUPPORTED_IMAGE_DATASETS:
233 _loader = getattr(datasets, dataset)
234 if dataset != "EMNIST":
235 train_dataset = _loader(root=self.data_path,
236 train=True,
237 transform=None,
238 download=True)
239 test_dataset = _loader(root=self.data_path,
240 train=False,
241 transform=None,
242 download=True)
243 else:
244 train_dataset = _loader(root=self.data_path,
245 train=True,
246 split="byclass",
247 transform=None,
248 download=True)
249 test_dataset = _loader(root=self.data_path,
250 train=False,
251 split="byclass",
252 transform=None,
253 download=True)
254
255 else:
256 raise ValueError("Dataset Not Supported: ", dataset)
257 return train_dataset, test_dataset
258
259 def get_data_transform(self, dataset, use_transform="simple"):
260 transform_list = [transforms.Resize(
261 (self.input_shape[0], self.input_shape[0])), ]
262
263 if use_transform == "simple":
264 transform_list += [transforms.RandomCrop(
265 32, padding=4), transforms.RandomHorizontalFlip(), ]
266
267 print("add simple data augmentation!")
268
269 transform_list.append(transforms.ToTensor())
270
271 if dataset in ["MNIST", "FashionMNIST", "EMNIST"]:
272 transform_list = [
273 transforms.Grayscale(3), ] + transform_list
274
275 transform_ = transforms.Compose(transform_list)
276 return transform_
277
278 def get_data_loader(self):
279
280 train_transform = self.get_data_transform(
281 self.args.dataset, use_transform=None)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected