NVIDIA
JAX
Container
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=25.04
    04/16/2025 9:44 PM UTC
    sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
    NVIDIA_PRODUCT_NAME=JAX
    04/16/2025 9:44 PM UTC
    sha256:12cc372d0b90f31da27a380961fe6a3dbde604b1dcd276ddd827c0c000f59525RUN
    URLREF_FLAX=https://github.com/google/flax.git#v0.10.5 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2025-04-16 BUILD_PATH_JAXLIB=/opt/jaxlibs pip-finalize.sh
    04/16/2025 9:44 PM UTC
    sha256:b0c7621c7577f60feee0d061232b6c8cd6d382db48480c38449313ab0deda8e7RUN
    RUN |7 URLREF_FLAX=https://github.com/google/flax.git#v0.10.5 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2025-04-16 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
    04/16/2025 9:40 PM UTC
    sha256:4411a344e09730bbaf183f13367a6cb8457c7fa3501a829c6a7cf08c0118dc71COPY
    /opt/transformer-engine /opt/transformer-engine
    04/16/2025 9:40 PM UTC
    sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
    SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine
    04/16/2025 9:40 PM UTC
    sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
    NVTE_FRAMEWORK=jax
    04/16/2025 9:40 PM UTC
    sha256:ab92178abe126ed44b238cf86a28757aca9abfc247e7819a221935f892b0aa5dRUN
    RUN |7 URLREF_FLAX=https://github.com/google/flax.git#v0.10.5 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2025-04-16 BUILD_PATH_JAXLIB=/opt/jaxlibs /bin/sh -c <<"EOF" bash -ex git-clone.sh ${URLREF_FLAX} ${SRC_PATH_FLAX} echo "-e file://${SRC_PATH_FLAX}" >> /opt/pip-tools.d/requirements-flax.in EOF # buildkit
    04/16/2025 9:40 PM UTC
    sha256:b78660fffbe1945f095579e2651739e8dc6bb2ad3f4cb459ed0c54a6d92fcf16RUN
    RUN |7 URLREF_FLAX=https://github.com/google/flax.git#v0.10.5 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2025-04-16 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 EOF # buildkit
    04/16/2025 9:40 PM UTC
    sha256:4f4fb700ef54461cfa02571ae0db9a0dc1e0cdb5577484a6d75e68dc38e8acc1RUN
    URLREF_FLAX=https://github.com/google/flax.git#v0.10.5 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE=2025-04-16 BUILD_PATH_JAXLIB=/opt/jaxlibs mkdir -p /opt/pip-tools.d
    04/16/2025 9:40 PM UTC
    ...