jax-development

v2026.09.24

Write, debug, review, profile, or shard JAX numerical code. Use when the hard part is JAX tracing, autodiff, control flow, PRNGs, compilation, array placement, or runtime performance. Do not impose JAX on a NumPy-only or unrelated GPU task.

GitHub
安装命令
npx skhub add tristanmanchester/jax-development
Markdown
SKILL.md

JAX development

Make the mathematical contract explicit before changing transformations or kernels. Record shapes, dtypes, numerical tolerances, randomness, differentiability, and required backend. Preserve useful existing architecture; a new API does not itself justify a rewrite. Use current public interfaces rather than compatibility shims.

Inspect, reproduce, measure

Resolve SKILL_DIR to the directory containing this file. Bundled scripts live there, not in the target project's scripts/. Run them in the project's actual Python environment. Inspect each script's help before use.

python "$SKILL_DIR/scripts/jax_env_report.py" --format json
python "$SKILL_DIR/scripts/jax_project_scan.py" /absolute/project --format json
python "$SKILL_DIR/scripts/jax_compile_probe.py" --help
python "$SKILL_DIR/scripts/jax_recompile_explorer.py" --help

Environment probing initialises the backend and can allocate resources. The report lists relevant environment variable names, never their values, and returns nonzero for import/backend/smoke-test failure. It is not an exhaustive secret scanner: review paths, devices, and private project output before sharing. Static scans produce leads, not proof that a transformation is wrong. Importing, tracing, lowering, and benchmarking a module execute code; use trusted inputs.

Reduce failures to the smallest reproducer retaining the relevant shape, dtype, static argument, transform order, and backend. Compare against a simple reference and check numerical/gradient invariants before measuring speed. Change one hypothesis at a time and record the observed result.

Select the relevant reference

The source baseline records current changes that affect these longer references. Check the installed release rather than treating online latest/unreleased docs as the API of a locked project. Useful templates and reproducer/evaluation assets remain bundled; run selected examples against the target environment.

Correctness and performance rules

Keep ordinary transformed functions pure; explicit jax.ref state is a separate supported model with its own transform/effect restrictions, not permission for hidden Python mutation. Thread typed PRNG keys and avoid reuse. Use structured control flow for traced decisions; static Python loops can be appropriate when small, so do not mechanically replace every loop with scan.

Keep host transfers and synchronisation deliberate. Separate input preparation, trace/compile, dispatch, device execution, output materialisation, and communication. Check x64 and matmul precision explicitly. jax.numpy.empty no longer promises zero-initialised storage on the current release; use zeros when zeros are needed.

Donation consumes input buffers. For independent timing trials, construct fresh inputs for every call; never reuse a donated warm-up argument. The benchmark harness now requires an input factory and propagates synchronisation errors:

python "$SKILL_DIR/scripts/jax_benchmark_harness.py" \
  --file /absolute/project/benchmark_case.py --function step --factory make_inputs \
  --jit --donate-argnums 0 --repeat 20

make_inputs() returns (args, kwargs) with equivalent, fresh inputs. Preparation and its synchronisation are outside the timer; output synchronisation is inside. Keep shapes/dtypes/shardings/static values fixed and verify outputs separately. First-call time includes compilation/execution and may use caches; it is not pure compile time. For stateful training throughput, write a separate benchmark carrying returned state forward, and report that different workload. No old JSON/arrayify or unsafe donation compatibility path remains.

Prefer global-view code with deliberate sharding before manual shard_map when it meets the objective. NamedSharding placement is not synonymous with explicit sharding-in-types. Current mesh context uses jax.set_mesh(mesh), not with mesh. Test global/local shapes, replication, collectives, gradients, and output sharding; a one-device run cannot validate multi-host communication.

Finish with evidence

Return the diagnosis, patch/example, correctness checks, measured timings with hardware and workload, and remaining backend limitations. Do not claim benchmarks or compilation succeeded when a helper returned partial/error output. For helper regressions run python -m unittest discover -s "$SKILL_DIR/tests" -v.

发现
标签

此技能尚未发布标签。

版本
最新版本元数据

版本

v2026.09.24

发布时间

2026年9月24日

分类

未分类

许可证

未指定

源路径

jax-development

默认分支

main

最新提交

3323bc9

Tree SHA

9837a32