update
This commit is contained in:
@@ -0,0 +1,9 @@
|
||||
import jax.numpy as jnp
|
||||
from jax import grad, jit, vmap
|
||||
|
||||
def sum_logistic(x):
|
||||
return jnp.sum(1.0 / (1.0 + jnp.exp(-x)))
|
||||
|
||||
x_small = jnp.arange(3.)
|
||||
derivative_fn = grad(sum_logistic)
|
||||
print(derivative_fn(x_small))
|
||||
Reference in New Issue
Block a user