PyTorch torch.index_copy Function


Pytorch torch 参考手册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)

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)

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)

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.


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions