PyTorch中替代index_add_的高效index_max_实现方法问询
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:
1. Use torch.Tensor.scatter_reduce_ (Recommended for PyTorch 1.12+)
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 yourindex_add_usage).reduce='max': Tells the method to take the maximum value for each index group.include_self=False: Ensures we only use values fromchild(not the initial values inparent). If you wanted to compute the max betweenparent's original values and the aggregatedchildvalues, set this toTrue.
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

