Skip to main content

DrJAX - Differentiable MapReduce Primitives in JAX

DrJAX is a library designed to embed a MapReduce programming model into JAX. DrJAX has multiple objectives.

  1. Create a simple JAX-based authoring surface for MapReduce computations.
  2. Leverage JAX's sharding mechanisms to enable highly optimized execution of MapReduce computations, especially in large-scale datacenter settings.
  3. Full differentiability of DrJAX computations, including differentiating through communication primitives like broadcasts and reductions.

DrJAX is designed to make it easy to author and execute parallel computations in the datacenter. DrJAX is tailored towards large-scale parallel and distributed computations, including computations involving larger models, and ensuring that they can be run efficiently. DrJAX embeds primitives like those defined by TensorFlow Federated using the mapping capabilities and primitive extensions of JAX.

System design

For details on DrJAX's system design, check out our paper.

Citation

To cite this repository, please use the following BibTeX citation:

@inproceedings{rush2024drjax,
  title={DrJAX: Scalable and Differentiable MapReduce Primitives in JAX},
  author={Rush, J Keith and Charles, Zachary and Garrett, Zachary and Augenstein, Sean and Mitchell, Nicole Elyse},
  booktitle={2nd Workshop on Advancing Neural Network Training: Computational Efficiency, Scalability, and Resource Optimization (WANT@ ICML 2024)}
}

Disclaimers

This is not an officially supported Google product.

If you're interested in learning more about responsible AI practices, please see Google AI's Responsible AI Practices.

Dataset Grouper is Apache 2.0 licensed. See the LICENSE file.

Metadata

Release files for drjax 0.2.1

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Built distribution (wheel)

Table of built distributions (wheels) for drjax 0.2.1
File Interpreter ABI Platform
drjax-0.2.1-py3-none-any.whl Python 3 none any Details

Release files / drjax-0.2.1-py3-none-any.whl

Download URL drjax-0.2.1-py3-none-any.whl
Size 24.7 kB
Tags Python 3
SHA-256 checksum
How to use checksums
9ec93541a7ba0af864d599cd3e14770a19e291a6291b8b79ddf6b92e8b173494
BLAKE2b-256 checksum
How to use checksums
280b77703d77dd2a87675044ae15d5e8e8abccdac6add6dd5a1907e7ef0e7028
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/7.0.0 CPython/3.13.14

Release history Release notifications | RSS feed

This release

0.2.1 This release

1 release file

0.2.0

1 release file

0.1.4

1 release file

0.1.3

1 release file

0.1.2

1 release file

0.1.1

1 release file

0.1.0

1 release file

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page