PyTorch torch.index_add Function
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)
# 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)
# 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)
# 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)
# 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.
Other Extensions