PyTorch torch.index_reduce Function
Pytorch torch Reference Manual
torch.index_reduceis a function in PyTorch used to aggregate the values of a source tensor to specified index positions in a specified manner. It, along the specified dimensiondim, atindexspecified index positions, aggregate in the specified mannersourcethe values of.
Function Definition
torch.index_reduce(input, dim, index, source, reduce='mean', *, include_self=True)
Parameters:
input(Tensor): Input tensor.dim(int): The dimension of the index.index(Tensor): A 1D integer tensor specifying the positions to aggregate to.source(Tensor): The source tensor, the values to be aggregated.reduce(str): The aggregation method, optional values are 'mean', 'prod', 'amax', 'amin'. Default is 'mean'.include_self(bool, optional): Whether to include the original values at the index positions themselves in the aggregation. Default is True.
Return Value:
torch.Tensor: Returns the aggregated tensor.
Usage Examples
Example
import torch
# Create the input tensor
input = torch.randn(3, 3)
# Create indices and source
index = torch.tensor([0, 0, 0])
source = torch.tensor([1.0, 2.0, 3.0])
# Use mean aggregation
output = torch.index_reduce(input, dim=0, index=index, source=source, reduce='mean')
print("Input:")
print(input)
print("nIndex:", index)
print("Source:", source)
print("nmean aggregation result:")
print(output)
# Create the input tensor
input = torch.randn(3, 3)
# Create indices and source
index = torch.tensor([0, 0, 0])
source = torch.tensor([1.0, 2.0, 3.0])
# Use mean aggregation
output = torch.index_reduce(input, dim=0, index=index, source=source, reduce='mean')
print("Input:")
print(input)
print("nIndex:", index)
print("Source:", source)
print("nmean aggregation result:")
print(output)
The output result is:
输入:
tensor([[ 0.3456, -0.1234, 0.5678],
[-0.5678, 0.1234, -0.6789],
[ 0.7890, -0.3456, 0.1234]])
索引: tensor([0, 0, 0])
源: tensor([1., 2., 3.])
mean 聚合结果:
tensor([[ 2.5237, 0.2931, 0.3374],
[-0.5678, 0.1234, -0.6789],
[ 0.7890, -0.3456, 0.1234]])
Example
import torch
# Test different aggregation methods
input = torch.ones(3)
index = torch.tensor([0, 0, 0])
source = torch.tensor([2.0, 4.0, 6.0])
# prod aggregation
output_prod = torch.index_reduce(input, 0, index, source, reduce='prod')
print("prod aggregation:", output_prod)
# amax aggregation
output_max = torch.index_reduce(input, 0, index, source, reduce='amax')
print("amax aggregation:", output_max)
# amin aggregation
output_min = torch.index_reduce(input, 0, index, source, reduce='amin')
print("amin aggregation:", output_min)
# Test different aggregation methods
input = torch.ones(3)
index = torch.tensor([0, 0, 0])
source = torch.tensor([2.0, 4.0, 6.0])
# prod aggregation
output_prod = torch.index_reduce(input, 0, index, source, reduce='prod')
print("prod aggregation:", output_prod)
# amax aggregation
output_max = torch.index_reduce(input, 0, index, source, reduce='amax')
print("amax aggregation:", output_max)
# amin aggregation
output_min = torch.index_reduce(input, 0, index, source, reduce='amin')
print("amin aggregation:", output_min)
The output result is:
prod 聚合: tensor([48., 1., 1.]) amax 聚合: tensor([7., 1., 1.]) amin 聚合: tensor([3., 1., 1.])
Example
import torch
# include_self parameter
input = torch.tensor([1.0, 10.0, 100.0])
index = torch.tensor([0, 0])
source = torch.tensor([5.0, 5.0])
# Include self values (default)
output1 = torch.index_reduce(input, 0, index, source, reduce='mean', include_self=True)
print("include_self=True:", output1)
# Exclude self values
output2 = torch.index_reduce(input, 0, index, source, reduce='mean', include_self=False)
print("include_self=False:", output2)
# include_self parameter
input = torch.tensor([1.0, 10.0, 100.0])
index = torch.tensor([0, 0])
source = torch.tensor([5.0, 5.0])
# Include self values (default)
output1 = torch.index_reduce(input, 0, index, source, reduce='mean', include_self=True)
print("include_self=True:", output1)
# Exclude self values
output2 = torch.index_reduce(input, 0, index, source, reduce='mean', include_self=False)
print("include_self=False:", output2)
The output result is:
include_self=True: tensor([ 3.6667, 10.0000, 100.0000]) include_self=False: tensor([ 5., 10., 100.])
Example
import torch
# 2D tensor application
input = torch.zeros(3, 4)
index = torch.tensor([0, 2, 2])
source = torch.randn(3, 4)
# mean aggregation
output = torch.index_reduce(input, dim=0, index=index, source=source, reduce='mean')
print("Input shape:", input.shape)
print("Index:", index)
print("Source shape:", source.shape)
print("Result shape:", output.shape)
print("nResult:")
print(output)
# 2D tensor application
input = torch.zeros(3, 4)
index = torch.tensor([0, 2, 2])
source = torch.randn(3, 4)
# mean aggregation
output = torch.index_reduce(input, dim=0, index=index, source=source, reduce='mean')
print("Input shape:", input.shape)
print("Index:", index)
print("Source shape:", source.shape)
print("Result shape:", output.shape)
print("nResult:")
print(output)
Note:torch.index_reduceIt does not modify the original input tensor, but returns a new tensor. Multiple indices can point to the same position, and the values will be aggregated according to the specified aggregation method.include_self=FalseIt is useful in some scenarios, as it can exclude the original values and aggregate only newly added values.
Other Extensions