x = torch.zeros(2, 1, 2, 1, 2)
x.size()
>>> torch.Size([2, 1, 2, 1, 2])
y = torch.squeeze(x) # remove 1
y.size()
>>> torch.Size([2, 2, 2])
y = torch.squeeze(x, 0)
y.size()
>>> torch.Size([2, 1, 2, 1, 2])
y = torch.squeeze(x, 1)
y.size()
>>> torch.Size([2, 2, 1, 2])