diff --git a/pkgs/development/python-modules/jax/default.nix b/pkgs/development/python-modules/jax/default.nix index e1d21af094f9..ce9a7758bc61 100644 --- a/pkgs/development/python-modules/jax/default.nix +++ b/pkgs/development/python-modules/jax/default.nix @@ -1,39 +1,64 @@ -{ buildPythonPackage, fetchFromGitHub, lib -# propagatedBuildInputs -, absl-py, numpy, opt-einsum -# checkInputs -, jaxlib, pytestCheckHook +{ lib +, absl-py +, buildPythonPackage +, fetchFromGitHub +, jaxlib +, numpy +, opt-einsum +, pytestCheckHook +, pythonOlder +, scipy +, typing-extensions }: buildPythonPackage rec { pname = "jax"; - version = "0.2.21"; + version = "0.2.24"; + format = "setuptools"; + + disabled = pythonOlder "3.7"; - # Fetching from pypi doesn't allow us to run the test suite. See https://discourse.nixos.org/t/pythonremovetestsdir-hook-being-run-before-checkphase/14612/3. src = fetchFromGitHub { owner = "google"; repo = pname; rev = "jax-v${version}"; - sha256 = "05w157h6jv20k8w2gnmlxbycmzf24lr5v392q0c5v0qcql11q7pn"; + sha256 = "1mmn1m4mprpwqlb1smjfdy3f74zm9p3l9dhhn25x6jrcj2cgc5pi"; }; # jaxlib is _not_ included in propagatedBuildInputs because there are # different versions of jaxlib depending on the desired target hardware. The # JAX project ships separate wheels for CPU, GPU, and TPU. Currently only the # CPU wheel is packaged. - propagatedBuildInputs = [ absl-py numpy opt-einsum ]; + propagatedBuildInputs = [ + absl-py + numpy + opt-einsum + scipy + typing-extensions + ]; + + checkInputs = [ + jaxlib + pytestCheckHook + ]; - checkInputs = [ jaxlib pytestCheckHook ]; # NOTE: Don't run the tests in the expiremental directory as they require flax # which creates a circular dependency. See https://discourse.nixos.org/t/how-to-nix-ify-python-packages-with-circular-dependencies/14648/2. # Not a big deal, this is how the JAX docs suggest running the test suite # anyhow. - pytestFlagsArray = [ "-W ignore::DeprecationWarning" "tests/" ]; + pytestFlagsArray = [ + "-W ignore::DeprecationWarning" + "tests/" + ]; + + pythonImportsCheck = [ + "jax" + ]; meta = with lib; { description = "Differentiate, compile, and transform Numpy code"; - homepage = "https://github.com/google/jax"; - license = licenses.asl20; + homepage = "https://github.com/google/jax"; + license = licenses.asl20; maintainers = with maintainers; [ samuela ]; }; }