I want to implement a vectorized version of the following function using numpy or pytorch:
def scatter_max_2(src, index, out):
src_shape = src.shape
for i in range(src_shape[0]):
for j in range(src_shape[1]):
out[i][index[i][j]] = max(out[i][index[i][j]],src[i][j])
return out
src = torch.tensor([[2, 0, 1, 4, 3], [0, 2, 1, 3, 4]])
index = torch.tensor([[4, 5, 4, 2, 3], [0, 0, 2, 2, 1]])
out = torch.zeros(2, 6, dtype=src.dtype)
out = scatter_max_2(src,index,out)
print(out)
Output:
tensor([[0, 0, 4, 3, 2, 0],
[2, 4, 3, 0, 0, 0]])
The function is a simpler implementation of scatter_max from https://github.com/rusty1s/pytorch_scatter
from torch_scatter import scatter_max
src = torch.tensor([[2, 0, 1, 4, 3], [0, 2, 1, 3, 4]])
index = torch.tensor([[4, 5, 4, 2, 3], [0, 0, 2, 2, 1]])
out, argmax = scatter_max(src, index, dim=-1)
print(out)
tensor([[0, 0, 4, 3, 2, 0],
[2, 4, 3, 0, 0, 0]])
However, they use torch.ops.torch_scatter.scatter_max from pytorch version 1.3 to implement the function and I could not find its implementation in pytorch github repo. An illustration of the function can be found here https://pytorch-scatter.readthedocs.io/en/1.3.0/functions/max.html