Pytorch operation for moving vectors between torch tensors

Viewed 66

Assuming we have the torch tensors:

A: with shape BxHxW and values in {0,1}, where 0 and 1 are classes
B: with shape Bx2xD and real values, where D is the dimensionality of our vector

We want to create a new tensor of shape BxDxHxW that holds in each index specified in the spatial dimension (HxW), the vector that corresponds to its class (specified by A).

Is there a function in pytorch that implements that? I tried torch scatter but think this is not the case.

1 Answers

You are actually looking for the reverse operation, namely gathering values from one tensor using indices contained in another. Here is a canonical answer to deal with this kind of indexing scenario and apply torch.gather without much trouble.

Let's set up a minimal example with dummy data:

>>> b = 2; d = 3; h = 2; w = 1
>>> A = torch.randint(0, 2, (b,h,w)) # bhw
>>> B = torch.rand(b,2,d) # b2d
  1. Define the indexing rule you want to perform according to your problem, here:

    # out[b, d, h, w] = B[b, A[b, h, w]]
    
  2. We are looking for some kind of indexing of B's 2nd dimension using the values in A. When applying torch.gather all three tensors (input, indexer, and output) must have the same number of dimensions and the same dimension sizes except for the dimension that is being indexed, i.e. here dim=1. Observing our case, we have to stick to this pattern:

    # out[b, 1, d, h, w] = B[b, A[b, 1, d, h, w], d, h, w]
    
  3. So in order to account for this change, we need to unsqueeze/expand additional dimensions on both our input and index tensors. Therefore, to stick with the above shapes, we can do:

    First, we unsqueeze two dimension on A:

    >>> A_ = A[:,None,None].expand(-1,1,d,-1,-1)
    

    Second, we unsqueeze two dimensions on B:

    >>> B_ = B[..., None, None].expand(-1,-1,-1,h,w)
    

    Note, that expanding a dimension does not perform a copy. It is merely a view on the tensor's underlying data. At this step A_ ends up having a shape of (b, 1, d, h, w), while B_ has a shape of (b, 2, d, h, w).

  4. Now, we can simply apply torch.gather on dim=1 using A_ and B_:

    >>> out = B_.gather(dim=1, index=A_)
    

    We had to use a singleton dimension for dim=1, so we can squeeze it on the resulting tensor. This is your desired result shaped (b, d, h, w):

    >>> out[:,0]
    
Related