MCPcopy Create free account
hub / github.com/Mew233/pairwise / dataloader

Function dataloader

pairwise/dataloader.py:28–68  ·  view source on GitHub ↗

should be starting from X, and Y X_{}_trainval, X_{}_test, Y_{}_trainval, Y_{}_test

(*args,**kwargs)

Source from the content-addressed store, hash-verified

26
27# def dataloader(X_train_val_set, X_test_set, Y_train_val_set, Y_test_set):
28def dataloader(*args,**kwargs):
29 """
30 should be starting from X, and Y
31 X_{}_trainval, X_{}_test, Y_{}_trainval, Y_{}_test
32 """
33 temp_loader = {}
34 for name, input in kwargs.items():
35 # First, format
36 input_name = name.split('_')
37 if input_name[0].startswith('Y'):
38 input = input.astype('float32')
39 input = torch.from_numpy(input)
40 input = input.unsqueeze(1)
41 temp_loader[name] = input
42
43 elif input_name[0].startswith('X'):
44 #input is tabular format
45 input = input.astype('float32')
46 input = torch.from_numpy(input)
47 temp_loader[name] = input
48
49 elif input_name[0].startswith('dummy'):
50 input = input.astype('float32')
51 input = torch.from_numpy(input)
52 temp_loader[name] = input
53
54 # Second, to tensordataset
55 temp_loader_trainval, temp_loader_test = [], []
56 for key, val in temp_loader.items():
57 load_name = key.split('_')
58 # train_val_dataset
59 if load_name[-1].endswith('trainval'):
60 temp_loader_trainval.append(val)
61 elif load_name[-1].endswith('test'):
62 temp_loader_test.append(val)
63
64 train_val_dataset = torch.utils.data.TensorDataset(*temp_loader_trainval)
65 test_dataset = torch.utils.data.TensorDataset(*temp_loader_test)
66 test_loader = DataLoader(test_dataset, batch_size=256,shuffle = False, sampler=sampler.SequentialSampler(test_dataset))
67 #list(BatchSampler(SequentialSampler(range(10)), batch_size=3, drop_last=False))
68 return train_val_dataset, test_loader
69
70
71def dataloader_graph(*args,**kwargs):

Callers 1

trainingFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected