# Scenario 2: Many weight sets, one feature vector (ensemble of models) # weights: map over axis 0, features: don't map (None) ensemble_predict = jax.vmap(dot_product, in_axes=(0, None)) many_weights = jnp.array([ [1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0], ]) single_features = jnp.array([1.0, 2.0, 3.0]) results = ensemble_predict(many_weights, single_features) print(results) # [1., 2., 3.] __ __