PyTorch torch.index_reduce Function


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

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)

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)

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)

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.


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions