PyTorch torch.scatter_add Function


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

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)

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)

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.


Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions