Avoid Loop for Selecting Dimensions in a Batch Tensor in PyTorch

Viewed 462

I have a batch tensor and another tensor having indices of the dimensions to select from batch tensor. At present, I am looping around batch tensor as shown below in the code snippet:

import torch

# create tensors to represent our data in torch format
batch_size = 8
batch_data = torch.rand(batch_size, 3, 240, 320)

# notice that channels_id has 8 elements, i.e., = batch_size
channels_id = torch.tensor([2, 0, 2, 1, 0, 2, 1, 0])

This is how I am selecting dimensions inside a for loop and then stacking to convert a single tensor:

batch_out = torch.stack([batch_i[channel_i] for batch_i, channel_i in zip(batch_data, channels_id)])
batch_out.size()  # prints torch.Size([8, 240, 320])

It works fine. However, is there a better PyTorch way to achieve the same?

1 Answers

As per the hint from @Shai, I could make it work using the torch.gather function. Below is the complete code:

import torch

# create tensors to represent our data in torch format
batch_size = 8
batch_data = torch.rand(batch_size, 3, 240, 320)

# notice that channels_id has 8 elements, i.e., batch_size
channels_id = torch.tensor([2, 0, 2, 1, 0, 2, 1, 0])

# resizing channels_id to (8 , 1, 240, 320)
channels_id = channels_id.view(-1, 1, 1, 1).repeat((1, 1) + batch_data.size()[-2:])

batch_out = torch.gather(batch_data, 1, channels_id).squeeze()
batch_out.size()  # prints torch.Size([8, 240, 320])
Related