عن المشروع
توفر JAX نظاماً لتحويلات الدوال القابلة للتركيب لبرامج Python و NumPy. وقد صُممت لتوسيع نطاق الحوسبة العددية عبر مسرعات الأجهزة بما في ذلك GPUs و TPUs باستخدام XLA (Accelerated Linear Algebra) للتجميع.
تشمل القدرات الرئيسية ما يلي:
- التفاضل التلقائي: باستخدام `jax.grad` ، تدعم التفاضل بنمط العكس ونمط الأمام من خلال حلقات Python، والتفريعات، والتكرار، والإغلاقات، مما يسمح باشتقاقات من رتب أعلى.
- التجميع في الوقت المناسب (JIT): باستخدام `jax.jit` ، تقوم بتجميع الدوال النقية عبر XLA لتحسين أداء التنفيذ على المسرعات.
- التوجيه التلقائي (Auto-vectorization): باستخدام `jax.vmap` ، تقوم بتعيين الدوال على طول محاور المصفوفة عن طريق دفع الحلقات إلى العمليات الأولية، مما يلغي الحاجة إلى الإدارة اليدوية لأبعاد الدفعات.
- التوسع والتوازي: تدعم أوضاع توسع مختلفة بما في ذلك التوازي التلقائي القائم على المجمع، والتقسيم الصريح (sharding)، والبرمجة اليدوية لكل جهاز باستخدام المجموعات الصريحة.
تتوافق JAX مع منصات متعددة بما في ذلك Linux و macOS و Windows، مع دعم لـ CPU و NVIDIA GPU و Google TPU و AMD GPU، ودعم تجريبي لـ Apple و Intel GPUs.