3 ms·
import torch x = torch.randn(4000, 4000, device='cuda:0') y = torch.randn(4000, 4000, device='cuda:0') import time torch.cuda.synchronize() t0 = time
by t-vi 5y ago
import torch
x = torch.randn(4000, 4000, device='cuda:0')
y = torch.randn(4000, 4000, device='cuda:0')
import time
torch.cuda.synchronize()
t0 = time.perf_counter()
z = torch.zeros(4000, 4000, device='cuda:0')
for j in range(4000):
z += (x[:, j] * z[j, :]).relu()
torch.cuda.synchronize()
t1 = time.perf_counter()
print(t1-t0)
seem like 640ms-ish for me (on a rtx3090). How much faster is Julia, then? Is it double or single digit ms?
Of course, the backward will be more work because you need to recompute which bits were 0.
I can appreciate that Python is slow and that Julia has very nice properties dissolving the boundaries of operators that are awefully hard in PyTorch & Co, but this is one of these examples where the story seems rather flawed (you don't need the for loop over 64 billion elements, just over the 4000 j-coordinates, and the Python overhead here isn't too bad, either). As far as I can tell, the large overhead here is from materializing the intermediate results in global memory. Personally, I would hope that the fusers might be able to fuse in the not to far future, so the entire thing can run on the GPU in one kernel, but until then, if Julia can do that, it'll have a huge advantage here.
Shameless plug: In my courses about efficient PyTorch, I always give the rule of thumb "avoid operations on individual elements, but things operating on hundreds+ elements should be fine". And the reason isn't as much just "Python is slow" but the overhead of the tensor objects.
- ChrisRackauckas 5y ago>I always give the rule of thumb "avoid operations on individual elements, but things operating on hundreds+ elements should be fine". And the reason isn't as much just "Python is slow" but the overhead of the tensor objects. That's just not possible in a lot of cases, like when defining a large nonlinear ODE which is precisely described by the actions of individual elements. Indeed, it's not the ML case, but it is a case where PyTorch's assumptions are too strong.