DrJAX - Differentiable MapReduce Primitives in JAX
DrJAX is a library designed to embed a MapReduce programming model into JAX. DrJAX has multiple objectives.
- Create a simple JAX-based authoring surface for MapReduce computations.
- Leverage JAX's sharding mechanisms to enable highly optimized execution of MapReduce computations, especially in large-scale datacenter settings.
- 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)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| 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
|