depot/third_party/nixpkgs/pkgs/development/python-modules/jax/test-cuda.nix

18 lines
285 B
Nix
Raw Normal View History

{ jax
, jaxlib
, pkgs
}:
pkgs.writers.writePython3Bin "jax-test-cuda" { libraries = [ jax jaxlib ]; } ''
import jax
from jax import random
assert jax.devices()[0].platform == "gpu"
rng = random.PRNGKey(0)
x = random.normal(rng, (100, 100))
x @ x
print("success!")
''