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:5ac8fc273357d34092bb94d5e9af447ff2892bde1d730c1ee19ccb1cf5c71a7bRUN
URLREF_FLAX=https://github.com/google/flax.git#0e5a0b4de10c352c8374d2b611ad77f64ec9afc6 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE= BUILD_PATH_JAXLIB=/opt/jaxlibs NVIDIA_PRODUCT_NAME=JAX NVIDIA_JAX_VERSION=24.10 pip-finalize.sh
10/23/2024 6:31 AM UTC
sha256:cedd8b01320cd2093f64fef00918012a642b5012962363dd2ace45f3ea61b81eRUN
RUN |9 URLREF_FLAX=https://github.com/google/flax.git#0e5a0b4de10c352c8374d2b611ad77f64ec9afc6 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE= BUILD_PATH_JAXLIB=/opt/jaxlibs NVIDIA_PRODUCT_NAME=JAX NVIDIA_JAX_VERSION=24.10 /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
10/23/2024 6:28 AM UTC
sha256:81fa6e6a3a3cd736c79aede44074742f47c55af0b87f09ebfd3ef107537df1dcCOPY
/opt/transformer-engine /opt/transformer-engine
10/23/2024 6:28 AM UTC
sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine
10/23/2024 6:28 AM UTC
sha256:a3ed95caeb02ffe68cdd9fd84406680ae93d633cb16422d00e8a7c22955b46d4ENV
NVTE_FRAMEWORK=jax
10/23/2024 6:28 AM UTC
sha256:ca6a82e67c2a3092bc18e8ffb7e9aceba2eccba167e5332ed55e1625ec0d9316RUN
RUN |9 URLREF_FLAX=https://github.com/google/flax.git#0e5a0b4de10c352c8374d2b611ad77f64ec9afc6 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE= BUILD_PATH_JAXLIB=/opt/jaxlibs NVIDIA_PRODUCT_NAME=JAX NVIDIA_JAX_VERSION=24.10 /bin/sh -c <<"EOF" bash -ex git-clone.sh ${URLREF_FLAX} ${SRC_PATH_FLAX} sed -i "s|optax|optax>=0.2.3|" ${SRC_PATH_FLAX}/pyproject.toml sed -i "s|orbax-checkpoint|orbax-checkpoint<0.7.0|" ${SRC_PATH_FLAX}/pyproject.toml echo "-e file://${SRC_PATH_FLAX}" >> /opt/pip-tools.d/requirements-flax.in EOF # buildkit
10/23/2024 6:28 AM UTC
sha256:23adc001197f5d4252dc446309f95fb4d7be27de5bb250deeed1adb064a51dc7RUN
RUN |9 URLREF_FLAX=https://github.com/google/flax.git#0e5a0b4de10c352c8374d2b611ad77f64ec9afc6 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE= BUILD_PATH_JAXLIB=/opt/jaxlibs NVIDIA_PRODUCT_NAME=JAX NVIDIA_JAX_VERSION=24.10 /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}" >> /opt/pip-tools.d/requirements-jax.in echo "numpy<2.0.0" >> /opt/pip-tools.d/requirements-jax.in EOF # buildkit
10/23/2024 6:28 AM UTC
sha256:4f4fb700ef54461cfa02571ae0db9a0dc1e0cdb5577484a6d75e68dc38e8acc1RUN
URLREF_FLAX=https://github.com/google/flax.git#0e5a0b4de10c352c8374d2b611ad77f64ec9afc6 SRC_PATH_JAX=/opt/jax SRC_PATH_XLA=/opt/xla SRC_PATH_FLAX=/opt/flax SRC_PATH_TRANSFORMER_ENGINE=/opt/transformer-engine BUILD_DATE= BUILD_PATH_JAXLIB=/opt/jaxlibs NVIDIA_PRODUCT_NAME=JAX NVIDIA_JAX_VERSION=24.10 mkdir -p /opt/pip-tools.d
10/23/2024 6:28 AM UTC
sha256:9d496a8c9300b54d205b7ca3d2070abbdf965f2ee531a975e4a5cb62cbe8c61aADD
build-jax.sh local_cuda_arch test-jax.sh /usr/local/bin/
10/23/2024 6:28 AM UTC
sha256:395393b03ab41619d8bb8e37ede865685ea56fdb9e41d82bbc63500efffa99caCOPY
/opt/manifest.d/git-clone.yaml /opt/manifest.d/git-clone.yaml
10/23/2024 6:28 AM UTC
...