drjax
DrJAX - Scalable and Differentiable MapReduce Primitives in JAX.
Decision gist · record as of 2026-08-14
Yes, if you are building distributed or federated machine learning systems in JAX and need MapReduce-style primitives with automatic differentiation. The package is actively maintained, carries no known vulnerabilities, and has low install friction. However, it remains early-stage (first release January 2025) with modest adoption; verify that its API and performance characteristics match your production requirements before committing to it.AI-flagged interpretation of the facts on this page — verify before relying
Before you install
- Requires Python 3.9 or later; jax and jaxlib must be installed as runtime dependencies.
- Low friction installation with a pure Python wheel.
- The package is actively maintained with a recent release and no known vulnerabilities, though it remains early-stage with modest adoption.
License · maintenance · safety
permissive license (permissive) — Apache 2.0 permissive license allows commercial and derivative use with minimal restrictions; you must include license and copyright notices in distributions.
last release 2026-06-15 (60 days) · last repo commit 2026-07-08 · 19 stars
0 known vulnerabilities (OSV.dev, 2026-08-14) · 253,355 downloads/mo, #8,521 on PyPI
Alternatives
Verify before relying
pip install drjax
import drjax
import jax
# Define a MapReduce computation using DrJAX primitives
# and execute it with JAX's sharding mechanisms- Whether the package's MapReduce API is stable or subject to breaking changes in future releases.
- Performance characteristics and scalability limits in typical datacenter deployments.
- Availability of tutorials or examples beyond the research paper for practical adoption.
What it is and what it does
DrJAX is a JAX library that brings MapReduce-style distributed computing into the JAX ecosystem. It provides a simple authoring surface for parallel computations while leveraging JAX's sharding mechanisms for optimized execution in large-scale settings. The package is designed around three core goals: offering an accessible way to write MapReduce computations in JAX, enabling efficient execution via JAX's built-in parallelization, and maintaining full differentiability throughout—including through communication primitives like broadcasts and reductions.
The library embeds patterns similar to those in TensorFlow Federated but uses JAX's mapping and extension capabilities. It is tailored for scenarios involving large models and distributed training where both correctness and efficiency matter. The package depends on absl-py, jax, and jaxlib, and requires Python 3.9 or later. It is actively maintained and carries no known security vulnerabilities.
Use it for
- Implement distributed training loops for large models using MapReduce patterns with automatic differentiation.
- Build federated learning workflows where computations are distributed across multiple devices or nodes.
- Optimize parallel data processing pipelines in datacenters by leveraging JAX's sharding for efficient communication.
- Prototype distributed algorithms that require differentiable reductions and broadcasts across compute clusters.
- Combine MapReduce-style data aggregation with JAX's automatic differentiation for gradient-based distributed optimization.
Worth the install?
AI-flagged interpretation of the facts on this page. Verify before relying on it.
Yes, if you are building distributed or federated machine learning systems in JAX and need MapReduce-style primitives with automatic differentiation.
The package is actively maintained, carries no known vulnerabilities, and has low install friction. However, it remains early-stage (first release January 2025) with modest adoption; verify that its API and performance characteristics match your production requirements before committing to it.
Install
drjax on PyPI
Before you install
Low friction installation with a pure Python wheel. The package is actively maintained with a recent release and no known vulnerabilities, though it remains early-stage with modest adoption.
Requires Python 3.9 or later; jax and jaxlib must be installed as runtime dependencies.
License in practice
Apache 2.0 permissive license allows commercial and derivative use with minimal restrictions; you must include license and copyright notices in distributions.
Quickstart
pip install drjax
import drjax
import jax
# Define a MapReduce computation using DrJAX primitives
# and execute it with JAX's sharding mechanisms
Verify before relying
- Whether the package's MapReduce API is stable or subject to breaking changes in future releases.
- Performance characteristics and scalability limits in typical datacenter deployments.
- Availability of tutorials or examples beyond the research paper for practical adoption.
Package facts
| License | permissive license permissive |
| Python support | Supports the current Python release >=3.9 |
| Install friction | Low. Pure-Python wheel |
| Runtime dependencies | 3 packagesabsl-pyjaxjaxlib |
| Maintenance | Actively maintained 60 days since the last release |
| Last repo commit | |
| First released | |
| Downloads | 253,355 / month, #8,521 on PyPI 30-day window, as of 2026-08-14 |
| Known vulnerabilities | None known OSV.dev, checked 2026-08-14 |
| Classifiers | Intended Audience :: Science/ResearchLicense :: OSI Approved :: Apache Software License |
Evidence: drjax-0.2.0-py3-none-any.whl
Tags
Let your AI agent find packages like this
Example. Real query, live index.
You found this page by searching. An agent finds it by wishing: SkillFed indexes 14,416 PyPI packages by what they can do, searchable in plain language.
wish › “mapreduce jax”
- drjaxDrJAX embeds MapReduce programming primitives into JAX, enabling…
- ytsaurus-clientProvides a Python client library for YTsaurus, enabling programmatic…
- equinoxEquinox provides neural network and model building on top of JAX with…
Give your agent the search over MCP, or paste the wish link into any chat.
More Distributed Computing packages
gRPC Python is an HTTP/2-based RPC framework that enables you to define and call remote procedures across network boundaries using protocol buffers for serialization.
Install it if you need RPC communication in a distributed system or are integrating with existing gRPC services.
execnet lets you spawn and communicate with Python interpreters across local processes, remote hosts, and different platforms, using a simple API for task distribution and inter-process messaging.
However, the aging maintenance status (275 days since last release) means you should verify it meets your concurrency and performance needs before committing to a…
Cloudpickle extends Python's standard pickle module to serialize lambda functions, interactively-defined functions and classes, and other constructs that the default pickle cannot handle, making it suitable for cluster computing and remote code execution.
Install it if you need to serialize lambda functions, interactively-defined code, or non-standard Python constructs for cluster computing or distributed execution.
Provides a unified, open()-compatible Python API for streaming large files from remote storage (S3, GCS, Azure, HDFS, SFTP, HTTP) and local filesystems, with transparent compression support.
Install it if you work with large files on cloud storage or remote systems and want to avoid writing boilerplate around multiple SDKs.
Portalocker provides cross-platform file locking with support for exclusive and shared locks, plus Redis-based distributed locks and process-aware PID file locking.
Install it if you need file or process coordination; the optional extras (pywin32, redis) are only required for specific lock types.
Ray is a distributed computing framework that scales Python applications from a single machine to multi-node clusters, providing abstractions for parallel tasks, stateful actors, and shared objects.
See also rax · jax · clu · jaxellip · jaxlib · jaxlie · drjit · jax-cuda12-plugin · jax-cuda12-pjrt · rucio-clients