PyTorch torch.unravel_index Function
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 indexshape: 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)
# 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]])
Other Extensions