agilerl 2.7.0.dev0__tar.gz → 2.7.0.dev1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- agilerl-2.7.0.dev1/.github/workflows/linux-tests.yml +96 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/workflows/macos-tests.yml +10 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/workflows/windows-tests.yml +10 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.gitignore +2 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.pre-commit-config.yaml +2 -2
- agilerl-2.7.0.dev1/CLAUDE.md +139 -0
- agilerl-2.7.0.dev1/DQN_LEARNING_ALGORITHM_ANALYSIS.md +309 -0
- agilerl-2.7.0.dev1/DQN_LEARNING_ANALYSIS.md +168 -0
- agilerl-2.7.0.dev1/GPU_CLEANUP_ANALYSIS.md +541 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/PKG-INFO +1 -1
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/__init__.py +13 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/__init__.py +5 -1
- agilerl-2.7.0.dev1/agilerl/algorithms/cispo.py +26 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/base.py +2096 -839
- agilerl-2.7.0.dev1/agilerl/algorithms/core/fused_lora.py +146 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/optimizer_wrapper.py +112 -14
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/registry.py +4 -3
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/dpo.py +63 -103
- agilerl-2.7.0.dev1/agilerl/algorithms/grpo.py +1076 -0
- agilerl-2.7.0.dev1/agilerl/algorithms/gspo.py +26 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ippo.py +0 -43
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/maddpg.py +123 -42
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/matd3.py +149 -55
- agilerl-2.7.0.dev1/agilerl/algorithms/ppo_llm.py +905 -0
- agilerl-2.7.0.dev1/agilerl/algorithms/reinforce_llm.py +729 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/sft.py +25 -18
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/hpo/mutation.py +1 -1
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/hpo/tournament.py +2 -2
- agilerl-2.7.0.dev1/agilerl/llm_envs/__init__.py +37 -0
- agilerl-2.7.0.dev1/agilerl/llm_envs/base.py +261 -0
- agilerl-2.7.0.dev1/agilerl/llm_envs/preference.py +135 -0
- agilerl-2.7.0.dev1/agilerl/llm_envs/reasoning.py +163 -0
- agilerl-2.7.0.dev1/agilerl/llm_envs/search.py +120 -0
- agilerl-2.7.0.dev1/agilerl/llm_envs/sft.py +99 -0
- agilerl-2.7.0.dev1/agilerl/llm_envs/sync_vec_env.py +273 -0
- agilerl-2.7.0.dev1/agilerl/llm_envs/token_observation.py +349 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/protocols.py +24 -4
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/rollouts/on_policy.py +66 -1
- agilerl-2.7.0.dev1/agilerl/training/train_llm.py +1867 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/typing.py +15 -4
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/algo_utils.py +84 -20
- agilerl-2.7.0.dev1/agilerl/utils/llm_utils.py +867 -0
- agilerl-2.7.0.dev1/agilerl/utils/ppo_value_head.py +275 -0
- agilerl-2.7.0.dev1/agilerl/utils/probe_envs_llm.py +154 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/utils.py +435 -68
- agilerl-2.7.0.dev1/agilerl/wrappers/llm_envs.py +19 -0
- agilerl-2.7.0.dev1/benchmarking/benchmarking_llm_multiturn.py +120 -0
- agilerl-2.7.0.dev0/benchmarking/benchmarking_dpo.py → agilerl-2.7.0.dev1/benchmarking/benchmarking_llm_preference.py +3 -3
- agilerl-2.7.0.dev0/benchmarking/benchmarking_grpo.py → agilerl-2.7.0.dev1/benchmarking/benchmarking_llm_reasoning.py +36 -57
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_sft.py +2 -2
- agilerl-2.7.0.dev1/configs/training/llm_finetuning/cispo.yaml +38 -0
- {agilerl-2.7.0.dev0/configs/training → agilerl-2.7.0.dev1/configs/training/llm_finetuning}/grpo.yaml +4 -3
- agilerl-2.7.0.dev1/configs/training/llm_finetuning/grpo_multiturn.yaml +30 -0
- agilerl-2.7.0.dev1/configs/training/llm_finetuning/gspo.yaml +36 -0
- agilerl-2.7.0.dev1/configs/training/llm_finetuning/ppo_llm.yaml +47 -0
- agilerl-2.7.0.dev1/configs/training/llm_finetuning/reinforce_llm.yaml +42 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/config_load.py +22 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/configs/grpo_constant_target.yaml +24 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/configs/grpo_grid_navigation.yaml +26 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_conditional_target.yaml +31 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_constant_target.yaml +27 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_grid_navigation.yaml +29 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_multi_input.yaml +28 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_value_head.yaml +30 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm.py +186 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_stage_1.py +252 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_stage_2.py +283 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_stage_3.py +400 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_training_matrix.py +697 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/debugging_value.py +190 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/llm_debug_utils.py +16 -0
- agilerl-2.7.0.dev1/demos/llm/debugging/tiny_model.py +229 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/llm}/demo_llm_finetuning.py +6 -6
- agilerl-2.7.0.dev1/docs/api/algorithms/cispo.rst +137 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/dpo.rst +1 -2
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/grpo.rst +5 -2
- agilerl-2.7.0.dev1/docs/api/algorithms/gspo.rst +134 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/index.rst +30 -0
- agilerl-2.7.0.dev1/docs/api/algorithms/llmppo.rst +147 -0
- agilerl-2.7.0.dev1/docs/api/algorithms/llmreinforce.rst +145 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/sft.rst +1 -1
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/llm_utils.rst +1 -1
- agilerl-2.7.0.dev1/docs/api/wrappers/llm_envs.rst +12 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/debugging_rl/index.rst +61 -1
- agilerl-2.7.0.dev1/docs/llm_finetuning/llm_checkpoints.rst +148 -0
- agilerl-2.7.0.dev1/find_dqn_commit.sh +82 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/pyproject.toml +4 -3
- agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/multiturn:GRPO:pop1:no_tournament/attributes.pt +0 -0
- agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/multiturn:LLMPPO:pop1:no_tournament/attributes.pt +0 -0
- agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/multiturn:LLMReinforce:pop1:no_tournament/attributes.pt +0 -0
- agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/preference:DPO:pop1:no_tournament/attributes.pt +0 -0
- agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/reasoning:GRPO:pop1:no_tournament/attributes.pt +0 -0
- agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/reasoning:LLMPPO:pop1:no_tournament/attributes.pt +0 -0
- agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/reasoning:LLMReinforce:pop1:no_tournament/attributes.pt +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/conftest.py +16 -1
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/subprocess_runner.py +14 -7
- agilerl-2.7.0.dev1/tests/test_algorithms/conftest.py +38 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_core_base.py +1819 -472
- agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/conftest.py +115 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_llms/test_dpo.py +23 -123
- agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_fused_lora.py +126 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_llms/test_grpo.py +1130 -1109
- agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_ppo_llm.py +1094 -0
- agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_reinforce_llm.py +944 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_llms/test_sft.py +15 -11
- agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_vllm.py +135 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/test_maddpg.py +152 -57
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/test_matd3.py +183 -83
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_optimizer_wrapper.py +209 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_ilql.py +6 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_data.py +15 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_hpo/test_mutation.py +4 -2
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_hpo/test_tournament.py +70 -2
- agilerl-2.7.0.dev1/tests/test_rollouts/test_on_policy.py +383 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_train/test_train_llm.py +1323 -82
- agilerl-2.7.0.dev1/tests/test_utils/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_algo_utils.py +386 -2
- agilerl-2.7.0.dev1/tests/test_utils/test_llm_utils.py +1257 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_minari_utils.py +6 -2
- agilerl-2.7.0.dev1/tests/test_utils/test_ppo_value_head.py +484 -0
- agilerl-2.7.0.dev1/tests/test_utils/test_probe_envs_llm.py +268 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_utils.py +280 -23
- agilerl-2.7.0.dev1/tests/test_wrappers/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_llm_envs.py +13 -2
- agilerl-2.7.0.dev1/tests/test_wrappers/test_multiturn_wrappers.py +859 -0
- agilerl-2.7.0.dev1/tests/test_wrappers/test_ppo_test_method.py +79 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/utils.py +57 -1
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/uv.lock +1982 -1895
- agilerl-2.7.0.dev0/.github/workflows/linux-tests.yml +0 -55
- agilerl-2.7.0.dev0/agilerl/algorithms/grpo.py +0 -674
- agilerl-2.7.0.dev0/agilerl/training/train_llm.py +0 -1080
- agilerl-2.7.0.dev0/agilerl/utils/llm_utils.py +0 -252
- agilerl-2.7.0.dev0/agilerl/wrappers/llm_envs.py +0 -795
- agilerl-2.7.0.dev0/docs/api/wrappers/llm_envs.rst +0 -12
- agilerl-2.7.0.dev0/tests/test_algorithms/test_llms/conftest.py +0 -152
- agilerl-2.7.0.dev0/tests/test_rollouts/test_on_policy.py +0 -149
- agilerl-2.7.0.dev0/tests/test_utils/test_llm_utils.py +0 -370
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/ISSUE_TEMPLATE/bug_report.md +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/ISSUE_TEMPLATE/feature_request.md +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/PULL_REQUEST_TEMPLATE/pull_request_template.md +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/badges/arena-github-badge.svg +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/codeql/install_codeql.sh +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/codeql/run_codeql.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1/.github}/dependabot.yml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/workflows/codeql.yml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.readthedocs.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/CITATION.cff +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/CODE_OF_CONDUCT.md +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/CONTRIBUTING.md +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/LICENSE +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/README.md +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/bc_lm.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/cqn.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ddpg.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/dqn.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/dqn_rainbow.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ilql.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/neural_ts_bandit.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/neural_ucb_bandit.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ppo.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/td3.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/data.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/multi_agent_replay_buffer.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/replay_buffer.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/rollout_buffer.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/sampler.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/segment_tree.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/language_environment.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/rl_data.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/tokenizer.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/torch_datasets.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/hpo/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/base.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/bert.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/cnn.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/configs.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/custom_components.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/dummy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/gpt.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/lstm.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/mlp.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/multi_input.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/resnet.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/simba.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/actors.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/base.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/custom_modules.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/distributions.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/q_networks.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/value_networks.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/rollouts/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_bandits.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_multi_agent_off_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_multi_agent_on_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_off_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_offline.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_on_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/cache.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/evolvable_networks.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/ilql_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/log_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/minari_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/probe_envs.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/probe_envs_ma.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/sampling_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/torch_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/vector/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/vector/pz_async_vec_env.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/vector/pz_vec_env.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/agent.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/learning.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/make_evolvable.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/pettingzoo_wrappers.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_bandits.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_multi_agent_off_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_multi_agent_on_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_off_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_off_policy_distributed.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_offline.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_offline_distributed.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_on_policy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_rainbow.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_recurrent.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_resnet.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_simba.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/configs/ds_config.json +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/make_evolvable_benchmarking.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/networks.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/accelerate/accelerate.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/accelerate/grpo_accelerate_config.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/bandit/neural_ts.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/bandit/neural_ucb.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/cqn.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ddpg/ddpg.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ddpg/ddpg_lstm.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ddpg/ddpg_simba.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/dqn/dqn.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/dqn/dqn_lstm.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/dqn/dqn_rainbow.yaml +0 -0
- {agilerl-2.7.0.dev0/configs/training → agilerl-2.7.0.dev1/configs/training/llm_finetuning}/dpo.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/ippo.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/ippo_pong.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/maddpg.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/matd3.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_input.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ppo/ppo.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ppo/ppo_image.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ppo/ppo_recurrent.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/sft.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/td3.yaml +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/cartpole/cartpole_random_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/cartpole/cartpole_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/pendulum/pendulum_random_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/pendulum/pendulum_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev0/tests → agilerl-2.7.0.dev1/debugging}/__init__.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/bandits}/demo_bandit.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/multi_agent}/demo_multi_agent.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_custom_network.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_off_policy.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_off_policy_distributed.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_offline.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_offline_distributed.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy_rnn_cartpole.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy_rnn_memory.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy_rnn_minigrid.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_cartpole.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_lunar_lander.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_lunar_lander_rnn.py +0 -0
- {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_rnn_memory.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/Makefile +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/arena-github-badge.svg +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/css/custom.css +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/favicon.ico +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/js/expand_sidebar.js +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/logo_teal.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/logo_white.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/module.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/network.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/iris-thumbnail.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/pendigits-thumbnail.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/rainbow_performance.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/simba_thumbnail.png +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/base.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/cql.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ddpg.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/dqn.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/dqn_rainbow.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ilql.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ippo.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/maddpg.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/matd3.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/neural_ts.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/neural_ucb.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ppo.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/registry.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/td3.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/wrappers.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/data.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/multi_agent_replay_buffer.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/replay_buffer.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/rollout_buffer.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/sampler.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/segment_tree.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/hpo/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/hpo/mutation.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/hpo/tournament.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/base.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/bert.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/cnn.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/custom_activation.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/dummy.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/gpt.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/lstm.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/mlp.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/multi_input.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/resnet.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/simba.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/actors.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/base.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/q_networks.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/value_networks.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/rollouts/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/rollouts/on_policy.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/train.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/algo_utils.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/cache.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/evolvable_networks.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/ilql_utils.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/log_utils.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/minari_utils.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/probe_envs.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/torch_utils.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/utils.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/vector/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/vector/petting_zoo_async_vector_env.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/vector/petting_zoo_vector_env.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/agent.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/learning.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/make_evolvable.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/pettingzoo.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/bandits/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/conf.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/custom_algorithms/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/distributed_training/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/evo_hyperparam_opt/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/evolvable_networks/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/get_started/agilerl2changes.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/get_started/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/llm_finetuning/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/make.bat +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/multi_agent_training/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/off_policy/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/offline_training/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/on_policy/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/pomdp/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/releases/index.rst +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/requirements.txt +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/sitecustomize.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_algorithms → agilerl-2.7.0.dev1/tests}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/helper_functions.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/pz_vector_test_utils.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_algorithms/test_bandits → agilerl-2.7.0.dev1/tests/test_algorithms}/__init__.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_algorithms/test_llms → agilerl-2.7.0.dev1/tests/test_algorithms/test_bandits}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_bandits/test_neural_ts.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_bandits/test_neural_ucb.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_base.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_bc_lm.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_algorithms/test_multi_agent → agilerl-2.7.0.dev1/tests/test_algorithms/test_llms}/__init__.py +0 -0
- /agilerl-2.7.0.dev0/tests/test_algorithms/test_single_agent/__init__.py → /agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_llm_checkpoint.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_components → agilerl-2.7.0.dev1/tests/test_algorithms/test_multi_agent}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/conftest.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/test_ippo.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_registry.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_hpo → agilerl-2.7.0.dev1/tests/test_algorithms/test_single_agent}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_cqn.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_ddpg.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_dqn.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_dqn_rainbow.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_ppo.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_td3.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_modules → agilerl-2.7.0.dev1/tests/test_components}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_multi_agent_replay_buffer.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_replay_buffer.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_replay_data.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_rollout_buffer.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_sampler.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_segment_tree.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_networks → agilerl-2.7.0.dev1/tests/test_hpo}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_init.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_utils → agilerl-2.7.0.dev1/tests/test_modules}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_base.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_bert.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_cnn.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_configs.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_custom_activation.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_dummy.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_gpt.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_lstm.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_mlp.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_multi_input.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_resnet.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_simba.py +0 -0
- {agilerl-2.7.0.dev0/tests/test_wrappers → agilerl-2.7.0.dev1/tests/test_networks}/__init__.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_actors.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_base.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_distributions.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_q_networks.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_value_functions.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_protocols.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_train/test_train.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_cache.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_ilql_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_log_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_probe_envs.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_probe_envs_ma.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_sampling_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_torch_utils.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_utils_evolvable.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_vector/test_vector.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_agent.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_autoreset.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_bandit_env.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_make_evolvable.py +0 -0
- {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_skills.py +0 -0
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
---
|
|
2
|
+
name: Linux
|
|
3
|
+
|
|
4
|
+
on:
|
|
5
|
+
push:
|
|
6
|
+
branches: [main, nightly]
|
|
7
|
+
paths:
|
|
8
|
+
- agilerl/**
|
|
9
|
+
- tests/**
|
|
10
|
+
- .github/workflows/**
|
|
11
|
+
- pyproject.toml
|
|
12
|
+
pull_request:
|
|
13
|
+
paths:
|
|
14
|
+
- agilerl/**
|
|
15
|
+
- tests/**
|
|
16
|
+
- .github/workflows/**
|
|
17
|
+
- pyproject.toml
|
|
18
|
+
|
|
19
|
+
concurrency:
|
|
20
|
+
group: ${{ github.workflow }}-${{ github.head_ref || github.ref }}
|
|
21
|
+
cancel-in-progress: true
|
|
22
|
+
|
|
23
|
+
permissions:
|
|
24
|
+
contents: read
|
|
25
|
+
|
|
26
|
+
jobs:
|
|
27
|
+
tests:
|
|
28
|
+
runs-on: gha-runner-scale-set
|
|
29
|
+
strategy:
|
|
30
|
+
fail-fast: false
|
|
31
|
+
max-parallel: 4
|
|
32
|
+
matrix:
|
|
33
|
+
python-version: ['3.10', '3.11', '3.12', '3.13']
|
|
34
|
+
|
|
35
|
+
container:
|
|
36
|
+
image: pytorch/pytorch:2.7.1-cuda12.6-cudnn9-devel
|
|
37
|
+
options: --user root
|
|
38
|
+
|
|
39
|
+
# Workspace (/__w) is ~1GB with little free space; root (/) has plenty. Put cache and venv on /.
|
|
40
|
+
env:
|
|
41
|
+
UV_CACHE_DIR: /tmp/uv-cache
|
|
42
|
+
UV_PROJECT_ENVIRONMENT: /tmp/agilerl-venv
|
|
43
|
+
HF_HOME: /tmp/hf-cache
|
|
44
|
+
TORCHINDUCTOR_CACHE_DIR: /tmp/inductor-cache
|
|
45
|
+
|
|
46
|
+
steps:
|
|
47
|
+
- uses: actions/checkout@v4
|
|
48
|
+
- uses: astral-sh/setup-uv@v7
|
|
49
|
+
with:
|
|
50
|
+
enable-cache: true
|
|
51
|
+
python-version: ${{ matrix.python-version }}
|
|
52
|
+
|
|
53
|
+
- name: Cache HuggingFace models
|
|
54
|
+
uses: actions/cache@v4
|
|
55
|
+
with:
|
|
56
|
+
path: /tmp/hf-cache
|
|
57
|
+
key: hf-${{ matrix.python-version }}-${{ hashFiles('pyproject.toml') }}
|
|
58
|
+
restore-keys: hf-${{ matrix.python-version }}-
|
|
59
|
+
|
|
60
|
+
- name: Cache torch inductor compilations
|
|
61
|
+
uses: actions/cache@v4
|
|
62
|
+
with:
|
|
63
|
+
path: /tmp/inductor-cache
|
|
64
|
+
key: inductor-${{ matrix.python-version }}-${{ hashFiles('pyproject.toml') }}
|
|
65
|
+
restore-keys: inductor-${{ matrix.python-version }}-
|
|
66
|
+
|
|
67
|
+
- name: Install dependencies
|
|
68
|
+
# swig is needed to build box2d-py from source (no pre-built wheels for py3.10+).
|
|
69
|
+
run: |
|
|
70
|
+
uv sync --locked --all-groups --extra all
|
|
71
|
+
echo "$UV_PROJECT_ENVIRONMENT/bin" >> $GITHUB_PATH
|
|
72
|
+
|
|
73
|
+
- name: Reset coverage data
|
|
74
|
+
run: rm -f .coverage .coverage.*
|
|
75
|
+
|
|
76
|
+
- name: Run non-LLM tests (parallel)
|
|
77
|
+
run: uv run pytest -m "not llm" --exitfirst --cov=agilerl --cov-report= --durations=0 --durations-min=1.0
|
|
78
|
+
|
|
79
|
+
- name: Run LLM tests (sequential, exclusive GPU)
|
|
80
|
+
run: uv run pytest -n0 -m "llm" --exitfirst --cov=agilerl --cov-append --cov-report= --durations=0 --durations-min=1.0
|
|
81
|
+
|
|
82
|
+
- name: Combine coverage data
|
|
83
|
+
run: |
|
|
84
|
+
if ls .coverage.* >/dev/null 2>&1; then
|
|
85
|
+
uv run coverage combine
|
|
86
|
+
else
|
|
87
|
+
echo "No parallel coverage shards found; skipping combine."
|
|
88
|
+
fi
|
|
89
|
+
|
|
90
|
+
- name: Generate coverage report
|
|
91
|
+
run: uv run coverage xml -o coverage.xml
|
|
92
|
+
|
|
93
|
+
- name: Upload coverage reports to Codecov
|
|
94
|
+
uses: codecov/codecov-action@v3
|
|
95
|
+
env:
|
|
96
|
+
CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}
|
|
@@ -6,10 +6,20 @@ on:
|
|
|
6
6
|
branches: [main, nightly]
|
|
7
7
|
paths:
|
|
8
8
|
- agilerl/**
|
|
9
|
+
- tests/**
|
|
10
|
+
- .github/workflows/**
|
|
11
|
+
- pyproject.toml
|
|
9
12
|
pull_request:
|
|
10
13
|
branches: [main, nightly]
|
|
11
14
|
paths:
|
|
12
15
|
- agilerl/**
|
|
16
|
+
- tests/**
|
|
17
|
+
- .github/workflows/**
|
|
18
|
+
- pyproject.toml
|
|
19
|
+
|
|
20
|
+
concurrency:
|
|
21
|
+
group: ${{ github.workflow }}-${{ github.head_ref || github.ref }}
|
|
22
|
+
cancel-in-progress: true
|
|
13
23
|
|
|
14
24
|
permissions:
|
|
15
25
|
contents: read
|
|
@@ -6,10 +6,20 @@ on:
|
|
|
6
6
|
branches: [main, nightly]
|
|
7
7
|
paths:
|
|
8
8
|
- agilerl/**
|
|
9
|
+
- tests/**
|
|
10
|
+
- .github/workflows/**
|
|
11
|
+
- pyproject.toml
|
|
9
12
|
pull_request:
|
|
10
13
|
branches: [main, nightly]
|
|
11
14
|
paths:
|
|
12
15
|
- agilerl/**
|
|
16
|
+
- tests/**
|
|
17
|
+
- .github/workflows/**
|
|
18
|
+
- pyproject.toml
|
|
19
|
+
|
|
20
|
+
concurrency:
|
|
21
|
+
group: ${{ github.workflow }}-${{ github.head_ref || github.ref }}
|
|
22
|
+
cancel-in-progress: true
|
|
13
23
|
|
|
14
24
|
permissions:
|
|
15
25
|
contents: read
|
|
@@ -32,7 +32,7 @@ repos:
|
|
|
32
32
|
- --skip=*.css,*.js,*.map,*.scss,*.svg
|
|
33
33
|
- --ignore-words-list=magent,pres,roate
|
|
34
34
|
- repo: https://github.com/astral-sh/ruff-pre-commit
|
|
35
|
-
rev: v0.15.
|
|
35
|
+
rev: v0.15.11
|
|
36
36
|
hooks:
|
|
37
37
|
- id: ruff
|
|
38
38
|
name: Ruff Linter
|
|
@@ -47,7 +47,7 @@ repos:
|
|
|
47
47
|
|
|
48
48
|
- repo: https://github.com/astral-sh/uv-pre-commit
|
|
49
49
|
# uv version.
|
|
50
|
-
rev: 0.11.
|
|
50
|
+
rev: 0.11.7
|
|
51
51
|
hooks:
|
|
52
52
|
- id: uv-lock
|
|
53
53
|
|
|
@@ -0,0 +1,139 @@
|
|
|
1
|
+
# CLAUDE.md
|
|
2
|
+
|
|
3
|
+
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
|
|
4
|
+
|
|
5
|
+
## Project Overview
|
|
6
|
+
|
|
7
|
+
AgileRL is a deep reinforcement learning library focused on RLOps - MLOps for reinforcement learning. The key innovation is evolutionary hyperparameter optimization (HPO) that automatically tunes hyperparameters during training, reducing the need for separate HPO experiments.
|
|
8
|
+
|
|
9
|
+
## Development Commands
|
|
10
|
+
|
|
11
|
+
```bash
|
|
12
|
+
# Install in development mode
|
|
13
|
+
pip install -e .
|
|
14
|
+
|
|
15
|
+
# Install with LLM dependencies (for GRPO, DPO algorithms)
|
|
16
|
+
pip install -e ".[all]"
|
|
17
|
+
|
|
18
|
+
# Run all tests
|
|
19
|
+
pytest tests
|
|
20
|
+
|
|
21
|
+
# Run tests with coverage
|
|
22
|
+
pytest --cov=agilerl --cov-report=xml
|
|
23
|
+
|
|
24
|
+
# Run a single test file
|
|
25
|
+
pytest tests/test_algorithms/test_dqn.py
|
|
26
|
+
|
|
27
|
+
# Run a specific test
|
|
28
|
+
pytest tests/test_algorithms/test_dqn.py::test_dqn_init -v
|
|
29
|
+
|
|
30
|
+
# Skip LLM-related tests (faster)
|
|
31
|
+
pytest -m "not llm"
|
|
32
|
+
|
|
33
|
+
# Run tests with fail-fast
|
|
34
|
+
pytest --exitfirst
|
|
35
|
+
|
|
36
|
+
# Lint and format (via pre-commit)
|
|
37
|
+
pre-commit run --all-files
|
|
38
|
+
|
|
39
|
+
# Install pre-commit hooks
|
|
40
|
+
pre-commit install
|
|
41
|
+
```
|
|
42
|
+
|
|
43
|
+
## Architecture
|
|
44
|
+
|
|
45
|
+
### Core Algorithm Hierarchy
|
|
46
|
+
|
|
47
|
+
All RL algorithms inherit from a common base in `agilerl/algorithms/core/base.py`:
|
|
48
|
+
|
|
49
|
+
- **EvolvableAlgorithm**: Base metaclass providing mutation registry, checkpointing, cloning, and evolutionary features
|
|
50
|
+
- **RLAlgorithm(EvolvableAlgorithm)**: Single-agent RL base with observation preprocessing and action handling
|
|
51
|
+
- **MultiAgentRLAlgorithm(EvolvableAlgorithm)**: Multi-agent RL base supporting agent-specific networks
|
|
52
|
+
|
|
53
|
+
### Evolvable Networks System
|
|
54
|
+
|
|
55
|
+
The evolution system uses two core concepts in `agilerl/modules/`:
|
|
56
|
+
|
|
57
|
+
- **EvolvableModule** (`modules/base.py`): Base class for neural network modules that support architectural mutations (add/remove layers, nodes, change activations). Uses `@mutation` decorator to mark methods that trigger network reconstruction.
|
|
58
|
+
- **EvolvableNetwork** (`networks/base.py`): Encoder-head architecture combining feature extraction with task-specific heads. Mutations can target encoder or head independently.
|
|
59
|
+
|
|
60
|
+
Network configurations are dataclasses in `agilerl/modules/configs.py` (MlpNetConfig, CnnNetConfig, etc.).
|
|
61
|
+
|
|
62
|
+
### HPO Components (`agilerl/hpo/`)
|
|
63
|
+
|
|
64
|
+
- **Mutations** (`mutation.py`): Applies architectural and hyperparameter mutations to agent populations
|
|
65
|
+
- **TournamentSelection** (`tournament.py`): Selects fittest agents based on evaluation scores
|
|
66
|
+
|
|
67
|
+
### Training Loop Functions (`agilerl/training/`)
|
|
68
|
+
|
|
69
|
+
Pre-built training loops that handle the full evolutionary training cycle:
|
|
70
|
+
- `train_off_policy`: DQN, DDPG, TD3, Rainbow
|
|
71
|
+
- `train_on_policy`: PPO
|
|
72
|
+
- `train_multi_agent_off_policy`: MADDPG, MATD3
|
|
73
|
+
- `train_multi_agent_on_policy`: IPPO
|
|
74
|
+
- `train_bandits`: NeuralUCB, NeuralTS
|
|
75
|
+
- `train_offline`: CQL
|
|
76
|
+
- `train_llm`: GRPO, DPO
|
|
77
|
+
|
|
78
|
+
### Registry System (`agilerl/algorithms/core/registry.py`)
|
|
79
|
+
|
|
80
|
+
Algorithms declare their evolvable components via `MutationRegistry`:
|
|
81
|
+
- **NetworkGroup**: Groups related networks (e.g., actor/critic, eval/target)
|
|
82
|
+
- **OptimizerConfig**: Maps optimizers to networks with hyperparameter bounds
|
|
83
|
+
- **HyperparameterConfig**: Defines mutable RL hyperparameters (lr, batch_size, etc.)
|
|
84
|
+
|
|
85
|
+
### Protocol Types (`agilerl/protocols.py`)
|
|
86
|
+
|
|
87
|
+
Runtime-checkable protocols define interfaces for:
|
|
88
|
+
- `EvolvableAlgorithmProtocol`: Required methods for evolutionary algorithms
|
|
89
|
+
- `EvolvableModuleProtocol`: Required methods for evolvable network modules
|
|
90
|
+
- `EvolvableNetworkProtocol`: Encoder-head network interface
|
|
91
|
+
|
|
92
|
+
## Key Patterns
|
|
93
|
+
|
|
94
|
+
### Creating Populations
|
|
95
|
+
|
|
96
|
+
```python
|
|
97
|
+
from agilerl.utils.utils import create_population
|
|
98
|
+
|
|
99
|
+
pop = create_population(
|
|
100
|
+
algo="DQN",
|
|
101
|
+
observation_space=obs_space,
|
|
102
|
+
action_space=act_space,
|
|
103
|
+
INIT_HP=init_hp_dict,
|
|
104
|
+
population_size=6,
|
|
105
|
+
device="cuda"
|
|
106
|
+
)
|
|
107
|
+
```
|
|
108
|
+
|
|
109
|
+
### Mutation Flow
|
|
110
|
+
|
|
111
|
+
1. Tournament selection picks winners from population
|
|
112
|
+
2. Clone winners to replace losers
|
|
113
|
+
3. Apply mutations via `Mutations.mutation()`:
|
|
114
|
+
- Network architecture mutations (layers, nodes)
|
|
115
|
+
- Hyperparameter mutations (lr, batch_size)
|
|
116
|
+
4. Rebuild optimizers for mutated networks
|
|
117
|
+
|
|
118
|
+
### Multi-Agent Support
|
|
119
|
+
|
|
120
|
+
Uses PettingZoo-style parallel API. Agent networks can be:
|
|
121
|
+
- **Homogeneous**: Shared architecture across agents
|
|
122
|
+
- **Heterogeneous**: Per-agent architectures via `ModuleDict`
|
|
123
|
+
|
|
124
|
+
## Testing Conventions
|
|
125
|
+
|
|
126
|
+
- Test files mirror source structure: `tests/test_<module>/test_<file>.py`
|
|
127
|
+
- Common fixtures in `tests/conftest.py` (spaces, networks, configs)
|
|
128
|
+
- Helper functions in `tests/helper_functions.py`
|
|
129
|
+
- LLM tests marked with `@pytest.mark.llm`
|
|
130
|
+
- Tests use session-scoped fixtures for spaces to avoid recreation overhead
|
|
131
|
+
|
|
132
|
+
## Branch Workflow
|
|
133
|
+
|
|
134
|
+
- `main`: Stable releases
|
|
135
|
+
- `nightly`: Active development (PRs target this branch)
|
|
136
|
+
|
|
137
|
+
## Linting
|
|
138
|
+
|
|
139
|
+
Uses Ruff for linting and formatting. Configuration in `pyproject.toml` with relaxed rules for test files.
|
|
@@ -0,0 +1,309 @@
|
|
|
1
|
+
# DQN Learning Algorithm Analysis
|
|
2
|
+
|
|
3
|
+
## Overview
|
|
4
|
+
Detailed analysis of the DQN learning algorithm implementation, focusing on the `learn()`, `update()`, and `soft_update()` methods.
|
|
5
|
+
|
|
6
|
+
## Algorithm Flow
|
|
7
|
+
|
|
8
|
+
### 1. `learn()` Method (lines 338-359)
|
|
9
|
+
|
|
10
|
+
```python
|
|
11
|
+
def learn(self, experiences: ExperiencesType) -> float:
|
|
12
|
+
obs = experiences["obs"]
|
|
13
|
+
actions = experiences["action"]
|
|
14
|
+
rewards = experiences["reward"]
|
|
15
|
+
next_obs = experiences["next_obs"]
|
|
16
|
+
dones = experiences["done"]
|
|
17
|
+
|
|
18
|
+
obs = self.preprocess_observation(obs)
|
|
19
|
+
next_obs = self.preprocess_observation(next_obs)
|
|
20
|
+
|
|
21
|
+
loss = self.update(obs, actions, rewards, next_obs, dones)
|
|
22
|
+
|
|
23
|
+
# soft update target network
|
|
24
|
+
self.soft_update()
|
|
25
|
+
return loss.item()
|
|
26
|
+
```
|
|
27
|
+
|
|
28
|
+
**Analysis**: ✅ Looks correct
|
|
29
|
+
- Extracts experiences correctly
|
|
30
|
+
- Preprocesses observations
|
|
31
|
+
- Calls `update()` to compute loss and backpropagate
|
|
32
|
+
- Calls `soft_update()` after each learning step
|
|
33
|
+
- Returns scalar loss value
|
|
34
|
+
|
|
35
|
+
### 2. `update()` Method (lines 286-336)
|
|
36
|
+
|
|
37
|
+
```python
|
|
38
|
+
def update(self, obs, actions, rewards, next_obs, dones) -> torch.Tensor:
|
|
39
|
+
with torch.no_grad():
|
|
40
|
+
if self.double: # Double Q-learning
|
|
41
|
+
q_idx = self.actor(next_obs).argmax(dim=1).unsqueeze(1)
|
|
42
|
+
q_target = (
|
|
43
|
+
self.actor_target(next_obs).gather(dim=1, index=q_idx).detach()
|
|
44
|
+
)
|
|
45
|
+
else:
|
|
46
|
+
q_target = self.actor_target(next_obs).max(axis=1)[0].unsqueeze(1)
|
|
47
|
+
|
|
48
|
+
# target, if terminal then y_j = rewards
|
|
49
|
+
y_j = rewards + self.gamma * q_target * (1 - dones)
|
|
50
|
+
|
|
51
|
+
if actions.ndim == 1:
|
|
52
|
+
actions = actions.unsqueeze(-1)
|
|
53
|
+
|
|
54
|
+
# Compute Q-values for actions taken and loss
|
|
55
|
+
q_eval = self.actor(obs).gather(1, actions.long())
|
|
56
|
+
loss: torch.Tensor = self.criterion(q_eval, y_j)
|
|
57
|
+
|
|
58
|
+
# zero gradients, perform a backward pass, and update the weights
|
|
59
|
+
self.optimizer.zero_grad()
|
|
60
|
+
if self.accelerator is not None:
|
|
61
|
+
self.accelerator.backward(loss)
|
|
62
|
+
else:
|
|
63
|
+
loss.backward()
|
|
64
|
+
|
|
65
|
+
self.optimizer.step()
|
|
66
|
+
return loss.detach()
|
|
67
|
+
```
|
|
68
|
+
|
|
69
|
+
## Issues Found
|
|
70
|
+
|
|
71
|
+
### ⚠️ Issue 1: Inconsistent `max()` Usage (Line 316)
|
|
72
|
+
|
|
73
|
+
**Problem**: Uses `axis=1` instead of `dim=1`
|
|
74
|
+
|
|
75
|
+
```python
|
|
76
|
+
q_target = self.actor_target(next_obs).max(axis=1)[0].unsqueeze(1)
|
|
77
|
+
```
|
|
78
|
+
|
|
79
|
+
**Impact**:
|
|
80
|
+
- PyTorch's `max()` accepts `axis` but it's deprecated
|
|
81
|
+
- Should use `dim=1` for consistency
|
|
82
|
+
- **However**: This shouldn't prevent learning, just causes deprecation warning
|
|
83
|
+
|
|
84
|
+
**Comparison**:
|
|
85
|
+
- Line 311: Uses `.argmax(dim=1)` ✅ (correct)
|
|
86
|
+
- Line 316: Uses `.max(axis=1)` ❌ (should be `dim=1`)
|
|
87
|
+
|
|
88
|
+
**Fix**:
|
|
89
|
+
```python
|
|
90
|
+
q_target = self.actor_target(next_obs).max(dim=1)[0].unsqueeze(1)
|
|
91
|
+
```
|
|
92
|
+
|
|
93
|
+
### ⚠️ Issue 2: Target Network Initialization Method
|
|
94
|
+
|
|
95
|
+
**Problem**: DQN uses a complex TensorDict-based initialization via `init_hook()`, while other algorithms use simple `load_state_dict()`
|
|
96
|
+
|
|
97
|
+
**DQN Approach** (lines 185-203):
|
|
98
|
+
```python
|
|
99
|
+
def init_hook(self) -> None:
|
|
100
|
+
param_vals: TensorDict = from_module(self.actor).detach()
|
|
101
|
+
target_params: TensorDict = param_vals.clone().lock_()
|
|
102
|
+
try:
|
|
103
|
+
target_params.to_module(self.actor_target)
|
|
104
|
+
except KeyError:
|
|
105
|
+
pass
|
|
106
|
+
finally:
|
|
107
|
+
self.param_vals = param_vals
|
|
108
|
+
self.target_params = target_params
|
|
109
|
+
```
|
|
110
|
+
|
|
111
|
+
**RainbowDQN/CQN Approach**:
|
|
112
|
+
```python
|
|
113
|
+
self.actor_target.load_state_dict(self.actor.state_dict())
|
|
114
|
+
```
|
|
115
|
+
|
|
116
|
+
**Potential Issues**:
|
|
117
|
+
1. The `lock_()` creates a locked TensorDict that's detached from computation graph
|
|
118
|
+
2. If `to_module()` fails silently (caught by `except KeyError: pass`), target network might not be initialized
|
|
119
|
+
3. The locked TensorDict might interfere with `soft_update()` parameter updates
|
|
120
|
+
|
|
121
|
+
**Impact**:
|
|
122
|
+
- If `to_module()` fails, target network starts with random weights instead of copying from actor
|
|
123
|
+
- This would cause incorrect Q-targets and prevent learning
|
|
124
|
+
- The silent exception handling makes this hard to detect
|
|
125
|
+
|
|
126
|
+
**Recommendation**: Add logging or assertion to verify target network is initialized:
|
|
127
|
+
```python
|
|
128
|
+
def init_hook(self) -> None:
|
|
129
|
+
param_vals: TensorDict = from_module(self.actor).detach()
|
|
130
|
+
target_params: TensorDict = param_vals.clone().lock_()
|
|
131
|
+
try:
|
|
132
|
+
target_params.to_module(self.actor_target)
|
|
133
|
+
except KeyError as e:
|
|
134
|
+
# Log the error instead of silently passing
|
|
135
|
+
warnings.warn(f"Failed to initialize target network: {e}. Using load_state_dict fallback.")
|
|
136
|
+
self.actor_target.load_state_dict(self.actor.state_dict())
|
|
137
|
+
finally:
|
|
138
|
+
self.param_vals = param_vals
|
|
139
|
+
self.target_params = target_params
|
|
140
|
+
```
|
|
141
|
+
|
|
142
|
+
### ⚠️ Issue 3: Missing Gradient Clipping
|
|
143
|
+
|
|
144
|
+
**Problem**: DQN doesn't clip gradients, while RainbowDQN does
|
|
145
|
+
|
|
146
|
+
**RainbowDQN** (line 442):
|
|
147
|
+
```python
|
|
148
|
+
clip_grad_norm_(self.actor.parameters(), 10.0)
|
|
149
|
+
```
|
|
150
|
+
|
|
151
|
+
**DQN**: No gradient clipping
|
|
152
|
+
|
|
153
|
+
**Impact**:
|
|
154
|
+
- Could lead to gradient explosion in some cases
|
|
155
|
+
- Not necessarily a bug, but could cause instability
|
|
156
|
+
|
|
157
|
+
**Recommendation**: Consider adding gradient clipping:
|
|
158
|
+
```python
|
|
159
|
+
from torch.nn.utils import clip_grad_norm_
|
|
160
|
+
|
|
161
|
+
# After loss.backward(), before optimizer.step()
|
|
162
|
+
clip_grad_norm_(self.actor.parameters(), max_norm=10.0)
|
|
163
|
+
self.optimizer.step()
|
|
164
|
+
```
|
|
165
|
+
|
|
166
|
+
### ✅ Correct Implementations
|
|
167
|
+
|
|
168
|
+
1. **Q-Learning Update Formula** (line 319): ✅ Correct
|
|
169
|
+
```python
|
|
170
|
+
y_j = rewards + self.gamma * q_target * (1 - dones)
|
|
171
|
+
```
|
|
172
|
+
|
|
173
|
+
2. **Double Q-Learning** (lines 310-314): ✅ Correct
|
|
174
|
+
- Uses actor to select action, target to evaluate
|
|
175
|
+
|
|
176
|
+
3. **Loss Computation** (line 326): ✅ Correct
|
|
177
|
+
```python
|
|
178
|
+
q_eval = self.actor(obs).gather(1, actions.long())
|
|
179
|
+
loss = self.criterion(q_eval, y_j)
|
|
180
|
+
```
|
|
181
|
+
|
|
182
|
+
4. **Gradient Flow** (lines 329-335): ✅ Correct
|
|
183
|
+
- Zero gradients
|
|
184
|
+
- Backward pass
|
|
185
|
+
- Optimizer step
|
|
186
|
+
|
|
187
|
+
5. **Soft Update** (lines 361-368): ✅ Correct formula
|
|
188
|
+
```python
|
|
189
|
+
target_param.data.copy_(
|
|
190
|
+
self.tau * eval_param.data + (1.0 - self.tau) * target_param.data
|
|
191
|
+
)
|
|
192
|
+
```
|
|
193
|
+
|
|
194
|
+
## Potential Learning Issues
|
|
195
|
+
|
|
196
|
+
### 1. Target Network Not Initialized Properly
|
|
197
|
+
|
|
198
|
+
**Most Likely Issue**: If `init_hook()` fails silently, target network has random weights, causing:
|
|
199
|
+
- Incorrect Q-targets
|
|
200
|
+
- No learning signal
|
|
201
|
+
- Random behavior
|
|
202
|
+
|
|
203
|
+
**How to Verify**:
|
|
204
|
+
```python
|
|
205
|
+
# After initialization, check if target network matches actor
|
|
206
|
+
actor_params = list(agent.actor.parameters())
|
|
207
|
+
target_params = list(agent.actor_target.parameters())
|
|
208
|
+
for a, t in zip(actor_params, target_params):
|
|
209
|
+
if not torch.allclose(a.data, t.data, atol=1e-6):
|
|
210
|
+
print("WARNING: Target network not initialized correctly!")
|
|
211
|
+
```
|
|
212
|
+
|
|
213
|
+
### 2. Tau Too Small
|
|
214
|
+
|
|
215
|
+
**Config**: `TAU: 0.001` (line 18)
|
|
216
|
+
|
|
217
|
+
**Impact**:
|
|
218
|
+
- Very slow target network updates
|
|
219
|
+
- Target network stays close to initial values for a long time
|
|
220
|
+
- Slower learning convergence
|
|
221
|
+
|
|
222
|
+
**Typical Values**:
|
|
223
|
+
- DQN papers often use `tau=0.01` or `tau=0.005`
|
|
224
|
+
- `tau=0.001` means only 0.1% update per step
|
|
225
|
+
|
|
226
|
+
**Recommendation**: Try `tau=0.01` or `tau=0.005`
|
|
227
|
+
|
|
228
|
+
### 3. Learning Rate
|
|
229
|
+
|
|
230
|
+
**Config**: `LR: 0.001` (line 12)
|
|
231
|
+
|
|
232
|
+
**Impact**:
|
|
233
|
+
- Might be too high for some environments
|
|
234
|
+
- Could cause instability
|
|
235
|
+
|
|
236
|
+
**Typical Values**:
|
|
237
|
+
- DQN often uses `lr=1e-4` to `lr=5e-4`
|
|
238
|
+
- `lr=0.001` is on the higher side
|
|
239
|
+
|
|
240
|
+
**Recommendation**: Try `lr=5e-4` or `lr=1e-4`
|
|
241
|
+
|
|
242
|
+
### 4. Learn Step Frequency
|
|
243
|
+
|
|
244
|
+
**Config**: `LEARN_STEP: 1` (line 17)
|
|
245
|
+
|
|
246
|
+
**Impact**:
|
|
247
|
+
- Learning every step (with 16 parallel envs, that's 16 steps per environment step)
|
|
248
|
+
- Very frequent learning might cause instability
|
|
249
|
+
- Typical DQN learns every 4-5 steps
|
|
250
|
+
|
|
251
|
+
**Recommendation**: Try `LEARN_STEP: 4` or `LEARN_STEP: 5`
|
|
252
|
+
|
|
253
|
+
## Summary of Critical Issues
|
|
254
|
+
|
|
255
|
+
1. **🔴 HIGH PRIORITY**: Target network initialization might fail silently
|
|
256
|
+
- Check if `init_hook()` actually initializes target network
|
|
257
|
+
- Add fallback to `load_state_dict()` if TensorDict method fails
|
|
258
|
+
|
|
259
|
+
2. **🟡 MEDIUM PRIORITY**: Inconsistent `max()` usage
|
|
260
|
+
- Change `axis=1` to `dim=1` for consistency
|
|
261
|
+
|
|
262
|
+
3. **🟡 MEDIUM PRIORITY**: Consider adding gradient clipping
|
|
263
|
+
- Prevents gradient explosion
|
|
264
|
+
|
|
265
|
+
4. **🟡 MEDIUM PRIORITY**: Hyperparameter tuning
|
|
266
|
+
- `tau=0.001` might be too small
|
|
267
|
+
- `lr=0.001` might be too high
|
|
268
|
+
- `learn_step=1` might be too frequent
|
|
269
|
+
|
|
270
|
+
## Recommended Fixes
|
|
271
|
+
|
|
272
|
+
### Fix 1: Improve Target Network Initialization
|
|
273
|
+
|
|
274
|
+
```python
|
|
275
|
+
def init_hook(self) -> None:
|
|
276
|
+
"""Resets module parameters for the detached and target networks."""
|
|
277
|
+
param_vals: TensorDict = from_module(self.actor).detach()
|
|
278
|
+
target_params: TensorDict = param_vals.clone().lock_()
|
|
279
|
+
|
|
280
|
+
try:
|
|
281
|
+
target_params.to_module(self.actor_target)
|
|
282
|
+
# Verify initialization succeeded
|
|
283
|
+
actor_first_param = next(self.actor.parameters()).data
|
|
284
|
+
target_first_param = next(self.actor_target.parameters()).data
|
|
285
|
+
if not torch.allclose(actor_first_param, target_first_param, atol=1e-5):
|
|
286
|
+
raise RuntimeError("Target network initialization verification failed")
|
|
287
|
+
except (KeyError, RuntimeError) as e:
|
|
288
|
+
warnings.warn(f"TensorDict initialization failed ({e}), using load_state_dict fallback")
|
|
289
|
+
self.actor_target.load_state_dict(self.actor.state_dict())
|
|
290
|
+
finally:
|
|
291
|
+
self.param_vals = param_vals
|
|
292
|
+
self.target_params = target_params
|
|
293
|
+
```
|
|
294
|
+
|
|
295
|
+
### Fix 2: Fix max() Usage
|
|
296
|
+
|
|
297
|
+
```python
|
|
298
|
+
# Line 316
|
|
299
|
+
q_target = self.actor_target(next_obs).max(dim=1)[0].unsqueeze(1)
|
|
300
|
+
```
|
|
301
|
+
|
|
302
|
+
### Fix 3: Add Gradient Clipping (Optional)
|
|
303
|
+
|
|
304
|
+
```python
|
|
305
|
+
# After line 333 (loss.backward())
|
|
306
|
+
from torch.nn.utils import clip_grad_norm_
|
|
307
|
+
clip_grad_norm_(self.actor.parameters(), max_norm=10.0)
|
|
308
|
+
self.optimizer.step()
|
|
309
|
+
```
|