TensorFlow
TensorFlow (in graph-mode) generates a computation graph and tf.gradients is implemented as an operation on the graph, which outputs a new graph for the gradient.
For execution, there are different options: It could directly operate on the graph, or it could compile it to XLA, or you could also transform it to TFLite, etc.
JAX
JAX operates on functions. jax.grad gets a function and returns a function.
For execution, there are different options, but I think the most common is to compile to XLA.
Question
So, this both sounds very similar to me. A computation graph is just another way to represent a function. Is there any conceptual difference?
I often see that people say that jax.vmap is one big advantage of JAX, but you can do just the same in TensorFlow, e.g. tf.vectorized_map.
Or phrased a bit different: Is there any algorithm which you could implement in JAX but not in TF in the same way, or vice versa? Or which would be efficient in one case but not in the other?
This question is really only about the conceptual aspect, which should be a purely objective thing to answer: Either the answer is yes, they are conceptually the same (equally powerful), maybe with a short explanation, or no, they are conceptually different, with an example what is possible in one framework but conceptually not possible on the other. I don't want to have any discussion here on any subjective aspect, e.g. that maybe tf.vectorized_map is a bit buggy, that TF documentation is worse, or whatever.