Communitygithub.com

G1Joshi/Agent-Skills

Expert JAX assistance covering Autograd, XLA compilation (`jit`), vectorization (`vmap`), and parallelization (`pmap`). Use when building high-performance numerical computing and cutting-edge deep learning research.

Agent-Skills 是什么?

Agent-Skills is a Claude Code agent skill that expert JAX assistance covering Autograd, XLA compilation (`jit`), vectorization (`vmap`), and parallelization (`pmap`). Use when building high-performance numerical computing and cutting-edge deep learning research.

兼容平台~Claude Code~Codex CLI~Cursor
npx skills add https://github.com/G1Joshi/Agent-Skills/tree/HEAD/skills/ai-ml/jax

在你喜欢的 AI 中提问

打开一个已预加载此 Agent Skill 的新对话。

文档

JAX

JAX is "NumPy on steroids". It combines Autograd (automatic differentiation) with XLA (compilation). 2025 sees Flax NNX (PyTorch-style OOP) becoming standard.

When to Use

  • High-Performance Numerical & Scientific Computing: Autograd and XLA compilation for accelerated linear algebra on GPU/TPU.
  • Custom Deep Learning Research: Building neural networks with Flax, Haiku, or Equinox without framework bloat.
  • Composable Function Transformations: Seamlessly chaining jit (compilation), grad (derivatives), and vmap (vectorization).
  • Multi-Device Distributed Training: Parallelizing computation across clusters using pmap and shard maps (jax.experimental.shard_map).

Quick Start

import jax
import jax.numpy as jnp

# JIT-compiled matrix multiplication and automatic differentiation
@jax.jit
def loss_fn(w, x, y):
    pred = jnp.dot(x, w)
    return jnp.mean((pred - y) ** 2)

# Compute value and gradients simultaneously
grad_fn = jax.grad(loss_fn)

w = jnp.array([1.0, 2.0])
x = jnp.array([[1.0, 0.5], [2.0, 1.0]])
y = jnp.array([2.0, 4.0])

grads = grad_fn(w, x, y)
print("Gradients:", grads)

Core Concepts

#Composable Function Transformations: jit, grad & vmap

Combining just-in-time compilation, automatic differentiation, and automated batching:

import jax
import jax.numpy as jnp

# Pure function definition
def loss_fn(w, x, y):
    pred = jnp.dot(x, w)
    return jnp.mean((pred - y) ** 2)

# Transform 1: Automatic gradient with respect to weights
grad_fn = jax.grad(loss_fn)

# Transform 2: JIT compilation via XLA for extreme GPU speedup
fast_grad_fn = jax.jit(grad_fn)

# Transform 3: Automatic vectorization across batches
# batch_loss_fn handles 2D matrices automatically without loops
vmapped_loss = jax.vmap(loss_fn, in_axes=(None, 0, 0))

# Sample computation
key = jax.random.PRNGKey(42)
w = jax.random.normal(key, (3,))
x = jnp.array([[1.0, 2.0, 3.0]])
y = jnp.array([5.0])

grad_val = fast_grad_fn(w, x, y)
print("Computed Gradient:", grad_val)

#Explicit PRNG Key Management

State-free pseudorandom number generation:

import jax

# Initialize root PRNG key
key = jax.random.PRNGKey(2026)

# JAX keys must be explicitly split; never reused
key, subkey1, subkey2 = jax.random.split(key, 3)

data_normal = jax.random.normal(subkey1, shape=(4, 4))
data_uniform = jax.random.uniform(subkey2, shape=(4, 4))

print("Normal sample mean:", jnp.mean(data_normal))

#Stateful Model with Equinox & Pure Functions

Expressive, object-oriented neural networks built on pure JAX functions:

import equinox as eqx
import jax
import jax.numpy as jnp

class SimpleMLP(eqx.Module):
    linear1: eqx.nn.Linear
    linear2: eqx.nn.Linear

    def __init__(self, in_features, out_features, key):
        k1, k2 = jax.random.split(key)
        self.linear1 = eqx.nn.Linear(in_features, 64, key=k1)
        self.linear2 = eqx.nn.Linear(64, out_features, key=k2)

    def __call__(self, x):
        x = jax.nn.relu(self.linear1(x))
        return self.linear2(x)

key = jax.random.PRNGKey(0)
model = SimpleMLP(in_features=10, out_features=2, key=key)
sample_input = jnp.ones((10,))
output = model(sample_input)
print("Model output:", output)

Common Patterns

Batch Vectorization with vmap

Problem: Writing manual Python loops to evaluate functions over batches slows down execution.

