rng = jax.random.PRNGKey(0) rng, inp_rng, init_rng = jax.random.split(rng, 3) delta = 0.42 factor = 0.42 @jax.jit def data_augmentation(image): new_image = pix.adjust_brightness(image=image, delta=delta) new_image = pix.random_brightness(image=new_image, max_delta=delta, key=inp_rng) new_image = pix.flip_up_down(image=image) new_image = pix.flip_left_right(image=new_image) new_image = pix.rot90(k=1, image=new_image) # k = number of times the rotation is applied return new_image __ __