how to rotate images in a batch separately in pytorch

Viewed 655

I am randomly rotating 3D images in pytorch using torch.rot90 but this rotates all the images in the batch in the same way. I would like to find a differentiable way to randomly rotate each image in a different axis. here is the code which rotates each image to the same orientation:

#x = next batch

k = torch.randint(0, 4, (1,)).item()
dims = [0,0]
dims[0] = dims[1] = torch.randint(2, 5, (1,))
while dims[0] == dims[1]:#this makes sure the two axes aren't the same
    dims[1] = torch.randint(2, 5, (1,))

x = torch.rot90(x, k, dims)

# x is now a batch of 3D images that have all been rotated in the same random orientation
2 Answers

You could split the data in the batch randomly into 3 subsets, and apply each dimensional rotation respectively.

Let me expand on iacob's answer. Firstly, let me go over the parameters of rot90 function. Other than the input tensor, it expects k and dims where k is the number of rotations to be done, and dims is a list or tuple containing two dimensions on how the tensor to be rotated. If a tensor is 4D for example, dims could be [0, 3] or (1,2) or [2,3] etc. They have to be valid axes and it should contain two numbers. You don't really need to create tensors for this parameter or k. It is important to note that, depending on the given dims, output shape can drastically change:

x = torch.rand(15, 3, 4,6)
y1 = torch.rot90(x[0:5], 1, [1,3])     
y2 = torch.rot90(x[5:10], 1, [1,2])    
y3 = torch.rot90(x[10:15], 1, [2,3])   
print(y1.shape)   # torch.Size([5, 6, 4, 3])
print(y2.shape)   # torch.Size([5, 4, 3, 6])
print(y3.shape)   # torch.Size([5, 3, 6, 4])

Similar to iacob's answer, here we apply 3 different rotations to slices of the input. Note that how the output dimensions are all different, due to nature of rotations over different dimensions. You can't really join these results into one tensor, unless you have a really specific input size, for example Batch x 10 x 10 x 10 where rotating over combinations of 1,2,3 axes will always return same dimensions. You can however use each of these different sized output separately as inputs to different modules, layers etc.

I personally can't think of a use case where your random axes rotation can be used. If you can elaborate on why you are trying to do this, I can try to give some better solutions.

Related