# Scenario 1: One set of weights, many feature vectors # weights: don't map (None), features: map over axis 0 batch_predict = jax.vmap(dot_product, in_axes=(None, 0)) weights = jnp.array([1.0, 2.0, 3.0]) features = jnp.array([ [1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0], ]) results = batch_predict(weights, features) print(results) # [1., 2., 3.] __ __