PyTorch torch.index_copy Function
Pytorch torch Reference Manual
torch.index_copyis a function in PyTorch used to copy the source tensor to the specified index positions. It, along the specified dimensiondim, atindexthe specified index positions, copiessourcethe value of.
andtorch.index_addThe difference is,index_copyis overwriting rather than accumulating.
Function Definition
torch.index_copy(input, dim, index, source)
Parameters:
input(Tensor): input tensor.dim(int): the dimension of the index.index(Tensor): a one-dimensional integer tensor specifying the positions to copy to.source(Tensor): the source tensor, the values to copy.
Return value:
torch.Tensor: returns the modified tensor.
Usage Example
Example
import torch
# Create the input tensor
input = torch.randn(4, 5)
# Create the index and source
index = torch.tensor([0, 2, 3])
source = torch.randn(3, 5)
# Copy along dim=0
output = torch.index_copy(input, dim=0, index=index, source=source)
print(Input:)
print(input)
print(nSource:)
print(source)
print(nThe result after copying to positions [0, 2, 3]:)
print(output)
# Create the input tensor
input = torch.randn(4, 5)
# Create the index and source
index = torch.tensor([0, 2, 3])
source = torch.randn(3, 5)
# Copy along dim=0
output = torch.index_copy(input, dim=0, index=index, source=source)
print(Input:)
print(input)
print(nSource:)
print(source)
print(nThe result after copying to positions [0, 2, 3]:)
print(output)
The output is:
输入:
tensor([[ 0.3456, -0.1234, 0.5678, -0.2345, 0.8901],
[-0.5678, 0.1234, -0.6789, 0.2345, -0.1234],
[ 0.7890, -0.3456, 0.1234, -0.5678, 0.3456],
[-0.1234, 0.4567, -0.8901, 0.6789, -0.5678]])
源:
tensor([[-1.2345, 0.5678, -1.2345, 0.5678, -1.2345],
[ 1.5678, -0.6789, 1.5678, -0.6789, 1.5678],
[-0.8901, 1.2345, -0.8901, 1.2345, -0.8901]])
复制到位置 [0, 2, 3] 后的结果:
tensor([[-1.2345, 0.5678, -1.2345, 0.5678, -1.2345],
[-0.5678, 0.1234, -0.6789, 0.2345, -0.1234],
[ 1.5678, -0.6789, 1.5678, -0.6789, 1.5678],
[-0.8901, 1.2345, -0.8901, 1.2345, -0.8901]])
Example
import torch
# Copy along dim=1
input = torch.zeros(3, 5)
index = torch.tensor([1, 3])
source = torch.tensor([[10, 20, 30, 40, 50],
[60, 70, 80, 90, 100]])
output = torch.index_copy(input, dim=1, index=index, source=source)
print(Input:)
print(input)
print(nSource:)
print(source)
print(nThe result after copying to positions [1, 3]:)
print(output)
# Copy along dim=1
input = torch.zeros(3, 5)
index = torch.tensor([1, 3])
source = torch.tensor([[10, 20, 30, 40, 50],
[60, 70, 80, 90, 100]])
output = torch.index_copy(input, dim=1, index=index, source=source)
print(Input:)
print(input)
print(nSource:)
print(source)
print(nThe result after copying to positions [1, 3]:)
print(output)
The output is:
输入:
tensor([[0., 0., 0., 0., 0.],
[0., 0., 0., 0., 0.],
[0., 0., 0., 0., 0.]])
源:
tensor([[ 10., 20., 30., 40., 50.],
[ 60., 70., 80., 90., 100.]])
复制到位置 [1, 3] 后的结果:
tensor([[ 0., 10., 0., 20., 0.],
[ 0., 60., 0., 70., 0.],
[ 0., 0., 0., 0., 0.]])
Example
import torch
# Build a large tensor
# Suppose we need to merge the results of multiple small batches into one large batch
# Target tensor
batch_size = 8
feature_dim = 4
output = torch.zeros(batch_size, feature_dim)
# Simulate the results of multiple small batches
mini_batches = [
torch.randn(2, feature_dim),
torch.randn(3, feature_dim),
torch.randn(1, feature_dim)
]
# The position where each batch should be placed
indices = [0, 2, 5]
# Copy each batch in sequence
for idx, batch in zip(indices, mini_batches):
# Create an index of corresponding size
index = torch.arange(idx, idx + len(batch))
output = torch.index_copy(output, dim=0, index=index, source=batch)
print(Final output shape:, output.shape)
print(output)
# Build a large tensor
# Suppose we need to merge the results of multiple small batches into one large batch
# Target tensor
batch_size = 8
feature_dim = 4
output = torch.zeros(batch_size, feature_dim)
# Simulate the results of multiple small batches
mini_batches = [
torch.randn(2, feature_dim),
torch.randn(3, feature_dim),
torch.randn(1, feature_dim)
]
# The position where each batch should be placed
indices = [0, 2, 5]
# Copy each batch in sequence
for idx, batch in zip(indices, mini_batches):
# Create an index of corresponding size
index = torch.arange(idx, idx + len(batch))
output = torch.index_copy(output, dim=0, index=index, source=batch)
print(Final output shape:, output.shape)
print(output)
The output is:
最终输出形状: torch.Size([8, 4])
tensor([[ 0.1234, -0.5678, 0.8901, -0.2345],
[ 0.6789, -0.1234, -0.5678, 0.3456],
[ 1.2345, -0.8901, 0.1234, -0.6789],
[-0.3456, 0.5678, -0.1234, 0.8901],
[ 1.5678, -0.2345, 0.6789, -0.1234],
[-0.8901, 0.3456, 0.5678, -0.8901],
[ 0.0000, 0.0000, 0.0000, 0.0000],
[ 0.0000, 0.0000, 0.0000, 0.0000]])
Note:torch.index_copydoes not modify the original input tensor, but returns a new tensor.index_copyis an overwrite operation, unliketorch.index_addthe accumulation operation is different.
Other Extensions