# broadcasting in jax jxArrOne = jnp.array([[1, 2], [3, 4]]) jxArrTwo = jnp.array([10, 20]) print(jxArrOne + jxArrTwo) # output: [[11, 22], [13, 24]] __ __