| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -7,14 +7,14 @@ | |||
| 7 | 7 | ||
| 8 | 8 | class TestJAX(unittest.TestCase): | |
| 9 | 9 | def tanh(self, x): | |
| 10 | + import jax.numpy as np | ||
| 10 | 11 | y = np.exp(-2.0 * x) | |
| 11 | 12 | return (1.0 - y) / (1.0 + y) | |
| 12 | 13 | ||
| 13 | 14 | @gpu_test | |
| 14 | 15 | def test_JAX(self): | |
| 15 | 16 | # importing inside the gpu-only test because these packages can't be | |
| 16 | 17 | # imported on the CPU image since they are not present there. | |
| 17 | - import jax.numpy as np | ||
| 18 | 18 | from jax import grad, jit | |
| 19 | 19 | ||
| 20 | 20 | grad_tanh = grad(self.tanh) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments