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 從公開儲存庫中收錄這些內容。

檢舉或申請下架