Solution: Use jax.vmap to automatically vectorize single-example functions:

def predict_single(w, x):
    return jnp.dot(x, w)

# Automatically vectorize across batch dimension (axis 0 of x)
predict_batch = jax.vmap(predict_single, in_axes=(None, 0))

batch_x = jnp.ones((100, 10))
weights = jnp.ones(10)
predictions = predict_batch(weights, batch_x) # Output shape: (100,)

Best Practices (2026)

  • Do write purely functional code: zero side effects, no in-place array mutations (x.at[idx].set(val) instead of x[idx] = val).
  • Do wrap performance-critical functions with @jax.jit to trigger XLA compilation.
  • Do split PRNG keys explicitly (jax.random.split(key)) every time random numbers are generated.
  • Do use modern high-level libraries like Equinox or Flax Linen rather than writing raw parameter dictionaries.
  • Don't use standard Python conditionals (if x > 0:) inside JIT functions on dynamic tracers; use jax.lax.cond.
  • Don't reuse PRNG keys; reusing keys generates statistically correlated random numbers.
  • Don't mutate global variables inside functions transformed with jit, grad, or vmap.

Troubleshooting

ErrorCauseSolution
ConcretizationTypeError: Abstract tracer value used in Python if/whilePython conditional branching depends on dynamic JAX array values inside jit.Use jax.lax.cond or jax.lax.while_loop instead of standard Python if.
JAX arrays are immutableIn-place assignment attempted (arr[0] = 5).Use functional update syntax: arr = arr.at[0].set(5).
PRNGKey reuse warning / identical randomnessReusing identical PRNG key produces identical random numbers.Split keys before each random operation: key, subkey = jax.random.split(key).

References

Individual skills in this repo

This repo contains 11 individual skills — each has its own dedicated page.

G1Joshi/Agent-Skills

Expert Dask distributed computing assistance covering Dask DataFrames, Arrays, Futures, and cluster scaling. Use when analyzing datasets too large for pandas on a single machine or multi-node cluster.

G1Joshi/Agent-Skills

Expert dbt (data build tool) assistance covering SQL modeling, Jinja macros, tests, documentation, and semantic layer. Use when building analytics engineering pipelines on BigQuery, Snowflake, or PostgreSQL.

G1Joshi/Agent-Skills

Expert Ray distributed computing assistance covering Ray Core (actors, tasks), Ray Train, Ray Tune, and Ray Serve. Use when scaling Python compute and ML training across multi-node clusters.

G1Joshi/Agent-Skills

Expert Git version control assistance covering branching, interactive rebase, cherry-pick, submodules, worktrees, and conflict resolution. Use when managing source code history and collaboration workflows.

G1Joshi/Agent-Skills

Expert K9s CLI assistance covering terminal Kubernetes cluster navigation, real-time log streaming, port forwarding, and pod debugging. Use when managing and troubleshooting Kubernetes clusters with speed.

G1Joshi/Agent-Skills

Expert SWC assistance covering high-performance Rust-based JavaScript/TypeScript compilation, .swcrc configuration, Jest testing via @swc/jest, and minification. Use when replacing Babel with SWC for faster builds, accelerating test execution, or compiling modern ECMAScript features.

G1Joshi/Agent-Skills

Expert Tig assistance covering text-mode interface for Git, interactive staging, commit graph visualization, diff exploration, and blame navigation. Use when navigating Git commit history in the terminal, staging hunks interactively, browsing file changes, or reviewing revisions.

G1Joshi/Agent-Skills

Expert Vim assistance covering modal editing, .vimrc configuration, registers, macros, search/replace, and plugin management via vim-plug. Use when editing text efficiently in terminal environments, writing Vimscript, recording macros, or configuring core Vim settings.

G1Joshi/Agent-Skills

Expert Zed editor assistance covering high-performance Rust-based text editing, multi-buffer editing, language server protocols (LSP), and AI assistant integrations. Use when configuring Zed settings.json, setting up language extensions, collaborating in real-time channels, or optimizing editor startup speed.

G1Joshi/Agent-Skills

Expert Zsh assistance covering shell customization, Oh My Zsh plugins, Zinit plugin manager, prompt engineering (Starship/Powerlevel10k), and shell scripting. Use when configuring .zshrc, writing Zsh automation scripts, optimizing shell startup time, or configuring tab-completion.

G1Joshi/Agent-Skills

Expert [skill-name] assistance covering [feature 1], [feature 2], and [feature 3]. Use when [working with X], [debugging Y], or [implementing Z].

相关技能