JAX is basically numpy on steroids and lets you do a lot of non-standard things (like a differentiable physics simulation or something) that would be harder with Pytorch.
They are both "high-performance."
Pytorch is more geared towards traditional deep learning and has the utilities and idioms to support it.