Mode:
Duration:
1
Coding works best on desktop or with an external keyboard.
Coding works best on desktop or with an external keyboard.
A logistic regression implementation using JAX for binary classification.
import jax.numpy as jnp
from jax import grad
# Sample data
X = jnp.array([[0,0],[0,1],[1,0],[1,1]])
y = jnp.array([0,1,1,0]) # XOR example
# Initialize parameters
w = jnp.zeros(2)
b = 0.0
# Sigmoid function
def sigmoid(z):
return 1 / (1 + jnp.exp(-z))
# Loss function
def loss(w, b):
y_pred = sigmoid(jnp.dot(X, w) + b)
return -jnp.mean(y * jnp.log(y_pred) + (1-y) * jnp.log(1-y_pred))
grad_loss = grad(loss, argnums=(0,1))
# Gradient descent
for _ in range(1000):
dw, db = grad_loss(w, b)
w -= 0.1 * dw
b -= 0.1 * db
print('Learned weights:', w, 'bias:', b)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.