summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorSandro <sandro.jaeckel@gmail.com>2024-11-25 15:16:57 +0100
committerGitHub <noreply@github.com>2024-11-25 15:16:57 +0100
commitd1af4a2b819b5f263a4659be2a9cb045b655d84a (patch)
tree2294ff6c45bdf1e73b9b12c8d0fa02d465a156c3
parentvivaldi: 7.0.3495.6 -> 7.0.3495.18 (#358331) (diff)
parentpython3Packages.stable-baselines3: init at 2.3.2 (diff)
downloadnixpkgs-d1af4a2b819b5f263a4659be2a9cb045b655d84a.tar.gz
python3Packages.stable-baselines3: init at 2.3.2 (#355954)
-rw-r--r--pkgs/development/python-modules/stable-baselines3/default.nix88
-rw-r--r--pkgs/top-level/python-packages.nix2
2 files changed, 90 insertions, 0 deletions
diff --git a/pkgs/development/python-modules/stable-baselines3/default.nix b/pkgs/development/python-modules/stable-baselines3/default.nix
new file mode 100644
index 000000000000..60f36cd6c050
--- /dev/null
+++ b/pkgs/development/python-modules/stable-baselines3/default.nix
@@ -0,0 +1,88 @@
+{
+ lib,
+ ale-py,
+ buildPythonPackage,
+ cloudpickle,
+ fetchFromGitHub,
+ gymnasium,
+ matplotlib,
+ numpy,
+ opencv4,
+ pandas,
+ pillow,
+ psutil,
+ pygame,
+ pytestCheckHook,
+ pythonOlder,
+ rich,
+ setuptools,
+ tensorboard,
+ torch,
+ tqdm,
+}:
+buildPythonPackage rec {
+ pname = "stable-baselines3";
+ version = "2.3.2-unstable-2024-11-04";
+ pyproject = true;
+
+ disabled = pythonOlder "3.8";
+
+ src = fetchFromGitHub {
+ owner = "DLR-RM";
+ repo = "stable-baselines3";
+ # commit with updated dependencies since gymnasium is not compatible with the latest release:
+ # https://github.com/DLR-RM/stable-baselines3/pull/1837
+ rev = "8f0b488bc5a897f1ac2b95f493bcb6b7e92d311c";
+ hash = "sha256-zhmNZ86lowFJKes3i/TBBBsO8ZMuUUQsphQ98IsmHd4=";
+ };
+
+ pythonRelaxDeps = true;
+
+ build-system = [ setuptools ];
+
+ dependencies = [
+ ale-py
+ cloudpickle
+ gymnasium
+ matplotlib
+ numpy
+ opencv4
+ pandas
+ pillow
+ psutil
+ pygame
+ rich
+ tensorboard
+ torch
+ tqdm
+ ];
+
+ nativeCheckInputs = [
+ pytestCheckHook
+ torch
+ ];
+
+ pythonImportsCheck = [ "stable_baselines3" ];
+
+ disabledTestPaths = [
+ # Tests starts training a model, which takes too long
+ "tests/test_cnn.py"
+ "tests/test_dict_env.py"
+ "tests/test_her.py"
+ "tests/test_save_load.py"
+ ];
+
+ disabledTests = [
+ # Tests that attempt to access the filesystem
+ "test_make_atari_env"
+ "test_vec_env_monitor_kwargs"
+ ];
+
+ meta = {
+ description = "PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms";
+ homepage = "https://github.com/DLR-RM/stable-baselines3";
+ # changelog = "https://github.com/DLR-RM/stable-baselines3/releases/tag/v${version}";
+ license = lib.licenses.mit;
+ maintainers = with lib.maintainers; [ derdennisop ];
+ };
+}
diff --git a/pkgs/top-level/python-packages.nix b/pkgs/top-level/python-packages.nix
index 20baacb12c38..eaa610cde265 100644
--- a/pkgs/top-level/python-packages.nix
+++ b/pkgs/top-level/python-packages.nix
@@ -15241,6 +15241,8 @@ self: super: with self; {
ssort = callPackage ../development/python-modules/ssort { };
+ stable-baselines3 = callPackage ../development/python-modules/stable-baselines3 { };
+
stack-data = callPackage ../development/python-modules/stack-data { };
stamina = callPackage ../development/python-modules/stamina { };