Mode:
Duration:
1
Coding works best on desktop or with an external keyboard.
Coding works best on desktop or with an external keyboard.
Computing mean squared error using JAX for vectorized operations.
import jax.numpy as jnp
# Sample predictions and targets
y_true = jnp.array([1.0,2.0,3.0])
y_pred = jnp.array([1.1,1.9,3.2])
# Mean squared error
def mse(y_true, y_pred):
return jnp.mean((y_true - y_pred)**2)
print('MSE:', mse(y_true, 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.