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.
Files changed (441) hide show
  1. agilerl-2.7.0.dev1/.github/workflows/linux-tests.yml +96 -0
  2. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/workflows/macos-tests.yml +10 -0
  3. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/workflows/windows-tests.yml +10 -0
  4. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.gitignore +2 -0
  5. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.pre-commit-config.yaml +2 -2
  6. agilerl-2.7.0.dev1/CLAUDE.md +139 -0
  7. agilerl-2.7.0.dev1/DQN_LEARNING_ALGORITHM_ANALYSIS.md +309 -0
  8. agilerl-2.7.0.dev1/DQN_LEARNING_ANALYSIS.md +168 -0
  9. agilerl-2.7.0.dev1/GPU_CLEANUP_ANALYSIS.md +541 -0
  10. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/PKG-INFO +1 -1
  11. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/__init__.py +13 -0
  12. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/__init__.py +5 -1
  13. agilerl-2.7.0.dev1/agilerl/algorithms/cispo.py +26 -0
  14. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/base.py +2096 -839
  15. agilerl-2.7.0.dev1/agilerl/algorithms/core/fused_lora.py +146 -0
  16. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/optimizer_wrapper.py +112 -14
  17. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/registry.py +4 -3
  18. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/dpo.py +63 -103
  19. agilerl-2.7.0.dev1/agilerl/algorithms/grpo.py +1076 -0
  20. agilerl-2.7.0.dev1/agilerl/algorithms/gspo.py +26 -0
  21. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ippo.py +0 -43
  22. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/maddpg.py +123 -42
  23. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/matd3.py +149 -55
  24. agilerl-2.7.0.dev1/agilerl/algorithms/ppo_llm.py +905 -0
  25. agilerl-2.7.0.dev1/agilerl/algorithms/reinforce_llm.py +729 -0
  26. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/sft.py +25 -18
  27. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/hpo/mutation.py +1 -1
  28. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/hpo/tournament.py +2 -2
  29. agilerl-2.7.0.dev1/agilerl/llm_envs/__init__.py +37 -0
  30. agilerl-2.7.0.dev1/agilerl/llm_envs/base.py +261 -0
  31. agilerl-2.7.0.dev1/agilerl/llm_envs/preference.py +135 -0
  32. agilerl-2.7.0.dev1/agilerl/llm_envs/reasoning.py +163 -0
  33. agilerl-2.7.0.dev1/agilerl/llm_envs/search.py +120 -0
  34. agilerl-2.7.0.dev1/agilerl/llm_envs/sft.py +99 -0
  35. agilerl-2.7.0.dev1/agilerl/llm_envs/sync_vec_env.py +273 -0
  36. agilerl-2.7.0.dev1/agilerl/llm_envs/token_observation.py +349 -0
  37. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/protocols.py +24 -4
  38. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/rollouts/on_policy.py +66 -1
  39. agilerl-2.7.0.dev1/agilerl/training/train_llm.py +1867 -0
  40. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/typing.py +15 -4
  41. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/algo_utils.py +84 -20
  42. agilerl-2.7.0.dev1/agilerl/utils/llm_utils.py +867 -0
  43. agilerl-2.7.0.dev1/agilerl/utils/ppo_value_head.py +275 -0
  44. agilerl-2.7.0.dev1/agilerl/utils/probe_envs_llm.py +154 -0
  45. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/utils.py +435 -68
  46. agilerl-2.7.0.dev1/agilerl/wrappers/llm_envs.py +19 -0
  47. agilerl-2.7.0.dev1/benchmarking/benchmarking_llm_multiturn.py +120 -0
  48. agilerl-2.7.0.dev0/benchmarking/benchmarking_dpo.py → agilerl-2.7.0.dev1/benchmarking/benchmarking_llm_preference.py +3 -3
  49. agilerl-2.7.0.dev0/benchmarking/benchmarking_grpo.py → agilerl-2.7.0.dev1/benchmarking/benchmarking_llm_reasoning.py +36 -57
  50. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_sft.py +2 -2
  51. agilerl-2.7.0.dev1/configs/training/llm_finetuning/cispo.yaml +38 -0
  52. {agilerl-2.7.0.dev0/configs/training → agilerl-2.7.0.dev1/configs/training/llm_finetuning}/grpo.yaml +4 -3
  53. agilerl-2.7.0.dev1/configs/training/llm_finetuning/grpo_multiturn.yaml +30 -0
  54. agilerl-2.7.0.dev1/configs/training/llm_finetuning/gspo.yaml +36 -0
  55. agilerl-2.7.0.dev1/configs/training/llm_finetuning/ppo_llm.yaml +47 -0
  56. agilerl-2.7.0.dev1/configs/training/llm_finetuning/reinforce_llm.yaml +42 -0
  57. agilerl-2.7.0.dev1/demos/llm/debugging/config_load.py +22 -0
  58. agilerl-2.7.0.dev1/demos/llm/debugging/configs/grpo_constant_target.yaml +24 -0
  59. agilerl-2.7.0.dev1/demos/llm/debugging/configs/grpo_grid_navigation.yaml +26 -0
  60. agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_conditional_target.yaml +31 -0
  61. agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_constant_target.yaml +27 -0
  62. agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_grid_navigation.yaml +29 -0
  63. agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_multi_input.yaml +28 -0
  64. agilerl-2.7.0.dev1/demos/llm/debugging/configs/ppo_value_head.yaml +30 -0
  65. agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm.py +186 -0
  66. agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_stage_1.py +252 -0
  67. agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_stage_2.py +283 -0
  68. agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_stage_3.py +400 -0
  69. agilerl-2.7.0.dev1/demos/llm/debugging/debugging_llm_training_matrix.py +697 -0
  70. agilerl-2.7.0.dev1/demos/llm/debugging/debugging_value.py +190 -0
  71. agilerl-2.7.0.dev1/demos/llm/debugging/llm_debug_utils.py +16 -0
  72. agilerl-2.7.0.dev1/demos/llm/debugging/tiny_model.py +229 -0
  73. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/llm}/demo_llm_finetuning.py +6 -6
  74. agilerl-2.7.0.dev1/docs/api/algorithms/cispo.rst +137 -0
  75. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/dpo.rst +1 -2
  76. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/grpo.rst +5 -2
  77. agilerl-2.7.0.dev1/docs/api/algorithms/gspo.rst +134 -0
  78. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/index.rst +30 -0
  79. agilerl-2.7.0.dev1/docs/api/algorithms/llmppo.rst +147 -0
  80. agilerl-2.7.0.dev1/docs/api/algorithms/llmreinforce.rst +145 -0
  81. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/sft.rst +1 -1
  82. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/llm_utils.rst +1 -1
  83. agilerl-2.7.0.dev1/docs/api/wrappers/llm_envs.rst +12 -0
  84. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/debugging_rl/index.rst +61 -1
  85. agilerl-2.7.0.dev1/docs/llm_finetuning/llm_checkpoints.rst +148 -0
  86. agilerl-2.7.0.dev1/find_dqn_commit.sh +82 -0
  87. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/pyproject.toml +4 -3
  88. agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/multiturn:GRPO:pop1:no_tournament/attributes.pt +0 -0
  89. agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/multiturn:LLMPPO:pop1:no_tournament/attributes.pt +0 -0
  90. agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/multiturn:LLMReinforce:pop1:no_tournament/attributes.pt +0 -0
  91. agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/preference:DPO:pop1:no_tournament/attributes.pt +0 -0
  92. agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/reasoning:GRPO:pop1:no_tournament/attributes.pt +0 -0
  93. agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/reasoning:LLMPPO:pop1:no_tournament/attributes.pt +0 -0
  94. agilerl-2.7.0.dev1/saved_checkpoints/debug_matrix/reasoning:LLMReinforce:pop1:no_tournament/attributes.pt +0 -0
  95. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/conftest.py +16 -1
  96. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/subprocess_runner.py +14 -7
  97. agilerl-2.7.0.dev1/tests/test_algorithms/conftest.py +38 -0
  98. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_core_base.py +1819 -472
  99. agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/conftest.py +115 -0
  100. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_llms/test_dpo.py +23 -123
  101. agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_fused_lora.py +126 -0
  102. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_llms/test_grpo.py +1130 -1109
  103. agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_ppo_llm.py +1094 -0
  104. agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_reinforce_llm.py +944 -0
  105. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_llms/test_sft.py +15 -11
  106. agilerl-2.7.0.dev1/tests/test_algorithms/test_llms/test_vllm.py +135 -0
  107. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/test_maddpg.py +152 -57
  108. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/test_matd3.py +183 -83
  109. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_optimizer_wrapper.py +209 -0
  110. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_ilql.py +6 -0
  111. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_data.py +15 -0
  112. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_hpo/test_mutation.py +4 -2
  113. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_hpo/test_tournament.py +70 -2
  114. agilerl-2.7.0.dev1/tests/test_rollouts/test_on_policy.py +383 -0
  115. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_train/test_train_llm.py +1323 -82
  116. agilerl-2.7.0.dev1/tests/test_utils/__init__.py +0 -0
  117. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_algo_utils.py +386 -2
  118. agilerl-2.7.0.dev1/tests/test_utils/test_llm_utils.py +1257 -0
  119. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_minari_utils.py +6 -2
  120. agilerl-2.7.0.dev1/tests/test_utils/test_ppo_value_head.py +484 -0
  121. agilerl-2.7.0.dev1/tests/test_utils/test_probe_envs_llm.py +268 -0
  122. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_utils.py +280 -23
  123. agilerl-2.7.0.dev1/tests/test_wrappers/__init__.py +0 -0
  124. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_llm_envs.py +13 -2
  125. agilerl-2.7.0.dev1/tests/test_wrappers/test_multiturn_wrappers.py +859 -0
  126. agilerl-2.7.0.dev1/tests/test_wrappers/test_ppo_test_method.py +79 -0
  127. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/utils.py +57 -1
  128. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/uv.lock +1982 -1895
  129. agilerl-2.7.0.dev0/.github/workflows/linux-tests.yml +0 -55
  130. agilerl-2.7.0.dev0/agilerl/algorithms/grpo.py +0 -674
  131. agilerl-2.7.0.dev0/agilerl/training/train_llm.py +0 -1080
  132. agilerl-2.7.0.dev0/agilerl/utils/llm_utils.py +0 -252
  133. agilerl-2.7.0.dev0/agilerl/wrappers/llm_envs.py +0 -795
  134. agilerl-2.7.0.dev0/docs/api/wrappers/llm_envs.rst +0 -12
  135. agilerl-2.7.0.dev0/tests/test_algorithms/test_llms/conftest.py +0 -152
  136. agilerl-2.7.0.dev0/tests/test_rollouts/test_on_policy.py +0 -149
  137. agilerl-2.7.0.dev0/tests/test_utils/test_llm_utils.py +0 -370
  138. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/ISSUE_TEMPLATE/bug_report.md +0 -0
  139. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/ISSUE_TEMPLATE/feature_request.md +0 -0
  140. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/PULL_REQUEST_TEMPLATE/pull_request_template.md +0 -0
  141. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/badges/arena-github-badge.svg +0 -0
  142. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/codeql/install_codeql.sh +0 -0
  143. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/codeql/run_codeql.py +0 -0
  144. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1/.github}/dependabot.yml +0 -0
  145. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.github/workflows/codeql.yml +0 -0
  146. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/.readthedocs.yaml +0 -0
  147. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/CITATION.cff +0 -0
  148. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/CODE_OF_CONDUCT.md +0 -0
  149. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/CONTRIBUTING.md +0 -0
  150. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/LICENSE +0 -0
  151. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/README.md +0 -0
  152. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/bc_lm.py +0 -0
  153. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/core/__init__.py +0 -0
  154. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/cqn.py +0 -0
  155. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ddpg.py +0 -0
  156. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/dqn.py +0 -0
  157. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/dqn_rainbow.py +0 -0
  158. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ilql.py +0 -0
  159. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/neural_ts_bandit.py +0 -0
  160. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/neural_ucb_bandit.py +0 -0
  161. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/ppo.py +0 -0
  162. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/algorithms/td3.py +0 -0
  163. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/__init__.py +0 -0
  164. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/data.py +0 -0
  165. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/multi_agent_replay_buffer.py +0 -0
  166. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/replay_buffer.py +0 -0
  167. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/rollout_buffer.py +0 -0
  168. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/sampler.py +0 -0
  169. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/components/segment_tree.py +0 -0
  170. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/__init__.py +0 -0
  171. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/language_environment.py +0 -0
  172. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/rl_data.py +0 -0
  173. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/tokenizer.py +0 -0
  174. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/data/torch_datasets.py +0 -0
  175. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/hpo/__init__.py +0 -0
  176. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/__init__.py +0 -0
  177. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/base.py +0 -0
  178. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/bert.py +0 -0
  179. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/cnn.py +0 -0
  180. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/configs.py +0 -0
  181. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/custom_components.py +0 -0
  182. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/dummy.py +0 -0
  183. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/gpt.py +0 -0
  184. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/lstm.py +0 -0
  185. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/mlp.py +0 -0
  186. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/multi_input.py +0 -0
  187. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/resnet.py +0 -0
  188. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/modules/simba.py +0 -0
  189. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/__init__.py +0 -0
  190. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/actors.py +0 -0
  191. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/base.py +0 -0
  192. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/custom_modules.py +0 -0
  193. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/distributions.py +0 -0
  194. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/q_networks.py +0 -0
  195. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/networks/value_networks.py +0 -0
  196. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/rollouts/__init__.py +0 -0
  197. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/__init__.py +0 -0
  198. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_bandits.py +0 -0
  199. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_multi_agent_off_policy.py +0 -0
  200. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_multi_agent_on_policy.py +0 -0
  201. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_off_policy.py +0 -0
  202. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_offline.py +0 -0
  203. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/training/train_on_policy.py +0 -0
  204. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/__init__.py +0 -0
  205. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/cache.py +0 -0
  206. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/evolvable_networks.py +0 -0
  207. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/ilql_utils.py +0 -0
  208. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/log_utils.py +0 -0
  209. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/minari_utils.py +0 -0
  210. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/probe_envs.py +0 -0
  211. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/probe_envs_ma.py +0 -0
  212. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/sampling_utils.py +0 -0
  213. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/utils/torch_utils.py +0 -0
  214. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/vector/__init__.py +0 -0
  215. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/vector/pz_async_vec_env.py +0 -0
  216. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/vector/pz_vec_env.py +0 -0
  217. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/__init__.py +0 -0
  218. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/agent.py +0 -0
  219. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/learning.py +0 -0
  220. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/make_evolvable.py +0 -0
  221. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/pettingzoo_wrappers.py +0 -0
  222. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/agilerl/wrappers/utils.py +0 -0
  223. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_bandits.py +0 -0
  224. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_multi_agent_off_policy.py +0 -0
  225. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_multi_agent_on_policy.py +0 -0
  226. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_off_policy.py +0 -0
  227. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_off_policy_distributed.py +0 -0
  228. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_offline.py +0 -0
  229. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_offline_distributed.py +0 -0
  230. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_on_policy.py +0 -0
  231. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_rainbow.py +0 -0
  232. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_recurrent.py +0 -0
  233. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_resnet.py +0 -0
  234. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/benchmarking_simba.py +0 -0
  235. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/configs/ds_config.json +0 -0
  236. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/make_evolvable_benchmarking.py +0 -0
  237. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/benchmarking/networks.py +0 -0
  238. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/accelerate/accelerate.yaml +0 -0
  239. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/accelerate/grpo_accelerate_config.yaml +0 -0
  240. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/bandit/neural_ts.yaml +0 -0
  241. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/bandit/neural_ucb.yaml +0 -0
  242. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/cqn.yaml +0 -0
  243. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ddpg/ddpg.yaml +0 -0
  244. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ddpg/ddpg_lstm.yaml +0 -0
  245. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ddpg/ddpg_simba.yaml +0 -0
  246. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/dqn/dqn.yaml +0 -0
  247. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/dqn/dqn_lstm.yaml +0 -0
  248. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/dqn/dqn_rainbow.yaml +0 -0
  249. {agilerl-2.7.0.dev0/configs/training → agilerl-2.7.0.dev1/configs/training/llm_finetuning}/dpo.yaml +0 -0
  250. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/ippo.yaml +0 -0
  251. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/ippo_pong.yaml +0 -0
  252. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/maddpg.yaml +0 -0
  253. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_agent/matd3.yaml +0 -0
  254. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/multi_input.yaml +0 -0
  255. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ppo/ppo.yaml +0 -0
  256. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ppo/ppo_image.yaml +0 -0
  257. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/ppo/ppo_recurrent.yaml +0 -0
  258. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/sft.yaml +0 -0
  259. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/configs/training/td3.yaml +0 -0
  260. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/cartpole/cartpole_random_v1.1.0.h5 +0 -0
  261. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/cartpole/cartpole_v1.1.0.h5 +0 -0
  262. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/pendulum/pendulum_random_v1.1.0.h5 +0 -0
  263. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/data/pendulum/pendulum_v1.1.0.h5 +0 -0
  264. {agilerl-2.7.0.dev0/tests → agilerl-2.7.0.dev1/debugging}/__init__.py +0 -0
  265. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/bandits}/demo_bandit.py +0 -0
  266. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/multi_agent}/demo_multi_agent.py +0 -0
  267. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_custom_network.py +0 -0
  268. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_off_policy.py +0 -0
  269. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_off_policy_distributed.py +0 -0
  270. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_offline.py +0 -0
  271. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_offline_distributed.py +0 -0
  272. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy.py +0 -0
  273. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy_rnn_cartpole.py +0 -0
  274. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy_rnn_memory.py +0 -0
  275. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/demo_on_policy_rnn_minigrid.py +0 -0
  276. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_cartpole.py +0 -0
  277. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_lunar_lander.py +0 -0
  278. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_lunar_lander_rnn.py +0 -0
  279. {agilerl-2.7.0.dev0/demos → agilerl-2.7.0.dev1/demos/single_agent}/performance_flamegraph_rnn_memory.py +0 -0
  280. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/Makefile +0 -0
  281. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/__init__.py +0 -0
  282. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/arena-github-badge.svg +0 -0
  283. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/css/custom.css +0 -0
  284. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/favicon.ico +0 -0
  285. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/js/expand_sidebar.js +0 -0
  286. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/logo_teal.png +0 -0
  287. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/logo_white.png +0 -0
  288. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/module.png +0 -0
  289. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/network.png +0 -0
  290. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/iris-thumbnail.png +0 -0
  291. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/pendigits-thumbnail.png +0 -0
  292. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/rainbow_performance.png +0 -0
  293. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/_static/thumbnails/simba_thumbnail.png +0 -0
  294. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/base.rst +0 -0
  295. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/cql.rst +0 -0
  296. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ddpg.rst +0 -0
  297. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/dqn.rst +0 -0
  298. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/dqn_rainbow.rst +0 -0
  299. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ilql.rst +0 -0
  300. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ippo.rst +0 -0
  301. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/maddpg.rst +0 -0
  302. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/matd3.rst +0 -0
  303. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/neural_ts.rst +0 -0
  304. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/neural_ucb.rst +0 -0
  305. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/ppo.rst +0 -0
  306. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/registry.rst +0 -0
  307. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/td3.rst +0 -0
  308. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/algorithms/wrappers.rst +0 -0
  309. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/data.rst +0 -0
  310. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/index.rst +0 -0
  311. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/multi_agent_replay_buffer.rst +0 -0
  312. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/replay_buffer.rst +0 -0
  313. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/rollout_buffer.rst +0 -0
  314. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/sampler.rst +0 -0
  315. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/components/segment_tree.rst +0 -0
  316. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/hpo/index.rst +0 -0
  317. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/hpo/mutation.rst +0 -0
  318. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/hpo/tournament.rst +0 -0
  319. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/base.rst +0 -0
  320. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/bert.rst +0 -0
  321. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/cnn.rst +0 -0
  322. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/custom_activation.rst +0 -0
  323. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/dummy.rst +0 -0
  324. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/gpt.rst +0 -0
  325. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/index.rst +0 -0
  326. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/lstm.rst +0 -0
  327. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/mlp.rst +0 -0
  328. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/multi_input.rst +0 -0
  329. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/resnet.rst +0 -0
  330. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/modules/simba.rst +0 -0
  331. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/actors.rst +0 -0
  332. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/base.rst +0 -0
  333. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/index.rst +0 -0
  334. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/q_networks.rst +0 -0
  335. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/networks/value_networks.rst +0 -0
  336. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/rollouts/index.rst +0 -0
  337. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/rollouts/on_policy.rst +0 -0
  338. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/train.rst +0 -0
  339. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/algo_utils.rst +0 -0
  340. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/cache.rst +0 -0
  341. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/evolvable_networks.rst +0 -0
  342. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/ilql_utils.rst +0 -0
  343. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/index.rst +0 -0
  344. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/log_utils.rst +0 -0
  345. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/minari_utils.rst +0 -0
  346. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/probe_envs.rst +0 -0
  347. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/torch_utils.rst +0 -0
  348. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/utils/utils.rst +0 -0
  349. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/vector/index.rst +0 -0
  350. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/vector/petting_zoo_async_vector_env.rst +0 -0
  351. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/vector/petting_zoo_vector_env.rst +0 -0
  352. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/agent.rst +0 -0
  353. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/index.rst +0 -0
  354. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/learning.rst +0 -0
  355. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/make_evolvable.rst +0 -0
  356. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/api/wrappers/pettingzoo.rst +0 -0
  357. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/bandits/index.rst +0 -0
  358. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/conf.py +0 -0
  359. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/custom_algorithms/index.rst +0 -0
  360. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/distributed_training/index.rst +0 -0
  361. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/evo_hyperparam_opt/index.rst +0 -0
  362. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/evolvable_networks/index.rst +0 -0
  363. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/get_started/agilerl2changes.rst +0 -0
  364. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/get_started/index.rst +0 -0
  365. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/index.rst +0 -0
  366. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/llm_finetuning/index.rst +0 -0
  367. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/make.bat +0 -0
  368. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/multi_agent_training/index.rst +0 -0
  369. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/off_policy/index.rst +0 -0
  370. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/offline_training/index.rst +0 -0
  371. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/on_policy/index.rst +0 -0
  372. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/pomdp/index.rst +0 -0
  373. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/releases/index.rst +0 -0
  374. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/docs/requirements.txt +0 -0
  375. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/sitecustomize.py +0 -0
  376. {agilerl-2.7.0.dev0/tests/test_algorithms → agilerl-2.7.0.dev1/tests}/__init__.py +0 -0
  377. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/helper_functions.py +0 -0
  378. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/pz_vector_test_utils.py +0 -0
  379. {agilerl-2.7.0.dev0/tests/test_algorithms/test_bandits → agilerl-2.7.0.dev1/tests/test_algorithms}/__init__.py +0 -0
  380. {agilerl-2.7.0.dev0/tests/test_algorithms/test_llms → agilerl-2.7.0.dev1/tests/test_algorithms/test_bandits}/__init__.py +0 -0
  381. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_bandits/test_neural_ts.py +0 -0
  382. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_bandits/test_neural_ucb.py +0 -0
  383. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_base.py +0 -0
  384. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_bc_lm.py +0 -0
  385. {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
  386. /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
  387. {agilerl-2.7.0.dev0/tests/test_components → agilerl-2.7.0.dev1/tests/test_algorithms/test_multi_agent}/__init__.py +0 -0
  388. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/conftest.py +0 -0
  389. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_multi_agent/test_ippo.py +0 -0
  390. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_registry.py +0 -0
  391. {agilerl-2.7.0.dev0/tests/test_hpo → agilerl-2.7.0.dev1/tests/test_algorithms/test_single_agent}/__init__.py +0 -0
  392. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_cqn.py +0 -0
  393. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_ddpg.py +0 -0
  394. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_dqn.py +0 -0
  395. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_dqn_rainbow.py +0 -0
  396. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_ppo.py +0 -0
  397. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_algorithms/test_single_agent/test_td3.py +0 -0
  398. {agilerl-2.7.0.dev0/tests/test_modules → agilerl-2.7.0.dev1/tests/test_components}/__init__.py +0 -0
  399. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_multi_agent_replay_buffer.py +0 -0
  400. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_replay_buffer.py +0 -0
  401. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_replay_data.py +0 -0
  402. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_rollout_buffer.py +0 -0
  403. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_sampler.py +0 -0
  404. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_components/test_segment_tree.py +0 -0
  405. {agilerl-2.7.0.dev0/tests/test_networks → agilerl-2.7.0.dev1/tests/test_hpo}/__init__.py +0 -0
  406. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_init.py +0 -0
  407. {agilerl-2.7.0.dev0/tests/test_utils → agilerl-2.7.0.dev1/tests/test_modules}/__init__.py +0 -0
  408. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_base.py +0 -0
  409. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_bert.py +0 -0
  410. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_cnn.py +0 -0
  411. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_configs.py +0 -0
  412. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_custom_activation.py +0 -0
  413. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_dummy.py +0 -0
  414. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_gpt.py +0 -0
  415. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_lstm.py +0 -0
  416. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_mlp.py +0 -0
  417. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_multi_input.py +0 -0
  418. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_resnet.py +0 -0
  419. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_modules/test_simba.py +0 -0
  420. {agilerl-2.7.0.dev0/tests/test_wrappers → agilerl-2.7.0.dev1/tests/test_networks}/__init__.py +0 -0
  421. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_actors.py +0 -0
  422. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_base.py +0 -0
  423. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_distributions.py +0 -0
  424. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_q_networks.py +0 -0
  425. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_networks/test_value_functions.py +0 -0
  426. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_protocols.py +0 -0
  427. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_train/test_train.py +0 -0
  428. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_cache.py +0 -0
  429. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_ilql_utils.py +0 -0
  430. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_log_utils.py +0 -0
  431. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_probe_envs.py +0 -0
  432. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_probe_envs_ma.py +0 -0
  433. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_sampling_utils.py +0 -0
  434. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_torch_utils.py +0 -0
  435. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_utils/test_utils_evolvable.py +0 -0
  436. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_vector/test_vector.py +0 -0
  437. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_agent.py +0 -0
  438. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_autoreset.py +0 -0
  439. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_bandit_env.py +0 -0
  440. {agilerl-2.7.0.dev0 → agilerl-2.7.0.dev1}/tests/test_wrappers/test_make_evolvable.py +0 -0
  441. {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
@@ -109,6 +109,8 @@ celerybeat.pid
109
109
  # Environments
110
110
  .env
111
111
  .vscode/
112
+ .cursor/
113
+ .claude/
112
114
  .venv*/
113
115
  venv/
114
116
  env.bak/
@@ -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.10
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.6
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
+ ```