From ca22ed51026b5199addc31e5358d5507ebda9ce6 Mon Sep 17 00:00:00 2001 From: Yu-Hang Tang Date: Mon, 1 Dec 2025 21:46:12 -0800 Subject: [PATCH 1/4] bump to CUDA 13.0 --- .../dockerfile/oss.dockerfile | 17 ++++++++--------- jax-inference-offloading/setup.py | 4 ++-- 2 files changed, 10 insertions(+), 11 deletions(-) diff --git a/jax-inference-offloading/dockerfile/oss.dockerfile b/jax-inference-offloading/dockerfile/oss.dockerfile index dbbac60b0..0dfe22684 100644 --- a/jax-inference-offloading/dockerfile/oss.dockerfile +++ b/jax-inference-offloading/dockerfile/oss.dockerfile @@ -14,7 +14,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -ARG BASE_IMAGE=nvcr.io/nvidia/cuda-dl-base:25.06-cuda12.9-devel-ubuntu24.04 +ARG BASE_IMAGE=nvcr.io/nvidia/cuda-dl-base:25.11-cuda13.0-devel-ubuntu24.04 ARG URL_JIO=https://github.com/NVIDIA/JAX-Toolbox.git ARG REF_JIO=main ARG URL_TUNIX=https://github.com/google/tunix.git @@ -76,7 +76,6 @@ EOF RUN <<"EOF" bash -ex -o pipefail mkdir -p /opt/pip-tools.d pip freeze | grep wheel >> /opt/pip-tools.d/overrides.in -echo "jax[cuda12_local]" >> /opt/pip-tools.d/requirements.in echo "-e file://${SRC_PATH_JIO}" >> /opt/pip-tools.d/requirements.in echo "-e file://${SRC_PATH_TUNIX}" >> /opt/pip-tools.d/requirements.in cat "${SRC_PATH_JIO}/examples/requirements.in" >> /opt/pip-tools.d/requirements.in @@ -90,8 +89,8 @@ FROM mealkit AS final # Finalize installation RUN <<"EOF" bash -ex -o pipefail -export PIP_INDEX_URL=https://download.pytorch.org/whl/cu129 -export PIP_EXTRA_INDEX_URL="https://flashinfer.ai/whl/cu129 https://pypi.org/simple" +export PIP_INDEX_URL=https://download.pytorch.org/whl/cu130 +export PIP_EXTRA_INDEX_URL="https://flashinfer.ai/whl/cu130 https://pypi.org/simple" pushd /opt/pip-tools.d pip-compile -o requirements.txt $(ls requirements*.in) --constraint overrides.in # remove cuda wheels from install list since the container already has them @@ -99,11 +98,11 @@ sed -i 's/^nvidia-/# nvidia-/g' requirements.txt sed -i 's/# nvidia-nvshmem/nvidia-nvshmem/g' requirements.txt pip install --no-deps --src /opt -r requirements.txt # make pip happy about the missing torch dependencies -pip-mark-installed nvidia-cublas-cu12 nvidia-cuda-cupti-cu12 \ - nvidia-cuda-nvrtc-cu12 nvidia-cuda-runtime-cu12 nvidia-cudnn-cu12 \ - nvidia-cufft-cu12 nvidia-cufile-cu12 nvidia-curand-cu12 nvidia-cusolver-cu12 \ - nvidia-cusparse-cu12 nvidia-cusparselt-cu12 nvidia-nccl-cu12 \ - nvidia-nvjitlink-cu12 nvidia-nvtx-cu12 +pip-mark-installed nvidia-cublas-cu13 nvidia-cuda-cupti-cu13 \ + nvidia-cuda-nvrtc-cu13 nvidia-cuda-runtime-cu13 nvidia-cudnn-cu13 \ + nvidia-cufft-cu13 nvidia-cufile-cu13 nvidia-curand-cu13 nvidia-cusolver-cu13 \ + nvidia-cusparse-cu13 nvidia-cusparselt-cu13 nvidia-nccl-cu13 \ + nvidia-nvjitlink-cu13 nvidia-nvtx-cu13 popd rm -rf ~/.cache/* EOF diff --git a/jax-inference-offloading/setup.py b/jax-inference-offloading/setup.py index 9bd9a89b8..421f5b867 100644 --- a/jax-inference-offloading/setup.py +++ b/jax-inference-offloading/setup.py @@ -57,13 +57,13 @@ def run(self): version='0.0.1', packages=['jax_inference_offloading'], install_requires=[ - 'cupy-cuda12x', + 'cupy-cuda13x', 'cloudpickle', 'flax', 'grpcio==1.76.*', 'protobuf==6.33.*', 'huggingface-hub', - 'jax==0.8.0', + 'jax[cuda13_local]==0.8.0', 'jaxtyping', 'kagglehub', 'vllm==0.11.2', From dd8c943cd691928a423061e5de517b1044937c97 Mon Sep 17 00:00:00 2001 From: Yu-Hang Tang Date: Mon, 1 Dec 2025 22:19:47 -0800 Subject: [PATCH 2/4] pin vllm CUDA 13.0 wheel --- jax-inference-offloading/dockerfile/oss.dockerfile | 1 + 1 file changed, 1 insertion(+) diff --git a/jax-inference-offloading/dockerfile/oss.dockerfile b/jax-inference-offloading/dockerfile/oss.dockerfile index 0dfe22684..eab3f2971 100644 --- a/jax-inference-offloading/dockerfile/oss.dockerfile +++ b/jax-inference-offloading/dockerfile/oss.dockerfile @@ -76,6 +76,7 @@ EOF RUN <<"EOF" bash -ex -o pipefail mkdir -p /opt/pip-tools.d pip freeze | grep wheel >> /opt/pip-tools.d/overrides.in +echo "vllm @ https://github.com/vllm-project/vllm/releases/download/v0.11.2/vllm-0.11.2+cu130-cp38-abi3-manylinux1_x86_64.whl" >> /opt/pip-tools.d/requirements.in echo "-e file://${SRC_PATH_JIO}" >> /opt/pip-tools.d/requirements.in echo "-e file://${SRC_PATH_TUNIX}" >> /opt/pip-tools.d/requirements.in cat "${SRC_PATH_JIO}/examples/requirements.in" >> /opt/pip-tools.d/requirements.in From bc77eec045611ea19eade91d1dbc1dc0462f2a52 Mon Sep 17 00:00:00 2001 From: Yu-Hang Tang Date: Mon, 1 Dec 2025 22:25:14 -0800 Subject: [PATCH 3/4] update base image in CI --- .github/workflows/jio.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/jio.yaml b/.github/workflows/jio.yaml index abfbfaede..7855f2c44 100644 --- a/.github/workflows/jio.yaml +++ b/.github/workflows/jio.yaml @@ -67,7 +67,7 @@ jobs: ARCHITECTURE: ${{ matrix.ARCHITECTURE }} ARTIFACT_NAME: artifact-jio-build BADGE_FILENAME: badge-jio-build - BASE_IMAGE: nvcr.io/nvidia/cuda-dl-base:25.06-cuda12.9-devel-ubuntu24.04 + BASE_IMAGE: nvcr.io/nvidia/cuda-dl-base:25.11-cuda13.0-devel-ubuntu24.04 BUILD_DATE: ${{ needs.metadata.outputs.BUILD_DATE }} CONTAINER_NAME: jio DOCKERFILE: jax-inference-offloading/dockerfile/oss.dockerfile From 2b4ba4a54650051a10ada9ad338eab1a678b6653 Mon Sep 17 00:00:00 2001 From: Yu-Hang Tang Date: Mon, 1 Dec 2025 22:33:24 -0800 Subject: [PATCH 4/4] x86 only --- .github/workflows/jio.yaml | 2 +- jax-inference-offloading/dockerfile/oss.dockerfile | 4 +++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/.github/workflows/jio.yaml b/.github/workflows/jio.yaml index 7855f2c44..b1e52820f 100644 --- a/.github/workflows/jio.yaml +++ b/.github/workflows/jio.yaml @@ -53,7 +53,7 @@ jobs: build: needs: metadata strategy: - fail-fast: true + fail-fast: false matrix: ARCHITECTURE: [amd64, arm64] runs-on: [self-hosted, "${{ matrix.ARCHITECTURE }}", "small"] diff --git a/jax-inference-offloading/dockerfile/oss.dockerfile b/jax-inference-offloading/dockerfile/oss.dockerfile index eab3f2971..9b46df3a4 100644 --- a/jax-inference-offloading/dockerfile/oss.dockerfile +++ b/jax-inference-offloading/dockerfile/oss.dockerfile @@ -76,7 +76,9 @@ EOF RUN <<"EOF" bash -ex -o pipefail mkdir -p /opt/pip-tools.d pip freeze | grep wheel >> /opt/pip-tools.d/overrides.in -echo "vllm @ https://github.com/vllm-project/vllm/releases/download/v0.11.2/vllm-0.11.2+cu130-cp38-abi3-manylinux1_x86_64.whl" >> /opt/pip-tools.d/requirements.in +if [[ $(uname -m) == "x86_64" ]]; then + echo "vllm @ https://github.com/vllm-project/vllm/releases/download/v0.11.2/vllm-0.11.2+cu130-cp38-abi3-manylinux1_x86_64.whl" >> /opt/pip-tools.d/requirements.in +fi echo "-e file://${SRC_PATH_JIO}" >> /opt/pip-tools.d/requirements.in echo "-e file://${SRC_PATH_TUNIX}" >> /opt/pip-tools.d/requirements.in cat "${SRC_PATH_JIO}/examples/requirements.in" >> /opt/pip-tools.d/requirements.in