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:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
NVIDIA_JAX_VERSION=26.01
01/21/2026 9:39 PM UTC
sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
NVIDIA_PRODUCT_NAME=JAX
01/21/2026 9:39 PM UTC
sha256:6cbe5ff9e3087a48fb8b9d98f9c5f672cc1c333f3d26546f045a4bbc2202f829RUN
URLREF_FLAX=https://github.com/google/flax.git#v0.12.1 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2026-01-21 BUILD_PATH_JAXLIB=/opt/jaxlibs pip-finalize.sh
01/21/2026 9:39 PM UTC
sha256:1397cb8bd2c8a37058dd849a8e27d48499ae1e689be70c1d4de9485d0771f9c5COPY
pinned/26.01-devel/jax-amd64.txt /opt/pip-tools.d/requirements-pinned.txt
01/21/2026 9:38 PM UTC
sha256:210eac57dfb2cdb5509ad438d4e84c766cae8ae6c24850bc1c80688422192ec1RUN
RUN |7 URLREF_FLAX=https://github.com/google/flax.git#v0.12.1 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2026-01-21 BUILD_PATH_JAXLIB=/opt/jaxlibs /bin/sh -c <<"EOF" bash -ex ls ${SRC_PATH_TRANSFORMER_ENGINE}/dist/*.whl echo "transformer-engine @ file://$(ls ${SRC_PATH_TRANSFORMER_ENGINE}/dist/*.whl)" > /opt/pip-tools.d/requirements-te.in EOF # buildkit
01/21/2026 9:35 PM UTC
sha256:dcee13354a33d58b56d782ff65226491f231ad6f6d22a838f43ff8289ed53729COPY
/opt/transformer-engine /opt/transformer-engine
01/21/2026 9:35 PM UTC
sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine
01/21/2026 9:35 PM UTC
sha256:3bb3a705e8793f7b8a25668d960183f5dd5f72db7844304eb9cc266aac88aa31RUN
RUN |7 URLREF_FLAX=https://github.com/google/flax.git#v0.12.1 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2026-01-21 BUILD_PATH_JAXLIB=/opt/jaxlibs /bin/sh -c <<"EOF" bash -ex git-clone.sh ${URLREF_FLAX} ${SRC_PATH_FLAX} sed -i 's/jax>=0\.8\.1/jax/' ${SRC_PATH_FLAX}/pyproject.toml echo "-e file://${SRC_PATH_FLAX}" >> /opt/pip-tools.d/requirements-flax.in EOF # buildkit
01/21/2026 9:35 PM UTC
sha256:f4d0023bf0713c8046d2045a02ec9535651079212b4564d76a7a25d089052b80RUN
RUN |7 URLREF_FLAX=https://github.com/google/flax.git#v0.12.1 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2026-01-21 BUILD_PATH_JAXLIB=/opt/jaxlibs /bin/sh -c <<"EOF" bash -ex for component in $(ls ${BUILD_PATH_JAXLIB}); do echo "-e file://${BUILD_PATH_JAXLIB}/${component}" >> /opt/pip-tools.d/requirements-jax.in; done echo "-e file://${SRC_PATH_JAX}[k8s]" >> /opt/pip-tools.d/requirements-jax.in echo "kubernetes @ git+https://github.com/kubernetes-client/python.git#release-34.0" >> /opt/pip-tools.d/requirements-jax.in echo "urllib3>2.6.0" >> /opt/pip-tools.d/requirements-jax.in EOF # buildkit
01/21/2026 9:35 PM UTC
sha256:4f4fb700ef54461cfa02571ae0db9a0dc1e0cdb5577484a6d75e68dc38e8acc1RUN
URLREF_FLAX=https://github.com/google/flax.git#v0.12.1 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2026-01-21 BUILD_PATH_JAXLIB=/opt/jaxlibs mkdir -p /opt/pip-tools.d
01/21/2026 9:35 PM UTC
...