Optimizing torch mean over a dimension in a random batch

Viewed 138

I am looking for a way to optimize the following code in pytorch.

I have a function f defined over space x,y and time t.
In a random batch, I need to compute the average over all the same timestamps. I was able to achieve this with the following inefficient for-loop

import torch
# Space (x,y) and time (t) coordinates in a random batch
x = torch.Tensor([[0, 0, 1, 0],[3, 2, 2, 1],[1,3,5,5]]).T 
# compute a dummy function u = f(t,x,y)
f = (x**2 + 0.5)[:,:2]
# timestamps
t = x[:,0]

# get unique timestamps
val = torch.unique(t.squeeze())
for v in val:
    # compute a mask for all timestamp equal to v
    mask = t == v
    # average over the spatial coordinates
    f[mask,:] = torch.mean(f[mask,:], dim=0)
print(f)

Which results in

f = tensor([[0.5000, 5.1667],
            [0.5000, 5.1667],
            [1.5000, 4.5000],
            [0.5000, 5.1667]])

Is there a way to make this computation faster?

1 Answers

I think you are looking for index_add_:

avg_size = int(t.max().item()) + 1  # number of rows in output tensor
z = torch.zeros((avg_size, f.shape[1]), dtype=f.dtype)
s = torch.index_add(s, 0, t.long(), f)  # sum the elements of f
c = torch.index_add(s, 0, t.long(), torch.ones_like(f[:, :1]))  # count how many at each entry
out = s / c  # divide to get the mean
Related