NVIDIA
NVIDIA
JAX
Container
NVIDIA
NVIDIA
JAX

JAX is a framework for high-performance numerical computing and machine learning research. It includes Numpy-like APIs, automatic differentiation, XLA acceleration and simple primitives for scaling across GPUs and supports an ecosystem of libraries.

LayerLabelCreated
sha256:319fd71f759a7274317b18e20ffde6662ee6641251dbcadce0e89dfe5efa4cb1RUN
pip-finalize.sh
05/08/2024 4:54 AM UTC
sha256:7aa7e79933ab1613b5ad06eae13c5659c0ac354c9bd6d64ff83f8e6c502d7f9cRUN
RUN |4 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_TE=/opt/transformer-engine BUILD_DATE= /bin/sh -c <<"EOF" bash -ex -o pipefail pip install ninja &&
  rm -rf ~/.cache/pip get-source.sh -l transformer-engine -m ${MANIFEST_FILE} pushd ${SRC_PATH_TE} ## apply release-specific patches: cherry-pick changes after the last release tag git cherry-pick 8e672ff0758033c348e263dbcd6a4b3578c01161 git cherry-pick bfe21c3d68b0a9951e5716fb520045db53419c5e git cherry-pick 7c1828f80edc1405d4ef1a7780c9e0046beab5c7 # https://github.com/NVIDIA/TransformerEngine/pull/745 ## end of apply release-specific patches python setup.py bdist_wheel &&
  rm -rf build echo "transformer-engine @ file://$(ls ${SRC_PATH_TE}/dist/*.whl)" >> /opt/pip-tools.d/requirements-te.in EOF # buildkit
05/08/2024 4:53 AM UTC
sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
SRC_PATH_TE=/opt/transformer-engine
05/08/2024 4:49 AM UTC
sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
NVTE_FRAMEWORK=jax
05/08/2024 4:49 AM UTC
sha256:031fd916525ef34cf570fdff9dcd0241093f587861b058f4836538d7dd7aea9eRUN
SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_TE=/opt/transformer-engine BUILD_DATE= get-source.sh -l flax -m ${MANIFEST_FILE} -o /opt/pip-tools.d/requirements-flax.in
05/08/2024 4:49 AM UTC
sha256:4787c074a07b36ee37b507df1a578827955d0148801b041942b0962e5b1c8f3fRUN
RUN |4 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_TE=/opt/transformer-engine BUILD_DATE= /bin/sh -c <<"EOF" bash -ex # Encourage a newer numpy so that pip's dependency resolver will allow newer
# versions of other packages that rely on newer numpy, but also include fixes
# for compatibility with newer JAX versions. e.g. chex.
echo "numpy >= 1.24.1"                                  >> /opt/pip-tools.d/requirements-jax.in
echo "-e file://${SRC_PATH_JAX}"                        >> /opt/pip-tools.d/requirements-jax.in
echo "jaxlib @ file://$(ls ${SRC_PATH_JAX}/dist/*.whl)" >> /opt/pip-tools.d/requirements-jax.in
EOF # buildkit
05/08/2024 4:49 AM UTC
sha256:4f4fb700ef54461cfa02571ae0db9a0dc1e0cdb5577484a6d75e68dc38e8acc1RUN
SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_TE=/opt/transformer-engine BUILD_DATE= mkdir -p /opt/pip-tools.d
05/08/2024 4:49 AM UTC
sha256:f3365aaa072949bc2b85bb4ce0050e592033af29bdfbe15a7003c80ffe5031bdADD
build-jax.sh local_cuda_arch test-jax.sh /usr/local/bin/
05/08/2024 4:49 AM UTC
sha256:06e67634e07219b54752a7dfd1ad7e37782be941e1d3e05cc1e79b4d78655b6aCOPY
/opt/xla /opt/xla
05/08/2024 4:49 AM UTC
sha256:1db2f76cc90f07d86aca607278acece6e358b1736c7c9bc568096051dba5863cCOPY
/opt/jax /opt/jax
05/08/2024 4:49 AM UTC
...

NVIDIA uses cookies to improve your experience on our web site. We and our third-party partners also use cookies and other tools to collect and record information you provide as well as information about your interactions with our websites for performance improvement, analytics, and to assist in marketing efforts. By clicking "Accept All", you consent to our use of cookies and other tools as described in our Cookie Policy. You can manage your cookie settings by clicking on "Manage Settings." By continuing to use this site or by clicking one of the buttons below, you agree to our Terms of Service (which contains important waivers). Please see our Privacy Policy for more information on our privacy practices.