PyTorch torch.narrow Function
Pytorch torch Reference Manual
torch.narrowis a function in PyTorch used to return tensor slices. It returns a sliced view starting from a specified position with a specified length along a specified dimension.
Function Definition
torch.narrow(input, dim, start, length)
Parameters:
input(Tensor): The input tensor.dim(int): The dimension to slice.start(int): The starting index.length(int): The length of the slice.
Return Value:
torch.Tensor: Returns the sliced view of the tensor.
Usage Examples
Example
import torch
# Create a tensor
x = torch.tensor([[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]])
# Slice 2 rows starting from index 0 along the first dimension (rows)
y = torch.narrow(x, dim=0, start=0, length=2)
print("Original tensor:")
print(x)
print("nSliced result:")
print(y)
# Create a tensor
x = torch.tensor([[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]])
# Slice 2 rows starting from index 0 along the first dimension (rows)
y = torch.narrow(x, dim=0, start=0, length=2)
print("Original tensor:")
print(x)
print("nSliced result:")
print(y)
The output is:
原始张量:
tensor([[ 1, 2, 3, 4],
[ 5, 6, 7, 8],
[ 9, 10, 11, 12]])
切片结果:
tensor([[1, 2, 3, 4],
[5, 6, 7, 8]])
Example
import torch
# Create a tensor
x = torch.tensor([[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]])
# Slice 2 columns starting from index 1 along the second dimension (columns)
y = torch.narrow(x, dim=1, start=1, length=2)
print("Original tensor:")
print(x)
print("nSliced result:")
print(y)
# Create a tensor
x = torch.tensor([[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]])
# Slice 2 columns starting from index 1 along the second dimension (columns)
y = torch.narrow(x, dim=1, start=1, length=2)
print("Original tensor:")
print(x)
print("nSliced result:")
print(y)
The output is:
原始张量:
tensor([[ 1, 2, 3, 4],
[ 5, 6, 7, 8],
[ 9, 10, 11, 12]])
切片结果:
tensor([[ 2, 3],
[ 6, 7],
[10, 11]])
Example
import torch
# Create a 3D tensor
x = torch.randn(5, 6, 7)
# Slice 3 elements starting from index 2 along the first dimension
y = torch.narrow(x, dim=0, start=2, length=3)
print("Original shape:", x.shape)
print("Shape after slicing:", y.shape)
# Create a 3D tensor
x = torch.randn(5, 6, 7)
# Slice 3 elements starting from index 2 along the first dimension
y = torch.narrow(x, dim=0, start=2, length=3)
print("Original shape:", x.shape)
print("Shape after slicing:", y.shape)
The output is:
原始形状: torch.Size([5, 6, 7]) 切片后形状: torch.Size([3, 6, 7])
Note:torch.narrowA view of the original tensor is returned, not a copy. If you need a copy, you can usetorch.narrow_copy。
Other Extensions