Skip to main content

JAX quickstart

A minimal, runnable example. Install with pip install omnibias-jax.

The derivative tower

import jax.numpy as jnp
from omnibias.jax import get_activation

spec = get_activation("tanh")
z = jnp.array([0.5])
print(spec.fastpath(z, 3)) # 3rd derivative of tanh, closed form

A closed-form Laplacian, verified

import jax, jax.numpy as jnp
from omnibias.jax import neural_field_value_grad_laplacian

key = jax.random.PRNGKey(0)
D, H = 8, 32
kW, kb, kc, kx = jax.random.split(key, 4)
W = jax.random.normal(kW, (H, D)) / jnp.sqrt(D)
beta = jax.random.normal(kb, (H,))
c = jax.random.normal(kc, (H,))
b, x = 0.0, jax.random.normal(kx, (D,))

val, grad, lap = neural_field_value_grad_laplacian(x, W, beta, c, b, "tanh")

# Verify against autodiff (Laplacian = trace of Hessian).
def f(x):
return jnp.sum(c * jnp.tanh(W @ x + beta)) + b
lap_ad = jnp.trace(jax.hessian(f)(x))

print("closed form:", lap)
print("autodiff :", lap_ad)
print("abs diff :", abs(lap - lap_ad)) # ~1e-15

Going high-order

from omnibias.jax import neural_field_polylaplacian
d4 = neural_field_polylaplacian(x, W, beta, c, b, "tanh", k=2) # biharmonic

Next