You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中替代index_add_的高效index_max_实现方法问询

Efficient Implementation of index_max_ in PyTorch

Great question! The loop-based approach you have works correctly, but it's inefficient—especially when dealing with large tensors where iterating over each index group adds significant overhead. Here are two optimized solutions to achieve the index_max_ functionality without explicit loops or excessive memory allocation:

PyTorch 1.12 introduced the scatter_reduce_ method, which is designed for exactly this kind of indexed aggregation operation (including max, min, sum, etc.). It operates in-place, just like index_add_, making it a perfect drop-in replacement for your desired index_max_ functionality.

Example Code:

import torch

idx = torch.LongTensor([0, 0, 0, 0, 1, 1])
child = torch.FloatTensor([1, 3, 5, 10, 8, 1])
parent = torch.FloatTensor([0, 0])

# Perform in-place indexed max aggregation
parent.scatter_reduce_(
    dim=0,
    index=idx,
    src=child,
    reduce='max',
    include_self=False  # Set to True if you want to compare with parent's original values
)

print(parent)  # Output: tensor([10.,  8.])

Key Details:

  • dim=0: Specifies we're aggregating along the 0th dimension (matching your index_add_ usage).
  • reduce='max': Tells the method to take the maximum value for each index group.
  • include_self=False: Ensures we only use values from child (not the initial values in parent). If you wanted to compute the max between parent's original values and the aggregated child values, set this to True.

2. Optimized Grouped Max for Older PyTorch Versions

If you're stuck on a PyTorch version older than 1.12, you can reduce loop overhead by only iterating over unique index values (instead of every possible index up to max(idx)). This cuts down the number of iterations significantly, especially when your index tensor has few unique values.

Example Code:

import torch

idx = torch.LongTensor([0, 0, 0, 0, 1, 1])
child = torch.FloatTensor([1, 3, 5, 10, 8, 1])
parent = torch.FloatTensor([0, 0])

# Get unique index values to minimize loop iterations
unique_indices = torch.unique(idx)

# Compute max for each unique index group and assign to parent
max_values = torch.stack([child[idx == i].max() for i in unique_indices])
parent[unique_indices] = max_values

print(parent)  # Output: tensor([10.,  8.])

This approach is still technically a loop, but it runs only once per unique index (instead of once per possible index value), which can drastically improve performance for large tensors with sparse index groups.


内容的提问来源于stack exchange,提问作者Kiv

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.29 08:15:03