| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
2 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -44,6 +44,15 @@ RUN apt-get update && apt-get install -y --no-install-recommends \ | |||
| 44 | 44 | ln -s /usr/local/cuda/lib64/stubs/libcuda.so /usr/local/cuda/lib64/stubs/libcuda.so.1 && \ | |
| 45 | 45 | /tmp/clean-layer.sh | |
| 46 | 46 | ||
| 47 | + # Install JAX | ||
| 48 | + ENV JAX_PYTHON_VERSION=cp36 | ||
| 49 | + ENV JAX_CUDA_VERSION=cuda100 | ||
| 50 | + ENV JAX_PLATFORM=linux_x86_64 | ||
| 51 | + ENV JAX_BASE_URL="https://storage.googleapis.com/jax-releases" | ||
| 52 | + | ||
| 53 | + RUN pip install --upgrade $JAX_BASE_URL/$JAX_CUDA_VERSION/jaxlib-0.1.36-$JAX_PYTHON_VERSION-none-$JAX_PLATFORM.whl && \ | ||
| 54 | + pip install --upgrade jax | ||
| 55 | + | ||
| 47 | 56 | # Reinstall packages with a separate version for GPU support. | |
| 48 | 57 | COPY --from=tensorflow_whl /tmp/tensorflow_gpu/*.whl /tmp/tensorflow_gpu/ | |
| 49 | 58 | RUN pip uninstall -y tensorflow && \ | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,22 @@ | |||
| 1 | + import unittest | ||
| 2 | + | ||
| 3 | + import time | ||
| 4 | + | ||
| 5 | + from common import gpu_test | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + class TestJAX(unittest.TestCase): | ||
| 9 | + def tanh(self, x): | ||
| 10 | + import jax.numpy as np | ||
| 11 | + y = np.exp(-2.0 * x) | ||
| 12 | + return (1.0 - y) / (1.0 + y) | ||
| 13 | + | ||
| 14 | + @gpu_test | ||
| 15 | + def test_JAX(self): | ||
| 16 | + # importing inside the gpu-only test because these packages can't be | ||
| 17 | + # imported on the CPU image since they are not present there. | ||
| 18 | + from jax import grad, jit | ||
| 19 | + | ||
| 20 | + grad_tanh = grad(self.tanh) | ||
| 21 | + ag = grad_tanh(1.0) | ||
| 22 | + self.assertEqual(0.4199743, ag) | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments