Mode:
Duration:
1
Coding works best on desktop or with an external keyboard.
Coding works best on desktop or with an external keyboard.
Forward pass of a simple 2-layer neural network in JAX.
import jax.numpy as jnp
# Input
x = jnp.array([1.0, 2.0, 3.0])
# Network parameters
W1 = jnp.array([[0.1,0.2,0.3],[0.4,0.5,0.6]])
b1 = jnp.array([0.1,0.2])
W2 = jnp.array([[0.7,0.8]])
b2 = jnp.array([0.3])
# Forward pass
def relu(x):
return jnp.maximum(0, x)
h = relu(jnp.dot(W1, x) + b1)
y_pred = jnp.dot(W2, h) + b2
print('Output:', y_pred)JAX is an open-source Python library for high-performance numerical computing, combining NumPy-like API with automatic differentiation (autograd), GPU/TPU acceleration, and composable function transformations for machine learning and scientific computing.
Origin & Creator
JAX was developed by researchers at Google Research starting in 2018, building on Autograd and XLA (Accelerated Linear Algebra) to enable high-performance, differentiable programming.
Industrial Note
JAX is widely used in cutting-edge machine learning research, physics simulations, reinforcement learning, and differentiable programming, where composable gradients and hardware acceleration are essential.