# Using Jax

> Applies JAX patterns for scientific Python development. Use when working with JAX, distrax, numpyro, blackjax, or scientific computing. Covers vmap, JIT, RNG handling. Use when this capability is needed.

- Skill: `tomevault-io/using-jax` (Agent Skill, multi-file: 2 files)
- Install (CLI): `npx skillmds@latest add tomevault-io/using-jax`
- Raw SKILL.md: https://api.skillmd.com/api/skills/tomevault-io/using-jax/raw
- Safety review: pending (external: skill-scanner PASS, skillspector PASS)
- Works with: Claude Code, Claude.ai, OpenAI Codex
- Category: Coding & Dev Tools
- Author: tomevault-io (https://skillmd.com/u/tomevault-io)
- Updated: 2026-09-17
- Page: https://skillmd.com/skills/tomevault-io/using-jax

---


# JAX Scientific Computing

## Core Rules

1. **Pure functions** - No side effects
2. **JIT outer functions** - `@jax.jit` on hot paths
3. **vmap not loops** - `jax.vmap(fn)` instead of list comprehensions
4. **Split RNG keys** - Never reuse keys

## Patterns

```python
# RNG: always split
key, k1, k2 = jax.random.split(key, 3)

# Batching: vmap not loops
batched = jax.vmap(fn)(inputs)

# Loops: use scan
_, results = jax.lax.scan(step_fn, init, xs)
```

## Gotchas

- Arrays are immutable
- No Python control flow in JIT - use `jax.lax.cond`, `jax.lax.scan`
- Check NaNs: `jnp.isnan(x).any()`

---
> Converted and distributed by [TomeVault](https://tomevault.io/claim/yallup) — claim your Tome and manage your conversions.
<!-- tomevault:4.0:skill_md:2026-04-13 -->

