Implement scatter max with numpy or pytorch for two dimensional array

Viewed 282

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

0 Answers
Related