jax-skills
High-performance numerical computing and machine learning workflows using JAX. Supports array operations, automatic differentiation, JIT compilation, RNN-style scans, map/reduce operations, and gradient computations. Ideal for scientific computing, ML models, and dynamic array transformations.
Install / Use
npx skills add benchflow-ai/skillsbench --skill jax-skillsInstalls into whichever agent you are using.
SKILL.md
Installable skill definition
Quality Score
Category
AutomationSupported Platforms
Our assessment of jax-skills
jax-skills scores 89/100 on our quality scale, 1341st of 2,877 Automation skills we index (top 47%).
Its SKILL.md is 4.1 KB long, well organised into 26 sections with 8 code examples: a solid amount of guidance for an agent.
With 1,813 GitHub stars, it is one of the more widely adopted skills in the catalogue.
Maintenance, license and trust
- The repository was last updated about 2 months ago, so jax-skills is actively maintained.
- It is released under the Apache-2.0 license, a permissive license that allows use, modification and commercial use with attribution.
- Its trust signals score 100/100, with no cautions. These come from repository metadata, not a code audit — read the skill file before letting an agent act on it.
Safety scan
No issues foundOur scan of the whole file found no instruction hijacking, hidden characters, credential access, data exfiltration or destructive commands.
Automated pattern scan on 2026-10-06. It catches known dangerous patterns, not every risk — read a skill before letting an agent act on it.
jax-skills compared with similar skills
All 4 of these similar skills score higher than jax-skills; compare them before choosing.
| Skill | Score | Stars | Updated | Format |
|---|---|---|---|---|
| jax-skills (this skill)by benchflow-ai | 89 | 1.8k | 2mo ago | SKILL.md |
| Agent-Reachby Panniantong | 100 | 92.1k | 20d ago | CLAUDE.md |
| headroomby headroomlabs-ai | 100 | 74.5k | today | CLAUDE.md |
| Scraplingby D4Vinci | 100 | 85.9k | today | MCP Server |
| crawl4aiby unclecode | 100 | 84.8k | today | MCP Server |
Frequently asked questions
- How do I install jax-skills?
- Run
npx skills add benchflow-ai/skillsbench --skill jax-skills. The install tabs above show the steps for each supported agent. - Which AI agents does jax-skills work with?
- It is written for Universal, as a SKILL.md file. Other agents that read the same format can often use it too.
- Is jax-skills safe to use?
- Our scan of the whole file found no instruction hijacking, hidden characters, credential access, data exfiltration or destructive commands. It is Apache-2.0-licensed and scores 100/100 on trust signals. Skills are instructions an agent will follow, so read the file before installing it and do not approve commands you do not understand.
- Is jax-skills still maintained?
- The repository was last updated about 2 months ago, so jax-skills is actively maintained.
Skill content
View source on GitHubname: jax-skills description: "High-performance numerical computing and machine learning workflows using JAX. Supports array operations, automatic differentiation, JIT compilation, RNN-style scans, map/reduce operations, and gradient computations. Ideal for scientific computing, ML models, and dynamic array transformations." license: Proprietary. LICENSE.txt has complete terms
Requirements for Outputs
General Guidelines
Arrays
- All arrays MUST be compatible with JAX (
jnp.array) or convertible from Python lists. - Use
.npy,.npz, JSON, or pickle for saving arrays.
Operations
- Validate input types and shapes for all functions.
- Maintain numerical stability for all operations.
- Provide meaningful error messages for unsupported operations or invalid inputs.
JAX Skills
1. Loading and Saving Arrays
load(path)
Description: Load a JAX-compatible array from a file. Supports .npy and .npz.
Parameters:
path(str): Path to the input file.
Returns: JAX array or dict of arrays if .npz.
import jax_skills as jx
arr = jx.load("data.npy")
arr_dict = jx.load("data.npz")
save(data, path)
Description: Save a JAX array or Python array to .npy.
Parameters:
- data (array): Array to save.
- path (str): File path to save.
jx.save(arr, "output.npy")
2. Map and Reduce Operations
map_op(array, op)
Description: Apply elementwise operations on an array using JAX vmap. Parameters:
- array (array): Input array.
- op (str): Operation name ("square" supported).
squared = jx.map_op(arr, "square")
reduce_op(array, op, axis)
Description: Reduce array along a given axis. Parameters:
- array (array): Input array.
- op (str): Operation name ("mean" supported).
- axis (int): Axis along which to reduce.
mean_vals = jx.reduce_op(arr, "mean", axis=0)
3. Gradients and Optimization
logistic_grad(x, y, w)
Description: Compute the gradient of logistic loss with respect to weights. Parameters:
- x (array): Input features.
- y (array): Labels.
- w (array): Weight vector.
grad_w = jx.logistic_grad(X_train, y_train, w_init)
Notes:
- Uses jax.grad for automatic differentiation.
- Logistic loss: mean(log(1 + exp(-y * (x @ w)))).
4. Recurrent Scan
rnn_scan(seq, Wx, Wh, b)
Description: Apply an RNN-style scan over a sequence using JAX lax.scan. Parameters:
- seq (array): Input sequence.
- Wx (array): Input-to-hidden weight matrix.
- Wh (array): Hidden-to-hidden weight matrix.
- b (array): Bias vector.
hseq = jx.rnn_scan(sequence, Wx, Wh, b)
Notes:
- Returns sequence of hidden states.
- Uses tanh activation.
5. JIT Compilation
jit_run(fn, args)
Description: JIT compile and run a function using JAX. Parameters:
- fn (callable): Function to compile.
- args (tuple): Arguments for the function.
result = jx.jit_run(my_function, (arg1, arg2))
Notes:
- Speeds up repeated function calls.
- Input shapes must be consistent across calls.
Best Practices
- Prefer JAX arrays (jnp.array) for all operations; convert to NumPy only when saving.
- Avoid side effects inside functions passed to vmap or scan.
- Validate input shapes for map_op, reduce_op, and rnn_scan.
- Use JIT compilation (jit_run) for compute-heavy functions.
- Save arrays using .npy or pickle/json to avoid system-specific issues.
Example Workflow
import jax.numpy as jnp
import jax_skills as jx
# Load array
arr = jx.load("data.npy")
# Square elements
arr2 = jx.map_op(arr, "square")
# Reduce along axis
mean_arr = jx.reduce_op(arr2, "mean", axis=0)
# Compute logistic gradient
grad_w = jx.logistic_grad(X_train, y_train, w_init)
# RNN scan
hseq = jx.rnn_scan(sequence, Wx, Wh, b)
# Save result
jx.save(hseq, "hseq.npy")
Notes
-
This skill set is designed for scientific computing, ML model prototyping, and dynamic array transformations.
-
Emphasizes JAX-native operations, automatic differentiation, and JIT compilation.
-
Avoid unnecessary conversions to NumPy; only convert when interacting with external file formats.
Related Skills
Agent-Reach
92.1kGive your AI agent eyes to see the entire internet. Read & search Twitter, Reddit, YouTube, GitHub, Bilibili, XiaoHongShu — one CLI, zero API fees.
headroom
74.5kCompress tool outputs, logs, files, and RAG chunks before they reach the LLM. 20% fewer tokens for coding agents, 60-95% fewer tokens for JSON, same answers. Library, proxy, MCP server.
Scrapling
85.9k🕷️ An adaptive Web Scraping framework that handles everything from a single request to a full-scale crawl! Don't be shy, join here: https://discord.gg/EMgGbDceNQ and follow here for daily tips and tricks: https://x.com/Scrapling_dev
crawl4ai
84.8kOpen-source web crawler and scraper for LLMs and AI agents: any website into clean, LLM-ready Markdown. Run it yourself, or use Crawl4AI Cloud with one key.
Languages
Trust signals
From repository metadata: license, adoption, age and documentation. Not a code audit — see the Safety scan above for what the skill file itself contains.
