When I first read about JAX I thought it would kill Pytorch, but I'm not sure I can get on with an immutable language for tensor operations in deep learning.
If I have an array `x` and want to set index 0 to 10, I cannot do:
x[0] = 10
I instead have to do: y = x.at[0].set(10)
I'm sure I could get used to it, but it really puts me off.