Premium problem107. Write a Dataset

Medium Locked

Write a torch.utils.data.Dataset subclass that stores features and labels and implements __len__ and __getitem__, where __getitem__ returns the (feature, label) pair at a position.

Do not use the built-in TensorDataset. Return the tuple (length, feature_at_index, label_at_index) for the given index.

Those two methods are the entire protocol: anything implementing them can be handed to a DataLoader.

Input

features =
tensor([[ 0.,  1.,  2.],
        [ 3.,  4.,  5.],
        [ 6.,  7.,  8.],
        [ 9., 10., 11.]])
labels = tensor([0, 1, 0, 1])
index = 2

Output

(4, tensor([6., 7., 8.]), tensor(0))

Premium problem

This one's part of Premium. Unlock the full PyTorch track plus every other premium problem on the site.

Implement solve(...)