PyTorch torch.unravel_index Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.unravel_indexis a function in PyTorch used to unravel flattened indices into multi-dimensional indices. Unliketorch.ravel_indexconversely, it converts a position in a flattened array back into coordinates in a multi-dimensional array.

Function Definition

torch.unravel_index(indices, shape)

Parameter Description

  • indices: Flattened index
  • shape: Shape of the multi-dimensional array

Usage Example

Example

import torch

# Flattened index
indices = torch.tensor([0, 1, 5, 6, 7])

# Array shape
shape = (2, 4)

# Unravel into multi-dimensional index
result = torch.unravel_index(indices, shape)

print("Flattened index:", indices)
print("Array shape:", shape)
print("Multi-dimensional index:")
print(result)

The output is:

展平索引: tensor([0, 1, 5, 6, 7])
数组形状: (2, 4)
多维索引:
tensor([[0, 0, 1, 1, 1],
        [0, 1, 1, 0, 1]])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions