(weights, Y0s, sigma, mu_0t)
| 32 | |
| 33 | @jax.jit |
| 34 | def softmax_update(weights, Y0s, sigma, mu_0t): |
| 35 | mu_0tm1 = jnp.einsum("n,nij->ij", weights, Y0s) |
| 36 | return mu_0tm1, sigma |
| 37 | |
| 38 | |
| 39 | @jax.jit |
nothing calls this directly
no outgoing calls
no test coverage detected