FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Merge pull request #682 from Kaggle/add-jax · feibyte/docker-python@b18cb97 · GitHub

Commit b18cb97

Browse files
authored
Merge pull request Kaggle#682 from Kaggle/add-jax
Add JAX package
2 parents b017e4e + 5dc100d commit b18cb97

2 files changed

Lines changed: 31 additions & 0 deletions

File tree

‎gpu.Dockerfile‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,15 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
4444
ln -s /usr/local/cuda/lib64/stubs/libcuda.so /usr/local/cuda/lib64/stubs/libcuda.so.1 && \
4545
/tmp/clean-layer.sh
4646

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+
4756
# Reinstall packages with a separate version for GPU support.
4857
COPY --from=tensorflow_whl /tmp/tensorflow_gpu/*.whl /tmp/tensorflow_gpu/
4958
RUN pip uninstall -y tensorflow && \

‎tests/test_jax.py‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff 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)

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL