diff --git a/pkgs/development/python-modules/clu/default.nix b/pkgs/development/python-modules/clu/default.nix new file mode 100644 index 000000000000..383565e6ea80 --- /dev/null +++ b/pkgs/development/python-modules/clu/default.nix @@ -0,0 +1,81 @@ +{ + lib, + buildPythonPackage, + fetchFromGitHub, + + # build-system + setuptools, + + # dependencies + absl-py, + etils, + flax, + jax, + jaxlib, + ml-collections, + numpy, + packaging, + typing-extensions, + wrapt, + + # tests + keras, + pytestCheckHook, + tensorflow, + tensorflow-datasets, + torch, +}: + +buildPythonPackage rec { + pname = "clu"; + version = "0.0.12"; + pyproject = true; + + src = fetchFromGitHub { + owner = "google"; + repo = "CommonLoopUtils"; + tag = "v${version}"; + hash = "sha256-ntqRz3fCXMf0EDQsddT68Mdi105ECBVQpVsStzk2kvQ="; + }; + + build-system = [ + setuptools + ]; + + dependencies = [ + absl-py + etils + flax + jax + jaxlib + ml-collections + numpy + packaging + typing-extensions + wrapt + ] + ++ etils.optional-dependencies.epath; + + pythonImportsCheck = [ "clu" ]; + + nativeCheckInputs = [ + keras + pytestCheckHook + tensorflow + tensorflow-datasets + torch + ]; + + disabledTests = [ + # AssertionError: [Chex] Assertion assert_trees_all_close failed + "test_collection_mixed_async" + ]; + + meta = { + description = "Common training loops in JAX"; + homepage = "https://github.com/google/CommonLoopUtils"; + changelog = "https://github.com/google/CommonLoopUtils/blob/${src.tag}/CHANGELOG.md"; + license = lib.licenses.asl20; + maintainers = with lib.maintainers; [ GaetanLepage ]; + }; +} diff --git a/pkgs/top-level/python-packages.nix b/pkgs/top-level/python-packages.nix index 4ef277c2d54d..ba2fe90abbf5 100644 --- a/pkgs/top-level/python-packages.nix +++ b/pkgs/top-level/python-packages.nix @@ -2792,6 +2792,8 @@ self: super: with self; { cltk = callPackage ../development/python-modules/cltk { }; + clu = callPackage ../development/python-modules/clu { }; + clustershell = callPackage ../development/python-modules/clustershell { }; clx-sdk-xms = callPackage ../development/python-modules/clx-sdk-xms { };