Thank you! It works:
def jigsaw_to_image(x, grid_size=(2, 2)):
# x shape is batch_size x num_patches x c x jigsaw_h x jigsaw_w
batch_size, num_patches, c, jigsaw_h, jigsaw_w = x.size()
assert num_patches == grid_size[0] * grid_size[1]
x_image = x.view(batch_size, grid_size[0], grid_size[1], c, jigsaw_h, jigsaw_w)
output_h = grid_size[0] * jigsaw_h
output_w = grid_size[1] * jigsaw_w
x_image = x_image.permute(0, 3, 1, 4, 2, 5).contiguous()
x_image = x_image.view(batch_size, c, output_h, output_w)
return x_image