Select specific indexes of 3D Pytorch Tensor using a 1D long tensor that represents indexes

Viewed 1772

So I have a tensor that is M x B x C, where M is the number of models, B is the batch and C is the classes and each cell is the probability of a class for a given model and batch. Then I have a tensor of the correct answers which is just a 1D of size B we'll call "t". How do I use the 1D of size B to just return a M x B x 1, where the returned tensor is just the value at the correct class? Say the M x B x C tensor is called "blah" I've tried

blah[:, :, C]

for i in range(M):
    blah[i, :, C]

blah[:, C, :]

The top 2 just return the values of indexes t in the 3rd dimension of every slice. The last one returns the values at t indexes in the 2nd dimension. How do I do this?

3 Answers

We can get the desired result by combining advanced and basic indexing

import torch

# shape [2, 3, 4]
blah = torch.tensor([
    [[ 0,  1,  2,  3],
     [ 4,  5,  6,  7],
     [ 8,  9, 10, 11]],
    [[12, 13, 14, 15],
     [16, 17, 18, 19],
     [20, 21, 22, 23]]])

# shape [3]
t = torch.tensor([2, 1, 0])
b = torch.arange(blah.shape[1]).type_as(t)

# shape [2, 3, 1]
result = blah[:, b, t].unsqueeze(-1)

which results in

>>> result
tensor([[[ 2],
         [ 5],
         [ 8]],
        [[14],
         [17],
         [20]]])

Here is one way to do it:

Suppose a is your M x B x C shaped tensor. I am taking some representative values below,

>>> M = 3
>>> B = 5
>>> C = 4
>>> a = torch.rand(M, B, C)
>>> a
tensor([[[0.6222, 0.6703, 0.0057, 0.3210],
         [0.6251, 0.3286, 0.8451, 0.5978],
         [0.0808, 0.8408, 0.3795, 0.4872],
         [0.8589, 0.8891, 0.8033, 0.8906],
         [0.5620, 0.5275, 0.4272, 0.2286]],

        [[0.2419, 0.0179, 0.2052, 0.6859],
         [0.1868, 0.7766, 0.3648, 0.9697],
         [0.6750, 0.4715, 0.9377, 0.3220],
         [0.0537, 0.1719, 0.0013, 0.0537],
         [0.2681, 0.7514, 0.6523, 0.7703]],

        [[0.5285, 0.5360, 0.7949, 0.6210],
         [0.3066, 0.1138, 0.6412, 0.4724],
         [0.3599, 0.9624, 0.0266, 0.1455],
         [0.7474, 0.2999, 0.7476, 0.2889],
         [0.1779, 0.3515, 0.8900, 0.2301]]])

Let's say the 1D class tensor is t, which gives the true class of each example in the batch. So it is a 1D tensor of shape (B, ) having class labels in the range {0, 1, 2, ..., C-1}.

>>> t = torch.randint(C, size = (B, ))
>>> t
tensor([3, 2, 1, 1, 0])

So basically you want to select the indices corresponding to t from the innermost dimension of a. This can be achieved using fancy indexing and broadcasting combined as follows:

>>> i = torch.arange(M).reshape(M, 1, 1)
>>> j = torch.arange(B).reshape(1, B, 1)
>>> k = t.reshape(1, B, 1)

Note that once you index anything by (i, j, k), they are going to expand and take the shape (M, B, 1) which is the desired output shape. Now just indexing a by i, j and k gives:

>>> a[i, j, k]
tensor([[[0.3210],
         [0.8451],
         [0.8408],
         [0.8891],
         [0.5620]],

        [[0.6859],
         [0.3648],
         [0.4715],
         [0.1719],
         [0.2681]],

        [[0.6210],
         [0.6412],
         [0.9624],
         [0.2999],
         [0.1779]]])

So essentially, if you generate the index arrays conveying your access pattern beforehand, you can directly use them to extract some slice of the tensor.

You simply need to pass:

  • your index as the third slice
  • range(B) as the second slice
    (i.e. which element in the 2nd dim each 3rd dim index corresponds to)
blah[:,range(B),t]
Related