From dfaeb9f8d575515933b52943751153ec2761d831 Mon Sep 17 00:00:00 2001 From: troyyyyy <49091847+troyyyyy@users.noreply.github.com> Date: Fri, 18 Nov 2022 15:38:41 +0800 Subject: [PATCH] Update get_mnist_add.py --- datasets/mnist_add/get_mnist_add.py | 26 +++++++++++++++++++------- 1 file changed, 19 insertions(+), 7 deletions(-) diff --git a/datasets/mnist_add/get_mnist_add.py b/datasets/mnist_add/get_mnist_add.py index fcada50..1af834a 100644 --- a/datasets/mnist_add/get_mnist_add.py +++ b/datasets/mnist_add/get_mnist_add.py @@ -3,27 +3,39 @@ import torchvision from torch.utils.data import Dataset from torchvision.transforms import transforms -def get_data(file, img_dataset): +def get_data(file, img_dataset, get_pseudo_label): X = [] + if(get_pseudo_label): + Z = [] 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]]) + if(get_pseudo_label): + Z.append([img_dataset[int(line[0])][1], img_dataset[int(line[1])][1]]) Y.append(int(line[2])) - return X, Y + + if(get_pseudo_label): + return X, Z, Y + else: + return X, None, Y -def get_mnist_add(): +def get_mnist_add(train = True, get_pseudo_label = False): transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081, ))]) img_dataset = torchvision.datasets.MNIST(root='./datasets/mnist_add/', train=True, download=True, transform=transform) - train_X, train_Y = get_data('./datasets/mnist_add/train_data.txt', img_dataset) - test_X, test_Y = get_data('./datasets/mnist_add/test_data.txt', img_dataset) + if(train): + file = './datasets/mnist_add/train_data.txt' + else: + file = './datasets/mnist_add/test_data.txt' + + return get_data(file, img_dataset, get_pseudo_label) - return train_X, train_Y, test_X, test_Y if __name__ == "__main__": - train_X, train_Y, test_X, test_Y = get_mnist_add() + train_X, train_Y = get_mnist_add(train = True) + test_X, test_Y = get_mnist_add(train = False) print(len(train_X), len(test_X)) print(train_X[0][0].shape, train_X[0][1].shape, train_Y[0])