Jax Best Practices

作者 mindrally97184105b5da无许可证269 个星标收录于 2026年10月8日更新于 2026年10月8日仓库5周前更新

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

AI 生成的概览

提供 JAX 最佳实践指导,面向高性能数值计算与机器学习。

功能
该技能提供使用 JAX 进行高性能数值计算与机器学习的专家指导。内容涵盖函数式编程核心原则、jit、vmap、grad 等关键变换、性能优化,以及 pytree、自定义 vjp/jvp、分片和检查点等常见模式。它产出的是建议与推荐,而非代码或文件。
适用场景
在编写或审查 JAX 代码、需要关于变换、性能或惯用函数式模式的指导时使用。适合解答有关 JIT 编译、向量化、自动微分和多设备分片的问题。
运行要求
无需任何工具、软件包、运行时、凭据或网络访问;不包含脚本,仅为说明性指令。

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

来源与署名

来源:mindrally/skills位于jax-best-practices提交9718410

许可证: 无许可证

内容归原作者所有。SourceWeft 从公开仓库中收录这些内容。

举报或申请下架