from torch.utils.data import DataLoader dataset = SequencesDataset(sequences) def collate(batches): batch_list = [] for batch in batches: pairs = np.array(batch[0]) negs = np.array(batch[1]) negs = np.vstack((pairs[:, 0].repeat(negs.shape[1]), negs.ravel())).T pairs_arr = np.ones((pairs.shape[0], pairs.shape[1] + 1), dtype=int) pairs_arr[:, :-1] = pairs negs_arr = np.zeros((negs.shape[0], negs.shape[1] + 1), dtype=int) negs_arr[:, :-1] = negs all_arr = np.vstack((pairs_arr, negs_arr)) batch_list.append(all_arr) batch_array = np.vstack(batch_list) # Return item1, item2, label return (torch.LongTensor(batch_array[:, 0]), torch.LongTensor(batch_array[:, 1]), torch.FloatTensor(batch_array[:, 2])) # We can customize the batch size, number of works, shuffle, etc. train_dataloader = DataLoader(dataset, batch_size=128, shuffle=True, num_workers=8, collate_fn=collate)