| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 11f01ce commit a54fc96
2 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,8 +1,6 @@ | |||
| 1 | 1 | ARG BASE_IMAGE_TAG | |
| 2 | - ARG LIBTPU_IMAGE_TAG | ||
| 3 | 2 | ARG TENSORFLOW_VERSION | |
| 4 | 3 | ||
| 5 | - FROM gcr.io/cloud-tpu-v2-images/libtpu:${LIBTPU_IMAGE_TAG} as libtpu | ||
| 6 | 4 | FROM gcr.io/kaggle-images/python-tpu-tensorflow-whl:python-${BASE_IMAGE_TAG}-${TENSORFLOW_VERSION} AS tensorflow_whl | |
| 7 | 5 | FROM gcr.io/kaggle-images/python:${BASE_IMAGE_TAG} | |
| 8 | 6 | ||
@@ -12,20 +10,29 @@ ARG TORCH_VERSION | |||
| 12 | 10 | ||
| 13 | 11 | ENV ISTPUVM=1 | |
| 14 | 12 | ||
| 15 | - COPY --from=libtpu /libtpu.so /lib | ||
| 16 | - | ||
| 17 | 13 | COPY --from=tensorflow_whl /tmp/tensorflow_pkg/tensorflow*.whl /tmp/tensorflow_pkg/ | |
| 18 | 14 | RUN pip install /tmp/tensorflow_pkg/tensorflow*.whl && \ | |
| 19 | 15 | rm -rf /tmp/tensorflow_pkg && \ | |
| 20 | 16 | /tmp/clean-layer.sh | |
| 21 | 17 | ||
| 18 | + # LIBTPU installed here: | ||
| 19 | + ENV DEFAULT_LIBTPU=/opt/conda/lib/python3.7/site-packages/libtpu/libtpu.so | ||
| 20 | + ENV PYTORCH_LIBTPU=/opt/conda/lib/python3.7/site-packages/libtpu/torch-libtpu.so | ||
| 21 | + ENV JAX_LIBTPU=/opt/conda/lib/python3.7/site-packages/libtpu/jax-libtpu.so | ||
| 22 | + | ||
| 22 | 23 | # https://cloud.google.com/tpu/docs/pytorch-xla-ug-tpu-vm#changing_pytorch_version | |
| 23 | 24 | RUN pip uninstall -y torch && \ | |
| 24 | 25 | pip install torch==${TORCH_VERSION} && \ | |
| 25 | 26 | # The URL doesn't include patch version. i.e. must use 1.11 instead of 1.11.0 | |
| 26 | 27 | pip install torch_xla[tpuvm] -f https://storage.googleapis.com/tpu-pytorch/wheels/tpuvm/torch_xla-${TORCH_VERSION%.*}-cp37-cp37m-linux_x86_64.whl && \ | |
| 28 | + cp $DEFAULT_LIBTPU $PYTORCH_LIBTPU && \ | ||
| 27 | 29 | /tmp/clean-layer.sh | |
| 28 | 30 | ||
| 29 | 31 | # https://cloud.google.com/tpu/docs/jax-quickstart-tpu-vm#install_jax_on_your_cloud_tpu_vm | |
| 30 | 32 | RUN pip install "jax[tpu]>=0.2.16" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html && \ | |
| 33 | + cp $DEFAULT_LIBTPU $JAX_LIBTPU && \ | ||
| 31 | 34 | /tmp/clean-layer.sh | |
| 35 | + | ||
| 36 | + # Monkey-patch JAX & PYTORCH to load the correct libtpu.so when they are imported: | ||
| 37 | + RUN sed -i "s|^\(\(.*\)libtpu.configure_library_path.*\)|\1\n\2os.environ['TPU_LIBRARY_PATH'] = '${PYTORCH_LIBTPU}'|" /opt/conda/lib/python3.7/site-packages/torch_xla/__init__.py && \ | ||
| 38 | + sed -i "s|^\(\(.*\)libtpu.configure_library_path.*\)|\1\n\2os.environ['TPU_LIBRARY_PATH'] = '${JAX_LIBTPU}'|" /opt/conda/lib/python3.7/site-packages/jax/_src/cloud_tpu_init.py | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,4 @@ | |||
| 1 | 1 | # TODO(b/213335159): Use ci-pretest for BASE_IMAGE_TAG once stable. | |
| 2 | - BASE_IMAGE_TAG=v108 | ||
| 3 | - LIBTPU_IMAGE_TAG=libtpu_1.1.0_RC00 | ||
| 2 | + BASE_IMAGE_TAG=v115 | ||
| 4 | 3 | TENSORFLOW_VERSION=2.8.0 | |
| 5 | 4 | TORCH_VERSION=1.11.0 | |
| Back | FazBrowse Home | New Git URL |
0 commit comments