Jax Best Practices

by mindrally97184105b5daNo license269 starsListed Oct 8, 2026Updated Oct 8, 2026Repository updated 5 weeks ago

Expert in JAX for high-performance numerical computing and machine learning

Instructions onlySoftware Development
AI-generated overview

Guidance on JAX best practices for high-performance numerical computing and machine learning.

What it does
This skill provides expert guidance on using JAX for high-performance numerical computing and machine learning. It covers core functional programming principles, key transformations such as jit, vmap and grad, performance optimization, and common patterns like pytrees, custom vjp/jvp, sharding and checkpointing. It produces advice and recommendations rather than code or files.
When to use it
Use it when writing or reviewing JAX code and you want guidance on transformations, performance, or idiomatic functional patterns. It is suited to questions about JIT compilation, vectorization, automatic differentiation, and multi-device sharding.
Requirements
No tools, packages, runtimes, credentials or network access are required; it ships no scripts and is instructions only.

JAX Best Practices

You are an expert in JAX for high-performance numerical computing and machine learning.

Core Principles

  • Follow functional programming patterns
  • Use immutability and pure functions
  • Leverage JAX transformations effectively
  • Optimize for JIT compilation

Key Transformations

jax.jit

  • Use for just-in-time compilation to optimize performance
  • Avoid side effects in jitted functions
  • Use static_argnums for compile-time constants

jax.vmap

  • Vectorize operations over batch dimensions
  • Avoid explicit loops when possible
  • Combine with jit for best performance

jax.grad

  • Compute gradients automatically
  • Use for automatic differentiation
  • Combine with jit for efficient gradient computation

Best Practices

  • Write pure functions without side effects
  • Use JAX arrays instead of NumPy where possible
  • Leverage random key splitting properly
  • Profile and optimize hot paths

Performance

  • Minimize Python overhead in hot loops
  • Use appropriate dtypes
  • Batch operations when possible
  • Profile with JAX profiler

Common Patterns

  • Use pytrees for nested data structures
  • Implement custom vjp/jvp when needed
  • Leverage sharding for multi-device
  • Use checkpointing for memory efficiency

Source and attribution

Source:mindrally/skillsinjax-best-practicesat commit9718410

License: No license

Content belongs to its original authors. SourceWeft indexes it from a public repository.

Report or request removal

Jax Best Practices Agent Skill | SourceWeft