-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathtrain_test_loader.py
More file actions
46 lines (24 loc) · 1.1 KB
/
Copy pathtrain_test_loader.py
File metadata and controls
46 lines (24 loc) · 1.1 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
# -*- coding: utf-8 -*-
"""train_test_loader.ipynb
Automatically generated by Colaboratory.
Original file is located at
https://colab.research.google.com/drive/1zrE_6IEr176UrjuNdjU7-RjS15VSTinq
"""
import torch
import torchvision
def load(trainset,testset,seed=1,batch_size=128,num_workers=4,pin_memory=True):
#Get the Train and Test Set
# trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=train_transform)
# testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=test_transform)
SEED = 1
# CUDA?
cuda = torch.cuda.is_available()
# For reproducibility
torch.manual_seed(SEED)
if cuda:
torch.cuda.manual_seed(SEED)
# dataloader arguments - something you'll fetch these from cmdprmt
dataloader_args = dict(shuffle=True, batch_size=batch_size, num_workers=num_workers, pin_memory=pin_memory) if cuda else dict(shuffle=True, batch_size=64)
trainloader = torch.utils.data.DataLoader(trainset, **dataloader_args)
testloader = torch.utils.data.DataLoader(testset, **dataloader_args)
return trainloader, testloader