PyTorch torch.index_add Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.index_addis a function in PyTorch used to add the values of a source tensor to specified index positions. Along the specified dimensiondim, atindexadd at the specified index positionssourcethe values of.

Function Definition

torch.index_add(input, dim, index, source, *, alpha=1)

Parameters:

  • input(Tensor): Input tensor.
  • dim(int): The dimension of the index.
  • index(Tensor): A one-dimensional integer tensor specifying the positions to add to.
  • source(Tensor): Source tensor, the values to be added.
  • alpha(float, optional): Scaling factor for source, default is 1.

Return Value:

  • torch.Tensor: Returns the modified tensor.

Usage Examples

Example

import torch

# Create input tensor
input = torch.randn(4, 5)

# Create index and source
index = torch.tensor([0, 2, 3])
source = torch.randn(3, 5)

# Add along dim=0
output = torch.index_add(input, dim=0, index=index, source=source)

print("Input shape:", input.shape)
print("Index:", index)
print("Source shape:", source.shape)
print("Result shape:", output.shape)
print("nResult:")
print(output)

Output result:

输入形状: torch.Size([4, 5])
索引: tensor([0, 2, 3])
源形状: torch.Size([3, 5])
结果形状: torch.Size([4, 5])

结果:
tensor([[ 1.8435,  0.3463, -0.1024,  0.5678,  0.1234],
        [-0.5678,  0.8901, -0.2345,  0.6789, -0.1234],
        [ 2.3456,  0.4567,  0.7890, -0.3456,  0.5678],
        [-0.7890,  1.2345,  0.3456, -0.8901,  0.2345]])

Example

import torch

# Use alpha parameter to scale source
input = torch.zeros(5)
index = torch.tensor([0, 2, 4])
source = torch.tensor([10, 20, 30])

# alpha=2 means add after multiplying source by 2
output = torch.index_add(input, dim=0, index=index, source=source, alpha=2)

print("Input:", input)
print("Source:", source)
print("Result after alpha=2:", output)

Output result:

输入: tensor([0., 0., 0., 0., 0.])
源: tensor([10., 20., 30.])
alpha=2 后的结果: tensor([20.,  0., 40.,  0., 60.])

Example

import torch

# Add along another dimension
input = torch.zeros(3, 4, 5)
index = torch.tensor([1, 3])
source = torch.randn(2, 4, 5)

# Add along dim=1
output = torch.index_add(input, dim=1, index=index, source=source)

print("Input shape:", input.shape)
print("Index shape:", index.shape)
print("Source shape:", source.shape)
print("Result shape:", output.shape)

Output result:

输入形状: torch.Size([3, 4, 5])
索引形状: torch.Size([2])
源形状: torch.Size([2, 4, 5])
结果形状: torch.Size([3, 4, 5])

Example

import torch

# Application in neural networks: attention mechanism
# Assume there are multiple key-value pairs that need to be aggregated to the query

# Simulate query and key-value
num_queries = 2
num_kv = 4
dim = 3

# Query indices
query_idx = torch.tensor([0, 1])
# Corresponding values
values = torch.randn(num_queries, dim) * 10

# Output
output = torch.zeros(num_kv, dim)
# Add value to the corresponding position
output = torch.index_add(output, dim=0, index=query_idx, source=values)

print("Query indices:", query_idx)
print("Values:", values)
print("Aggregation result:", output)

Note:torch.index_adddoes not modify the original input tensor, but returns a new tensor. Ifindexthere are duplicate indices in [index], the values will be accumulated.alphaThe [alpha] parameter can be used to scale the source values.


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions