Files
FYS-STK4155/doc/src/week39/jax.py
T
Morten Hjorth-Jensen bff99c3750 update
2023-09-24 22:04:27 +02:00

11 lines
233 B
Python

import jax.nuympy as jnp
from jax import grad, jit, vmap
from jax import random
def sum_logistic(x):
return jnp.sum(1.0 / (1.0 + jnp.exp(-x)))
x_small = arange(3.)
derivative_fn = grad(sum_logistic)
print(derivative_fn(x_small))