smithery.ai

jax

Essential tools for using JAX in machine learning and mathematical analysis, covering core concepts, transformations, ML specifics, control flow, and parallelism.

First seen Mar 27, 2026

Installation

$ npx skills add https://smithery.ai

Also in this package

Other skills from smithery.ai · top by installs.

npx skills add https://smithery.ai

Browse all from smithery.ai

More details

Agent compatibility

Declared targets from SKILL.md / docs. Unmarked agents are not listed — the skill may still install via the CLI.

Claude Code Not declared
Cursor Not declared
Codex Not declared
GitHub Copilot Not declared
Windsurf Not declared
Gemini CLI Not declared
Cline Not declared
OpenCode Not declared

Package contents

Files included with this skill beyond the listing page.

  • skill md SKILL.md 1,060 B
  • docs SUMMARY.md 173 B

History

  1. First seen on skills.sh
  2. First recorded snapshot · 1 installs

SKILL.md

JAX Skill

JAX is Autograd and XLA, brought together for high-performance machine learning research.

Contents

  • [Concepts & Theory](reference.md)

- Immutability - The 4 Transformations - Pytrees

  • [Code Examples](examples.md)

- jit, grad, vmap, random usage - Control Flow (scan, cond, fori_loop) - Parallelism (sharding)

Common Workflows

1. Developing a new Model

  1. Define your parameters as a Pytree (dict/dataclass).
  2. Define your forward pass function (pure).
  3. Define your loss function.
  4. Use jax.valueandgrad to get gradients.
  5. Use jax.jit to speed up the update step.
  6. See [examples.md](examples.md) for snippets.

2. Debugging Shapes/NaNs

  1. Disable JIT: jax.config.update("jaxdisablejit", True) to debug with standard python tools.
  2. Use jax.debug.print inside JITted functions.