Building Convolutional Neural Networks in JAX