PyTorch torch.scatter_add Function
PyTorch torch Reference Manual
torch.scatter_addis a function in PyTorch used to add the values of the source tensor to specified positions. It willsrcthe values according toindexthe specified positions to add toinputin.
Function Definition
torch.scatter_add(input, dim, index, src)
Parameters:
input(Tensor): Input tensor.dim(int): The dimension to scatter.index(Tensor): Index tensor, specifying where to add the values of src to input.src(Tensor): Source tensor, the values to add.
Return Value:
torch.Tensor: Returns the modified tensor.
Usage Example
Example
import torch
# Create input tensor
input = torch.zeros(3, 5)
# Create index and source
index = torch.tensor([[0, 1, 2, 0, 0],
[1, 2, 0, 1, 2],
[2, 0, 1, 2, 0]])
src = torch.tensor([[1, 1, 1, 1, 1],
[2, 2, 2, 2, 2],
[3, 3, 3, 3, 3]])
# Scatter and accumulate along dim=0
output = torch.scatter_add(input, dim=0, index=index, src=src)
print("Input:")
print(input)
print("nIndex:")
print(index)
print("nSource:")
print(src)
print("nResult:")
print(output)
# Create input tensor
input = torch.zeros(3, 5)
# Create index and source
index = torch.tensor([[0, 1, 2, 0, 0],
[1, 2, 0, 1, 2],
[2, 0, 1, 2, 0]])
src = torch.tensor([[1, 1, 1, 1, 1],
[2, 2, 2, 2, 2],
[3, 3, 3, 3, 3]])
# Scatter and accumulate along dim=0
output = torch.scatter_add(input, dim=0, index=index, src=src)
print("Input:")
print(input)
print("nIndex:")
print(index)
print("nSource:")
print(src)
print("nResult:")
print(output)
The output result is:
输入:
tensor([[0., 0., 0., 0., 0.],
[0., 0., 0., 0., 0.],
[0., 0., 0., 0., 0.]])
索引:
tensor([[0, 1, 2, 0, 0],
[1, 2, 0, 1, 2],
[2, 0, 1, 2, 0]])
源:
tensor([[1., 1., 1., 1., 1.],
[2., 2., 2., 2., 2.],
[3., 3., 3., 3., 3.]])
结果:
tensor([[4., 1., 2., 4., 4.],
[2., 1., 2., 2., 2.],
[3., 3., 1., 3., 3.]])
Example
import torch
# Use dim=1
input = torch.zeros(3, 5)
index = torch.tensor([[0, 1, 2, 1, 0],
[1, 2, 0, 2, 1],
[0, 1, 1, 0, 2]])
src = torch.arange(1, 6).float()
output = torch.scatter_add(input, dim=1, index=index, src=src)
print("Scatter along dim=1:")
print(output)
# Use dim=1
input = torch.zeros(3, 5)
index = torch.tensor([[0, 1, 2, 1, 0],
[1, 2, 0, 2, 1],
[0, 1, 1, 0, 2]])
src = torch.arange(1, 6).float()
output = torch.scatter_add(input, dim=1, index=index, src=src)
print("Scatter along dim=1:")
print(output)
The output result is:
沿 dim=1 散布:
tensor([[ 6., 3., 3., 0., 0.],
[ 3., 6., 3., 0., 0.],
[ 2., 4., 5., 0., 0.]])
Example
import torch
# Application scenario for aggregating values from multiple positions
# For example, accumulating neighbor node features in graph neural networks
# Simulate initial features of 4 nodes
node_features = torch.zeros(4, 3)
# Simulate edge connections (source nodes point to target nodes)
edge_index = torch.tensor([0, 1, 2, 3, 0, 1]) # Source nodes of edges
edge_weights = torch.tensor([1.0, 2.0, 3.0, 1.5, 2.5, 0.5])
# Create weighted values of source node features for each edge
src_features = torch.randn(6, 3) * edge_weights.unsqueeze(1)
# Accumulate the features to the target nodes (simplified here; in practice, it should be based on the target nodes of the edges)
target_nodes = torch.tensor([0, 0, 1, 1, 2, 3])
index = target_nodes
output = torch.scatter_add(node_features, 0, index.unsqueeze(1).expand_as(src_features), src_features)
print("Node feature shape:", node_features.shape)
print("Accumulated features:", output)
# Application scenario for aggregating values from multiple positions
# For example, accumulating neighbor node features in graph neural networks
# Simulate initial features of 4 nodes
node_features = torch.zeros(4, 3)
# Simulate edge connections (source nodes point to target nodes)
edge_index = torch.tensor([0, 1, 2, 3, 0, 1]) # Source nodes of edges
edge_weights = torch.tensor([1.0, 2.0, 3.0, 1.5, 2.5, 0.5])
# Create weighted values of source node features for each edge
src_features = torch.randn(6, 3) * edge_weights.unsqueeze(1)
# Accumulate the features to the target nodes (simplified here; in practice, it should be based on the target nodes of the edges)
target_nodes = torch.tensor([0, 0, 1, 1, 2, 3])
index = target_nodes
output = torch.scatter_add(node_features, 0, index.unsqueeze(1).expand_as(src_features), src_features)
print("Node feature shape:", node_features.shape)
print("Accumulated features:", output)
Note:torch.scatter_adddoes not modify the original input tensor, but returns a new tensor. Multiple indices can point to the same position, and the values will be accumulated. This function istorch.gatherthe inverse operation of.
Other Extensions