Browse Source

Update get_mnist_add.py

pull/3/head
troyyyyy GitHub 3 years ago
parent
commit
4fa2f85945
No known key found for this signature in database GPG Key ID: 4AEE18F83AFDEB23
1 changed files with 12 additions and 16 deletions
  1. +12
    -16
      datasets/mnist_add/get_mnist_add.py

+ 12
- 16
datasets/mnist_add/get_mnist_add.py View File

@@ -3,25 +3,22 @@ import torchvision
from torch.utils.data import Dataset
from torchvision.transforms import transforms

def get_data(file, img_dataset):
X = []
Y = []
with open(file) as f:
for line in f:
line = line.strip().split(' ')
X.append((img_dataset[int(line[0])][0], img_dataset[int(line[1])][0]))
Y.append(int(line[2]))
return X, Y

def get_mnist_add():
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081, ))])
img_dataset = torchvision.datasets.MNIST(root='./', train=True, download=True, transform=transform)
train_X = []
train_Y = []
with open('./train_data.txt') as f:
for line in f:
line = line.strip().split(' ')
train_X.append((img_dataset[int(line[0])][0], img_dataset[int(line[1])][0]))
train_Y.append(int(line[2]))
test_X = []
test_Y = []
with open('./test_data.txt') as f:
for line in f:
line = line.strip().split(' ')
test_X.append((img_dataset[int(line[0])][0], img_dataset[int(line[1])][0]))
test_Y.append(int(line[2]))
train_X, train_Y = get_data('./train_data.txt', img_dataset)
test_X, test_Y = get_data('./test_data.txt', img_dataset)
return train_X, train_Y, test_X, test_Y

@@ -29,4 +26,3 @@ if __name__ == "__main__":
train_X, train_Y, test_X, test_Y = get_mnist_add()
print(len(train_X), len(test_X))
print(train_X[0][0].shape, train_X[0][1].shape, train_Y[0])

Loading…
Cancel
Save