# Scenario 3: Many weights, many features (parallel evaluation) # Both: map over axis 0 parallel_predict = jax.vmap(dot_product, in_axes=(0, 0)) results = parallel_predict(many_weights, features) print(results) # [1., 2., 3.] __ __