agilerl 2.7.0.dev2__tar.gz → 2.7.0.dev3__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 (439) hide show
  1. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/PKG-INFO +1 -1
  2. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/cispo.py +1 -4
  3. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/base.py +322 -96
  4. agilerl-2.7.0.dev3/agilerl/algorithms/core/llm_ops/__init__.py +36 -0
  5. {agilerl-2.7.0.dev2/agilerl/algorithms/core → agilerl-2.7.0.dev3/agilerl/algorithms/core/llm_ops}/fused_lora.py +5 -0
  6. agilerl-2.7.0.dev3/agilerl/algorithms/core/llm_ops/fused_loss.py +484 -0
  7. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/dpo.py +8 -4
  8. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/grpo.py +91 -56
  9. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/gspo.py +1 -4
  10. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ppo_llm.py +285 -59
  11. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/reinforce_llm.py +170 -35
  12. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/sft.py +4 -3
  13. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/sft.py +1 -1
  14. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/token_observation.py +14 -2
  15. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/protocols.py +8 -3
  16. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_llm.py +15 -0
  17. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/llm_utils.py +61 -113
  18. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/ppo_value_head.py +0 -3
  19. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/utils.py +26 -5
  20. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_llm_multiturn.py +57 -12
  21. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/cispo.yaml +17 -8
  22. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/grpo.yaml +12 -11
  23. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/grpo_multiturn.yaml +17 -6
  24. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/gspo.yaml +15 -6
  25. agilerl-2.7.0.dev3/configs/training/llm_finetuning/ppo_llm.yaml +44 -0
  26. agilerl-2.7.0.dev3/configs/training/llm_finetuning/reinforce_llm.yaml +40 -0
  27. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_stage_2.py +1 -1
  28. agilerl-2.7.0.dev3/docs/llm_finetuning/fused_logprobs.rst +109 -0
  29. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/llm_finetuning/index.rst +7 -0
  30. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/pyproject.toml +1 -1
  31. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_core_base.py +447 -5
  32. {agilerl-2.7.0.dev2/tests/test_algorithms/test_llms → agilerl-2.7.0.dev3/tests/test_algorithms/test_llm_ops}/test_fused_lora.py +16 -8
  33. agilerl-2.7.0.dev3/tests/test_algorithms/test_llm_ops/test_fused_loss.py +775 -0
  34. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_dpo.py +8 -0
  35. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_grpo.py +114 -376
  36. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_ppo_llm.py +201 -0
  37. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_reinforce_llm.py +116 -0
  38. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_sft.py +4 -4
  39. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_protocols.py +21 -0
  40. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_llm_utils.py +9 -63
  41. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_ppo_value_head.py +0 -8
  42. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_utils.py +103 -0
  43. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_llm_envs.py +9 -7
  44. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/uv.lock +1 -1
  45. agilerl-2.7.0.dev2/configs/training/llm_finetuning/ppo_llm.yaml +0 -47
  46. agilerl-2.7.0.dev2/configs/training/llm_finetuning/reinforce_llm.yaml +0 -42
  47. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/ISSUE_TEMPLATE/bug_report.md +0 -0
  48. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/ISSUE_TEMPLATE/feature_request.md +0 -0
  49. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/PULL_REQUEST_TEMPLATE/pull_request_template.md +0 -0
  50. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/badges/arena-github-badge.svg +0 -0
  51. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/codeql/install_codeql.sh +0 -0
  52. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/codeql/run_codeql.py +0 -0
  53. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/dependabot.yml +0 -0
  54. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/codeql.yml +0 -0
  55. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/linux-tests.yml +0 -0
  56. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/macos-tests.yml +0 -0
  57. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/windows-tests.yml +0 -0
  58. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.gitignore +0 -0
  59. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.pre-commit-config.yaml +0 -0
  60. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.readthedocs.yaml +0 -0
  61. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/CITATION.cff +0 -0
  62. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/CODE_OF_CONDUCT.md +0 -0
  63. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/CONTRIBUTING.md +0 -0
  64. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/LICENSE +0 -0
  65. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/README.md +0 -0
  66. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/__init__.py +0 -0
  67. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/__init__.py +0 -0
  68. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/bc_lm.py +0 -0
  69. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/__init__.py +0 -0
  70. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/optimizer_wrapper.py +0 -0
  71. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/registry.py +0 -0
  72. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/cqn.py +0 -0
  73. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ddpg.py +0 -0
  74. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/dqn.py +0 -0
  75. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/dqn_rainbow.py +0 -0
  76. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ilql.py +0 -0
  77. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ippo.py +0 -0
  78. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/maddpg.py +0 -0
  79. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/matd3.py +0 -0
  80. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/neural_ts_bandit.py +0 -0
  81. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/neural_ucb_bandit.py +0 -0
  82. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ppo.py +0 -0
  83. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/td3.py +0 -0
  84. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/__init__.py +0 -0
  85. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/data.py +0 -0
  86. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/multi_agent_replay_buffer.py +0 -0
  87. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/replay_buffer.py +0 -0
  88. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/rollout_buffer.py +0 -0
  89. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/sampler.py +0 -0
  90. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/segment_tree.py +0 -0
  91. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/__init__.py +0 -0
  92. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/language_environment.py +0 -0
  93. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/rl_data.py +0 -0
  94. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/tokenizer.py +0 -0
  95. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/torch_datasets.py +0 -0
  96. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/hpo/__init__.py +0 -0
  97. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/hpo/mutation.py +0 -0
  98. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/hpo/tournament.py +0 -0
  99. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/__init__.py +0 -0
  100. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/base.py +0 -0
  101. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/preference.py +0 -0
  102. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/reasoning.py +0 -0
  103. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/search.py +0 -0
  104. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/sync_vec_env.py +0 -0
  105. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/__init__.py +0 -0
  106. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/base.py +0 -0
  107. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/bert.py +0 -0
  108. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/cnn.py +0 -0
  109. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/configs.py +0 -0
  110. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/custom_components.py +0 -0
  111. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/dummy.py +0 -0
  112. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/gpt.py +0 -0
  113. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/lstm.py +0 -0
  114. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/mlp.py +0 -0
  115. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/multi_input.py +0 -0
  116. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/resnet.py +0 -0
  117. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/simba.py +0 -0
  118. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/__init__.py +0 -0
  119. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/actors.py +0 -0
  120. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/base.py +0 -0
  121. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/custom_modules.py +0 -0
  122. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/distributions.py +0 -0
  123. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/q_networks.py +0 -0
  124. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/value_networks.py +0 -0
  125. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/rollouts/__init__.py +0 -0
  126. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/rollouts/on_policy.py +0 -0
  127. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/__init__.py +0 -0
  128. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_bandits.py +0 -0
  129. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_multi_agent_off_policy.py +0 -0
  130. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_multi_agent_on_policy.py +0 -0
  131. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_off_policy.py +0 -0
  132. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_offline.py +0 -0
  133. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_on_policy.py +0 -0
  134. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/typing.py +0 -0
  135. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/__init__.py +0 -0
  136. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/algo_utils.py +0 -0
  137. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/cache.py +0 -0
  138. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/evolvable_networks.py +0 -0
  139. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/ilql_utils.py +0 -0
  140. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/log_utils.py +0 -0
  141. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/minari_utils.py +0 -0
  142. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/probe_envs.py +0 -0
  143. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/probe_envs_llm.py +0 -0
  144. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/probe_envs_ma.py +0 -0
  145. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/sampling_utils.py +0 -0
  146. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/torch_utils.py +0 -0
  147. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/vector/__init__.py +0 -0
  148. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/vector/pz_async_vec_env.py +0 -0
  149. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/vector/pz_vec_env.py +0 -0
  150. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/__init__.py +0 -0
  151. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/agent.py +0 -0
  152. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/learning.py +0 -0
  153. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/llm_envs.py +0 -0
  154. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/make_evolvable.py +0 -0
  155. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/pettingzoo_wrappers.py +0 -0
  156. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/utils.py +0 -0
  157. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_bandits.py +0 -0
  158. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_llm_preference.py +0 -0
  159. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_llm_reasoning.py +0 -0
  160. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_multi_agent_off_policy.py +0 -0
  161. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_multi_agent_on_policy.py +0 -0
  162. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_off_policy.py +0 -0
  163. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_off_policy_distributed.py +0 -0
  164. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_offline.py +0 -0
  165. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_offline_distributed.py +0 -0
  166. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_on_policy.py +0 -0
  167. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_rainbow.py +0 -0
  168. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_recurrent.py +0 -0
  169. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_resnet.py +0 -0
  170. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_sft.py +0 -0
  171. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_simba.py +0 -0
  172. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/configs/ds_config.json +0 -0
  173. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/make_evolvable_benchmarking.py +0 -0
  174. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/networks.py +0 -0
  175. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/accelerate/accelerate.yaml +0 -0
  176. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/accelerate/grpo_accelerate_config.yaml +0 -0
  177. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/bandit/neural_ts.yaml +0 -0
  178. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/bandit/neural_ucb.yaml +0 -0
  179. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/cqn.yaml +0 -0
  180. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ddpg/ddpg.yaml +0 -0
  181. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ddpg/ddpg_lstm.yaml +0 -0
  182. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ddpg/ddpg_simba.yaml +0 -0
  183. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/dqn/dqn.yaml +0 -0
  184. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/dqn/dqn_lstm.yaml +0 -0
  185. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/dqn/dqn_rainbow.yaml +0 -0
  186. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/dpo.yaml +0 -0
  187. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/ippo.yaml +0 -0
  188. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/ippo_pong.yaml +0 -0
  189. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/maddpg.yaml +0 -0
  190. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/matd3.yaml +0 -0
  191. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_input.yaml +0 -0
  192. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ppo/ppo.yaml +0 -0
  193. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ppo/ppo_image.yaml +0 -0
  194. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ppo/ppo_recurrent.yaml +0 -0
  195. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/sft.yaml +0 -0
  196. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/td3.yaml +0 -0
  197. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/cartpole/cartpole_random_v1.1.0.h5 +0 -0
  198. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/cartpole/cartpole_v1.1.0.h5 +0 -0
  199. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/pendulum/pendulum_random_v1.1.0.h5 +0 -0
  200. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/pendulum/pendulum_v1.1.0.h5 +0 -0
  201. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/bandits/demo_bandit.py +0 -0
  202. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/config_load.py +0 -0
  203. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/grpo_constant_target.yaml +0 -0
  204. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/grpo_grid_navigation.yaml +0 -0
  205. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_conditional_target.yaml +0 -0
  206. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_constant_target.yaml +0 -0
  207. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_grid_navigation.yaml +0 -0
  208. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_multi_input.yaml +0 -0
  209. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_value_head.yaml +0 -0
  210. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm.py +0 -0
  211. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_stage_1.py +0 -0
  212. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_stage_3.py +0 -0
  213. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_training_matrix.py +0 -0
  214. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_value.py +0 -0
  215. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/llm_debug_utils.py +0 -0
  216. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/tiny_model.py +0 -0
  217. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/demo_llm_finetuning.py +0 -0
  218. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/multi_agent/demo_multi_agent.py +0 -0
  219. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_custom_network.py +0 -0
  220. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_off_policy.py +0 -0
  221. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_off_policy_distributed.py +0 -0
  222. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_offline.py +0 -0
  223. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_offline_distributed.py +0 -0
  224. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy.py +0 -0
  225. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy_rnn_cartpole.py +0 -0
  226. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy_rnn_memory.py +0 -0
  227. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy_rnn_minigrid.py +0 -0
  228. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_cartpole.py +0 -0
  229. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_lunar_lander.py +0 -0
  230. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_lunar_lander_rnn.py +0 -0
  231. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_rnn_memory.py +0 -0
  232. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/Makefile +0 -0
  233. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/__init__.py +0 -0
  234. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/arena-github-badge.svg +0 -0
  235. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/css/custom.css +0 -0
  236. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/favicon.ico +0 -0
  237. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/js/expand_sidebar.js +0 -0
  238. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/logo_teal.png +0 -0
  239. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/logo_white.png +0 -0
  240. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/module.png +0 -0
  241. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/network.png +0 -0
  242. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/iris-thumbnail.png +0 -0
  243. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/pendigits-thumbnail.png +0 -0
  244. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/rainbow_performance.png +0 -0
  245. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/simba_thumbnail.png +0 -0
  246. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/base.rst +0 -0
  247. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/cispo.rst +0 -0
  248. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/cql.rst +0 -0
  249. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ddpg.rst +0 -0
  250. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/dpo.rst +0 -0
  251. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/dqn.rst +0 -0
  252. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/dqn_rainbow.rst +0 -0
  253. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/grpo.rst +0 -0
  254. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/gspo.rst +0 -0
  255. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ilql.rst +0 -0
  256. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/index.rst +0 -0
  257. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ippo.rst +0 -0
  258. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/llmppo.rst +0 -0
  259. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/llmreinforce.rst +0 -0
  260. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/maddpg.rst +0 -0
  261. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/matd3.rst +0 -0
  262. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/neural_ts.rst +0 -0
  263. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/neural_ucb.rst +0 -0
  264. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ppo.rst +0 -0
  265. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/registry.rst +0 -0
  266. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/sft.rst +0 -0
  267. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/td3.rst +0 -0
  268. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/wrappers.rst +0 -0
  269. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/data.rst +0 -0
  270. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/index.rst +0 -0
  271. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/multi_agent_replay_buffer.rst +0 -0
  272. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/replay_buffer.rst +0 -0
  273. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/rollout_buffer.rst +0 -0
  274. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/sampler.rst +0 -0
  275. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/segment_tree.rst +0 -0
  276. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/hpo/index.rst +0 -0
  277. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/hpo/mutation.rst +0 -0
  278. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/hpo/tournament.rst +0 -0
  279. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/base.rst +0 -0
  280. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/bert.rst +0 -0
  281. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/cnn.rst +0 -0
  282. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/custom_activation.rst +0 -0
  283. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/dummy.rst +0 -0
  284. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/gpt.rst +0 -0
  285. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/index.rst +0 -0
  286. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/lstm.rst +0 -0
  287. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/mlp.rst +0 -0
  288. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/multi_input.rst +0 -0
  289. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/resnet.rst +0 -0
  290. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/simba.rst +0 -0
  291. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/actors.rst +0 -0
  292. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/base.rst +0 -0
  293. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/index.rst +0 -0
  294. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/q_networks.rst +0 -0
  295. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/value_networks.rst +0 -0
  296. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/rollouts/index.rst +0 -0
  297. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/rollouts/on_policy.rst +0 -0
  298. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/train.rst +0 -0
  299. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/algo_utils.rst +0 -0
  300. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/cache.rst +0 -0
  301. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/evolvable_networks.rst +0 -0
  302. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/ilql_utils.rst +0 -0
  303. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/index.rst +0 -0
  304. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/llm_utils.rst +0 -0
  305. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/log_utils.rst +0 -0
  306. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/minari_utils.rst +0 -0
  307. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/probe_envs.rst +0 -0
  308. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/torch_utils.rst +0 -0
  309. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/utils.rst +0 -0
  310. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/vector/index.rst +0 -0
  311. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/vector/petting_zoo_async_vector_env.rst +0 -0
  312. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/vector/petting_zoo_vector_env.rst +0 -0
  313. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/agent.rst +0 -0
  314. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/index.rst +0 -0
  315. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/learning.rst +0 -0
  316. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/llm_envs.rst +0 -0
  317. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/make_evolvable.rst +0 -0
  318. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/pettingzoo.rst +0 -0
  319. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/bandits/index.rst +0 -0
  320. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/conf.py +0 -0
  321. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/custom_algorithms/index.rst +0 -0
  322. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/debugging_rl/index.rst +0 -0
  323. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/distributed_training/index.rst +0 -0
  324. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/evo_hyperparam_opt/index.rst +0 -0
  325. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/evolvable_networks/index.rst +0 -0
  326. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/get_started/agilerl2changes.rst +0 -0
  327. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/get_started/index.rst +0 -0
  328. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/index.rst +0 -0
  329. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/llm_finetuning/llm_checkpoints.rst +0 -0
  330. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/make.bat +0 -0
  331. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/multi_agent_training/index.rst +0 -0
  332. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/off_policy/index.rst +0 -0
  333. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/offline_training/index.rst +0 -0
  334. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/on_policy/index.rst +0 -0
  335. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/pomdp/index.rst +0 -0
  336. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/releases/index.rst +0 -0
  337. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/requirements.txt +0 -0
  338. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/sitecustomize.py +0 -0
  339. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/__init__.py +0 -0
  340. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/build_minari_fixture.py +0 -0
  341. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/build_tiny_llm_fixture.py +0 -0
  342. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/door/human-v2/data/main_data.hdf5 +0 -0
  343. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/door/human-v2/data/metadata.json +0 -0
  344. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/door/namespace_metadata.json +0 -0
  345. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/namespace_metadata.json +0 -0
  346. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/added_tokens.json +0 -0
  347. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/chat_template.jinja +0 -0
  348. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/config.json +0 -0
  349. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/generation_config.json +0 -0
  350. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/model.safetensors +0 -0
  351. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/special_tokens_map.json +0 -0
  352. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/tokenizer.json +0 -0
  353. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/tokenizer_config.json +0 -0
  354. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/conftest.py +0 -0
  355. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/helper_functions.py +0 -0
  356. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/pz_vector_test_utils.py +0 -0
  357. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/subprocess_runner.py +0 -0
  358. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/__init__.py +0 -0
  359. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/conftest.py +0 -0
  360. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bandits/__init__.py +0 -0
  361. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bandits/test_neural_ts.py +0 -0
  362. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bandits/test_neural_ucb.py +0 -0
  363. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_base.py +0 -0
  364. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bc_lm.py +0 -0
  365. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/__init__.py +0 -0
  366. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/conftest.py +0 -0
  367. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_llm_checkpoint.py +0 -0
  368. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_vllm.py +0 -0
  369. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/__init__.py +0 -0
  370. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/conftest.py +0 -0
  371. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/test_ippo.py +0 -0
  372. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/test_maddpg.py +0 -0
  373. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/test_matd3.py +0 -0
  374. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_optimizer_wrapper.py +0 -0
  375. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_registry.py +0 -0
  376. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/__init__.py +0 -0
  377. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_cqn.py +0 -0
  378. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_ddpg.py +0 -0
  379. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_dqn.py +0 -0
  380. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_dqn_rainbow.py +0 -0
  381. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_ilql.py +0 -0
  382. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_ppo.py +0 -0
  383. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_td3.py +0 -0
  384. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/__init__.py +0 -0
  385. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_multi_agent_replay_buffer.py +0 -0
  386. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_replay_buffer.py +0 -0
  387. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_replay_data.py +0 -0
  388. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_rollout_buffer.py +0 -0
  389. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_sampler.py +0 -0
  390. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_segment_tree.py +0 -0
  391. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_data.py +0 -0
  392. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_hpo/__init__.py +0 -0
  393. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_hpo/test_mutation.py +0 -0
  394. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_hpo/test_tournament.py +0 -0
  395. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_init.py +0 -0
  396. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/__init__.py +0 -0
  397. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_base.py +0 -0
  398. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_bert.py +0 -0
  399. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_cnn.py +0 -0
  400. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_configs.py +0 -0
  401. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_custom_activation.py +0 -0
  402. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_dummy.py +0 -0
  403. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_gpt.py +0 -0
  404. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_lstm.py +0 -0
  405. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_mlp.py +0 -0
  406. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_multi_input.py +0 -0
  407. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_resnet.py +0 -0
  408. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_simba.py +0 -0
  409. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/__init__.py +0 -0
  410. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_actors.py +0 -0
  411. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_base.py +0 -0
  412. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_distributions.py +0 -0
  413. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_q_networks.py +0 -0
  414. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_value_functions.py +0 -0
  415. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_rollouts/test_on_policy.py +0 -0
  416. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_train/test_train.py +0 -0
  417. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_train/test_train_llm.py +0 -0
  418. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/__init__.py +0 -0
  419. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_algo_utils.py +0 -0
  420. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_cache.py +0 -0
  421. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_ilql_utils.py +0 -0
  422. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_log_utils.py +0 -0
  423. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_minari_utils.py +0 -0
  424. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_probe_envs.py +0 -0
  425. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_probe_envs_llm.py +0 -0
  426. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_probe_envs_ma.py +0 -0
  427. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_sampling_utils.py +0 -0
  428. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_torch_utils.py +0 -0
  429. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_utils_evolvable.py +0 -0
  430. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_vector/test_vector.py +0 -0
  431. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/__init__.py +0 -0
  432. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_agent.py +0 -0
  433. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_autoreset.py +0 -0
  434. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_bandit_env.py +0 -0
  435. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_make_evolvable.py +0 -0
  436. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_multiturn_wrappers.py +0 -0
  437. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_ppo_test_method.py +0 -0
  438. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_skills.py +0 -0
  439. {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: agilerl
3
- Version: 2.7.0.dev2
3
+ Version: 2.7.0.dev3
4
4
  Summary: AgileRL is a deep reinforcement learning library focused on improving RL development through RLOps.
5
5
  Author-email: Nick Ustaran-Anderegg <dev@agilerl.com>
6
6
  License-Expression: Apache-2.0
@@ -2,7 +2,6 @@
2
2
 
3
3
  from __future__ import annotations
4
4
 
5
- from functools import partial
6
5
  from typing import Any
7
6
 
8
7
  from agilerl.algorithms.grpo import GRPO, _signatures_without_loss_type
@@ -14,11 +13,9 @@ class CISPO(GRPO):
14
13
  Paper: https://arxiv.org/abs/2506.13585
15
14
  """
16
15
 
17
- _init_with_cispo = partial(GRPO.__init__, loss_type="cispo")
18
-
19
16
  def __init__(self, *args: Any, **kwargs: Any) -> None:
20
17
  """Initialize a CISPO agent with fixed ``loss_type``."""
21
- self._init_with_cispo(self, *args, **kwargs)
18
+ super().__init__(*args, loss_type="cispo", **kwargs)
22
19
 
23
20
 
24
21
  _CISPO_CLASS_SIG, _CISPO_INIT_SIG = _signatures_without_loss_type()
@@ -24,7 +24,6 @@ from typing import (
24
24
  import dill
25
25
  import numpy as np
26
26
  import torch
27
- import torch.nn.functional as F
28
27
  from accelerate import Accelerator
29
28
  from accelerate.utils import broadcast_object_list, set_seed
30
29
  from gymnasium import spaces
@@ -121,13 +120,14 @@ if TYPE_CHECKING or HAS_DEEPSPEED:
121
120
  if TYPE_CHECKING or HAS_VLLM:
122
121
  from vllm import LLM, SamplingParams
123
122
 
124
- from agilerl.algorithms.core.fused_lora import (
123
+ from agilerl.algorithms.core.llm_ops.fused_lora import (
125
124
  clear_fused_adapter_routing,
126
125
  patch_lora_for_fused_forward,
127
126
  set_fused_adapter_routing,
128
127
  )
129
128
  from agilerl.utils.llm_utils import (
130
129
  align_deepspeed_lr,
130
+ build_completion_mask,
131
131
  create_model_from_name_or_path,
132
132
  gather_if_zero3,
133
133
  get_model_name_or_path,
@@ -1979,7 +1979,9 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
1979
1979
  :type pad_token_id: int
1980
1980
  :param pad_token: The pad token.
1981
1981
  :type pad_token: str
1982
- :param use_liger_loss: Whether to use Liger loss.
1982
+ :param use_liger_loss: Whether to use Liger loss. Defaults to ``False``.
1983
+ Passing ``True`` without ``liger-kernel`` installed warns and falls
1984
+ back to ``False``.
1983
1985
  :type use_liger_loss: bool
1984
1986
  :param lora_config: The LoRA config.
1985
1987
  :type lora_config: LoraConfigProtocol | None
@@ -2015,6 +2017,24 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
2015
2017
  :param reduce_memory_peak: Deprecated. Previously hinted peak-memory batching;
2016
2018
  ignored. Configure ``micro_batch_size_per_gpu`` and DeepSpeed instead.
2017
2019
  :type reduce_memory_peak: bool, optional
2020
+ :param cast_logprobs_to_fp32: When ``True`` (the default), the per-token
2021
+ log-probability reduction (``amax`` / ``gather`` / ``logsumexp``)
2022
+ runs in fp32 before being cast back to the input dtype. Applies
2023
+ uniformly to both the unfused ``(B, T, V)`` path
2024
+ (:meth:`_logprobs_from_logits`) and the fused linear log-prob
2025
+ path (:meth:`_logprobs_from_hidden_fused`) so the two paths
2026
+ produce numerically equivalent log-probs.
2027
+
2028
+ The default preserves prior behaviour exactly: the unfused path
2029
+ was already promoting to fp32 unconditionally before this flag
2030
+ existed. The flag exposes that promotion as configurable.
2031
+
2032
+ Setting ``False`` introduces a per-token bf16 quantisation error
2033
+ (~0.1 at ``V≈128k``) which can bias PPO/GRPO importance-sampling
2034
+ ratios. Use only if you've verified bf16 is acceptable for your
2035
+ vocab/shape — it saves ~18 GB on the unfused path at ``B=8,
2036
+ T=2048, V≈152k``, ~6 MB on the fused path.
2037
+ :type cast_logprobs_to_fp32: bool, optional
2018
2038
  """
2019
2039
 
2020
2040
  _separate_reference_adapter_deprecation_emitted = False
@@ -2052,6 +2072,8 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
2052
2072
  gradient_checkpointing: bool = True,
2053
2073
  torch_compiler: str | None = None,
2054
2074
  reduce_memory_peak: bool = False,
2075
+ use_fused_linear_logprobs: bool = False,
2076
+ cast_logprobs_to_fp32: bool = True,
2055
2077
  ) -> None:
2056
2078
  if not HAS_LLM_DEPENDENCIES:
2057
2079
  msg = "LLM dependencies are not installed. Please install them using `pip install agilerl[llm]`."
@@ -2196,6 +2218,8 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
2196
2218
  self.wrap = wrap
2197
2219
  self.use_separate_reference_adapter = use_separate_reference_adapter
2198
2220
  self._warn_separate_reference_adapter_deprecation()
2221
+ self.use_fused_linear_logprobs = use_fused_linear_logprobs
2222
+ self.cast_logprobs_to_fp32 = cast_logprobs_to_fp32
2199
2223
 
2200
2224
  selected_adapters = ("actor",)
2201
2225
  if use_separate_reference_adapter:
@@ -3424,6 +3448,7 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3424
3448
  """
3425
3449
  unwrapped = self._get_unwrapped_actor()
3426
3450
  total = fused_ids.shape[0]
3451
+ seq_len_out = fused_ids.shape[1] - 1
3427
3452
 
3428
3453
  position_ids = None
3429
3454
  if self.calc_position_embeddings:
@@ -3436,9 +3461,20 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3436
3461
  else [(s, min(s + batch_size, total)) for s in range(0, total, batch_size)]
3437
3462
  )
3438
3463
 
3439
- all_logprobs: list[torch.Tensor] = []
3440
- all_values: list[torch.Tensor] = []
3441
- for start, end in chunks:
3464
+ # Fused-linear-logprob path: replace lm_head with nn.Identity for the
3465
+ # no-grad forward, then compute per-token logprobs via a chunked
3466
+ # matmul over the lm_head weight. Skips materializing (B, T, V).
3467
+ # Only safe when grads are disabled — autograd graph would not
3468
+ # capture the manual matmul.
3469
+ use_fused_lp = self.use_fused_linear_logprobs and not torch.is_grad_enabled()
3470
+ if use_fused_lp:
3471
+ lm_head = self._get_lm_head()
3472
+ lm_head_weight = lm_head.weight
3473
+ lm_head_bias = lm_head.bias
3474
+
3475
+ def _process_chunk(
3476
+ start: int, end: int
3477
+ ) -> tuple[torch.Tensor, torch.Tensor | None]:
3442
3478
  set_fused_adapter_routing(unwrapped, routing[start:end])
3443
3479
  model_kwargs: dict = {
3444
3480
  "input_ids": fused_ids[start:end],
@@ -3448,38 +3484,84 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3448
3484
  if position_ids is not None:
3449
3485
  model_kwargs["position_ids"] = position_ids[start:end]
3450
3486
 
3451
- with self._amp_ctx():
3487
+ patch_ctx = (
3488
+ self._patch_lm_head_to_identity() if use_fused_lp else nullcontext()
3489
+ )
3490
+ with patch_ctx, self._amp_ctx():
3452
3491
  output = self.actor.forward(**model_kwargs)
3453
3492
 
3454
3493
  if isinstance(output, tuple):
3455
3494
  # Value-head models may return (loss, logits, value, ...); Peft/causal
3456
- # paths may return shorter tuples — only index when present.
3457
- logits = output[0]
3495
+ # paths may return shorter tuples — only index when present. With
3496
+ # lm_head identity-patched, output[0] is the last hidden state.
3497
+ first = output[0]
3458
3498
  value = output[2] if len(output) > 2 else None
3459
3499
  else:
3460
- logits = output.logits
3500
+ first = output.logits
3461
3501
  value = None
3462
-
3463
3502
  del output
3464
- logits = logits / self.temperature
3465
3503
 
3466
- all_logprobs.append(
3467
- LLMAlgorithm._memory_efficient_logits(
3504
+ if use_fused_lp:
3505
+ chunk_lp = LLMAlgorithm._logprobs_from_hidden_fused(
3506
+ first[:, :-1],
3507
+ lm_head_weight,
3508
+ lm_head_bias,
3509
+ fused_ids[start:end, 1:],
3510
+ temperature=self.temperature,
3511
+ cast_to_fp32=self.cast_logprobs_to_fp32,
3512
+ )
3513
+ del first
3514
+ else:
3515
+ logits = first / self.temperature
3516
+ del first
3517
+ chunk_lp = LLMAlgorithm._logprobs_from_logits(
3468
3518
  logits[:, :-1],
3469
3519
  fused_ids[start:end, 1:],
3520
+ cast_to_fp32=self.cast_logprobs_to_fp32,
3470
3521
  )
3522
+ del logits
3523
+
3524
+ chunk_v = (
3525
+ value[:, :-1] if (self.use_value_head and value is not None) else None
3471
3526
  )
3472
- if self.use_value_head and value is not None:
3473
- all_values.append(value[:, :-1])
3527
+ return chunk_lp, chunk_v
3474
3528
 
3475
- if self.use_value_head:
3476
- values = torch.cat(all_values, dim=0) if len(chunks) > 1 else all_values[0]
3477
- else:
3478
- values = None
3479
- logprobs = (
3480
- torch.cat(all_logprobs, dim=0) if len(chunks) > 1 else all_logprobs[0]
3481
- )
3482
- return logprobs, values
3529
+ # Single-chunk fast path: skip the buffer + copy entirely.
3530
+ if len(chunks) == 1:
3531
+ return _process_chunk(0, total)
3532
+
3533
+ # Multi-chunk path: pre-allocate output buffers once and write each
3534
+ # chunk in place via copy_(). Avoids holding the full list of chunk
3535
+ # tensors plus the concatenated buffer in memory at the same time
3536
+ # (which doubles peak memory in the torch.cat path).
3537
+ logprobs_out: torch.Tensor | None = None
3538
+ values_out: torch.Tensor | None = None
3539
+
3540
+ for start, end in chunks:
3541
+ chunk_lp, chunk_v = _process_chunk(start, end)
3542
+
3543
+ # Lazy-allocate on the first chunk so we inherit dtype/device
3544
+ # from the model output rather than guessing up front.
3545
+ if logprobs_out is None:
3546
+ logprobs_out = torch.empty(
3547
+ (total, seq_len_out),
3548
+ dtype=chunk_lp.dtype,
3549
+ device=chunk_lp.device,
3550
+ )
3551
+ logprobs_out[start:end].copy_(chunk_lp)
3552
+ del chunk_lp
3553
+
3554
+ if chunk_v is not None:
3555
+ if values_out is None:
3556
+ values_out = torch.empty(
3557
+ (total, seq_len_out),
3558
+ dtype=chunk_v.dtype,
3559
+ device=chunk_v.device,
3560
+ )
3561
+ values_out[start:end].copy_(chunk_v)
3562
+ del chunk_v
3563
+
3564
+ return logprobs_out, values_out
3483
3565
 
3484
3566
  def _fused_forward(
3485
3567
  self,
@@ -3629,6 +3711,7 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3629
3711
  :return: Log probabilities of the completion IDs.
3630
3712
  :rtype: torch.Tensor
3631
3713
  """
3714
+ use_fused_lp = self.use_fused_linear_logprobs and not torch.is_grad_enabled()
3632
3715
  with self.select_adapter("reference" if use_reference else "actor"):
3633
3716
  self.actor.train(mode=not eval_mode)
3634
3717
  num_samples = ids.shape[0]
@@ -3639,6 +3722,11 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3639
3722
  position_ids = attention_mask.long().cumsum(dim=-1) - 1
3640
3723
  position_ids.masked_fill_(mask=(attention_mask == 0), value=1)
3641
3724
 
3725
+ if use_fused_lp:
3726
+ lm_head = self._get_lm_head()
3727
+ lm_head_weight = lm_head.weight
3728
+ lm_head_bias = lm_head.bias
3729
+
3642
3730
  # Split the sample into batches
3643
3731
  log_probs = []
3644
3732
  for batch in range(0, num_samples, batch_size):
@@ -3653,18 +3741,33 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3653
3741
  if self.calc_position_embeddings:
3654
3742
  batch_position_ids = position_ids[batch:end_idx, :]
3655
3743
  batch_model_kwargs |= {"position_ids": batch_position_ids}
3656
- with self._amp_ctx():
3657
- output = self.actor.forward(**batch_model_kwargs)
3658
- logits = output[0] if isinstance(output, tuple) else output.logits
3659
- logits = logits / self.temperature
3660
-
3661
- log_prob = LLMAlgorithm._memory_efficient_logits(
3662
- logits[:, :-1],
3663
- batch_ids[:, 1:],
3744
+ patch_ctx = (
3745
+ self._patch_lm_head_to_identity() if use_fused_lp else nullcontext()
3664
3746
  )
3747
+ with patch_ctx, self._amp_ctx():
3748
+ output = self.actor.forward(**batch_model_kwargs)
3749
+ first = output[0] if isinstance(output, tuple) else output.logits
3750
+
3751
+ if use_fused_lp:
3752
+ log_prob = LLMAlgorithm._logprobs_from_hidden_fused(
3753
+ first[:, :-1],
3754
+ lm_head_weight,
3755
+ lm_head_bias,
3756
+ batch_ids[:, 1:],
3757
+ temperature=self.temperature,
3758
+ cast_to_fp32=self.cast_logprobs_to_fp32,
3759
+ )
3760
+ else:
3761
+ logits = first / self.temperature
3762
+ log_prob = LLMAlgorithm._logprobs_from_logits(
3763
+ logits[:, :-1],
3764
+ batch_ids[:, 1:],
3765
+ cast_to_fp32=self.cast_logprobs_to_fp32,
3766
+ )
3767
+ logits = None
3665
3768
 
3769
+ first = None
3666
3770
  batch_model_kwargs = None
3667
- logits = None
3668
3771
  log_probs.append(log_prob)
3669
3772
  return torch.cat(log_probs, dim=0)
3670
3773
 
@@ -3796,14 +3899,19 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3796
3899
  msg = "vLLM is required when use_vllm=True. Install AgileRL with vLLM support for this platform: `pip install agilerl[llm]`."
3797
3900
  raise ImportError(msg)
3798
3901
 
3902
+ max_token_cap = (
3903
+ self.max_output_tokens
3904
+ if self.max_output_tokens is not None
3905
+ else self.max_model_len
3906
+ )
3907
+
3799
3908
  def _trajectory_input_ids(prompt: dict[str, Any]) -> torch.Tensor:
3800
3909
  return cast(
3801
3910
  "torch.Tensor",
3802
3911
  prompt.get("trajectory_input_ids", prompt["input_ids"]),
3803
3912
  )
3804
3913
 
3805
- def _token_prompt_for_vllm(prompt: dict[str, Any]) -> dict[str, list[int]]:
3806
- ids = _trajectory_input_ids(prompt)
3914
+ def _token_prompt_for_vllm(ids: torch.Tensor) -> dict[str, list[int]]:
3807
3915
  return {"prompt_token_ids": ids.squeeze(0).tolist()}
3808
3916
 
3809
3917
  def _stitch_prefix(prompt: dict[str, Any], ref: torch.Tensor) -> torch.Tensor:
@@ -3812,7 +3920,7 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3812
3920
  return ref.new_zeros((ref.shape[0], 0))
3813
3921
  return cast("torch.Tensor", st)
3814
3922
 
3815
- def _vllm_max_new_tokens(model_prompt_len: int, max_token_cap: int) -> int:
3923
+ def _vllm_max_new_tokens(model_prompt_len: int) -> int:
3816
3924
  room = self.max_model_len - model_prompt_len
3817
3925
  if room <= 0:
3818
3926
  error_msg = f"Model prompt length ({model_prompt_len}) is greater than the model length ({self.max_model_len})"
@@ -3822,22 +3930,25 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3822
3930
  max_out = max(max_out, min(self.min_output_tokens, room))
3823
3931
  return min(max_out, room)
3824
3932
 
3825
- group_prompts = [prompt for prompt in prompts for _ in range(group_size)]
3826
- prompts_ids = [_trajectory_input_ids(p) for p in group_prompts]
3827
- stitch_prefixes = [
3828
- _stitch_prefix(p, prompts_ids[i]) for i, p in enumerate(group_prompts)
3829
- ]
3830
- token_prompts = [_token_prompt_for_vllm(p) for p in group_prompts]
3831
- max_token_cap = (
3832
- self.max_output_tokens
3833
- if self.max_output_tokens is not None
3834
- else self.max_model_len
3835
- )
3836
- max_output_tokens = [
3837
- _vllm_max_new_tokens(int(prompt_id.shape[1]), max_token_cap)
3838
- for prompt_id in prompts_ids
3933
+ # Compute the per-prompt work once per *unique* prompt (N items),
3934
+ # then alias by reference across each group (N·G items)
3935
+ unique_ids = [_trajectory_input_ids(p) for p in prompts]
3936
+ unique_tokens = [_token_prompt_for_vllm(ids) for ids in unique_ids]
3937
+ unique_max = [_vllm_max_new_tokens(int(ids.shape[1])) for ids in unique_ids]
3938
+ unique_stitch = [
3939
+ _stitch_prefix(p, ids) for p, ids in zip(prompts, unique_ids, strict=True)
3839
3940
  ]
3840
3941
 
3942
+ # Replicate by reference for the flat vLLM batch. Entries within a
3943
+ # group of `group_size` are aliased references to the same tensor / dict
3944
+ # — safe because downstream use is read-only is read-only w.r.t. these objects.
3945
+ # Do not introduce in-place ops on these aliases.
3946
+ group_prompts = [p for p in prompts for _ in range(group_size)]
3947
+ prompts_ids = [ids for ids in unique_ids for _ in range(group_size)]
3948
+ token_prompts = [tp for tp in unique_tokens for _ in range(group_size)]
3949
+ max_output_tokens = [m for m in unique_max for _ in range(group_size)]
3950
+ stitch_prefixes = [sp for sp in unique_stitch for _ in range(group_size)]
3951
+
3841
3952
  if self.vllm_config.tensor_parallel_size > 1:
3842
3953
  orig_size = len(token_prompts)
3843
3954
 
@@ -3922,10 +4033,17 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3922
4033
  prompts_ids = all_prompts_ids[tp_slice]
3923
4034
  stitch_prefixes = all_stitch_prefixes[tp_slice]
3924
4035
 
3925
- prompts_ids = [p.to(self.device, non_blocking=True) for p in prompts_ids]
3926
- stitch_prefixes = [
3927
- sp.to(self.device, non_blocking=True) for sp in stitch_prefixes
4036
+ # Transfer fromn host-to-device once per unique prompt, then re-alias across the group.
4037
+ unique_prompts_ids_dev = [
4038
+ prompts_ids[group_size * i].to(self.device, non_blocking=True)
4039
+ for i in range(len(prompts))
3928
4040
  ]
4041
+ unique_stitch_dev = [
4042
+ stitch_prefixes[group_size * i].to(self.device, non_blocking=True)
4043
+ for i in range(len(prompts))
4044
+ ]
4045
+ prompts_ids = [ids for ids in unique_prompts_ids_dev for _ in range(group_size)]
4046
+ stitch_prefixes = [sp for sp in unique_stitch_dev for _ in range(group_size)]
3929
4047
 
3930
4048
  completion_ids = [
3931
4049
  torch.cat(
@@ -3933,7 +4051,7 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3933
4051
  torch.cat(
3934
4052
  prompts_ids[group_size * i : group_size * (i + 1)],
3935
4053
  dim=0,
3936
- ).to(self.device),
4054
+ ),
3937
4055
  stack_and_pad_experiences(
3938
4056
  completion_ids[group_size * i : group_size * (i + 1)],
3939
4057
  padding_values=[self.pad_token_id],
@@ -3958,67 +4076,137 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
3958
4076
  int(cast("torch.Tensor", prompts[i]["input_ids"]).shape[1])
3959
4077
  for i in range(len(prompts))
3960
4078
  ]
3961
- completion_masks = []
3962
-
3963
- for i, completion_id in enumerate(completion_ids):
3964
- completion_mask = torch.zeros_like(
3965
- completion_id,
3966
- dtype=torch.bool,
3967
- device=self.device,
3968
- )
3969
- completion_mask[:, num_input_tokens[i] :] = True
3970
- completion_mask[completion_id == self.pad_token_id] = False
3971
- completion_mask = completion_mask[:, 1:]
3972
- completion_masks.append(completion_mask)
4079
+ completion_masks = [
4080
+ build_completion_mask(completion_id, num_input_tokens[i], self.pad_token_id)
4081
+ for i, completion_id in enumerate(completion_ids)
4082
+ ]
3973
4083
 
3974
4084
  return completion_ids, completion_masks
3975
4085
 
3976
4086
  @staticmethod
3977
- def _memory_efficient_logits(
4087
+ def _logprobs_from_logits(
3978
4088
  logits: torch.Tensor,
3979
4089
  index: torch.Tensor,
4090
+ cast_to_fp32: bool = True,
3980
4091
  _chunk_rows: int = 1,
3981
4092
  ) -> torch.Tensor:
3982
4093
  """Calculate log probabilities for previously generated token ids.
3983
4094
 
3984
- Processes a few rows at a time so peak memory stays bounded to
3985
- ``(_chunk_rows, seq_len, vocab_size)`` rather than the full batch,
3986
- avoiding OOM on large-vocabulary models while reducing Python loop
3987
- overhead compared to a strict row-by-row approach.
4095
+ Processes ``_chunk_rows`` rows at a time so peak memory stays bounded to
4096
+ ``(_chunk_rows, seq_len, vocab_size)`` rather than the full batch, avoiding
4097
+ OOM on large-vocabulary models. Default ``_chunk_rows=1`` minimizes the
4098
+ fp32 workspace at the cost of more kernel launches; raise to amortize
4099
+ launch overhead when memory headroom allows.
4100
+
4101
+ With ``cast_to_fp32=True``, the per-chunk reduction (``amax`` /
4102
+ ``gather`` / ``logsumexp``) runs in fp32 then casts the
4103
+ ``(B, seq_len)`` output back to *logits* dtype. Matches the precision
4104
+ of ``F.log_softmax`` over the same inputs to within the final bf16
4105
+ cast. With ``cast_to_fp32=False`` the reduction stays in *logits*
4106
+ dtype throughout — faster and lower peak (no fp32 workspace) at the
4107
+ cost of bf16-quantisation error in the reduction.
4108
+
4109
+ Logits are max-centered per row before ``logsumexp``, matching
4110
+ ``F.log_softmax`` stability either way.
3988
4111
 
3989
4112
  :param logits: Logits of shape ``(B, seq_len, vocab_size)``.
3990
4113
  :type logits: torch.Tensor
3991
4114
  :param index: Token IDs of shape ``(B, seq_len)``.
3992
4115
  :type index: torch.Tensor
4116
+ :param cast_to_fp32: Promote each chunk to fp32 before the reduction.
4117
+ :type cast_to_fp32: bool
3993
4118
  :return: Log probabilities of the completion IDs, shape ``(B, seq_len)``.
3994
4119
  :rtype: torch.Tensor
3995
4120
  """
3996
- # 1. Gather the raw logits for the specific token IDs.
3997
- # Shape reduces from (B, seq_len, vocab_size) immediately to (B, seq_len)
3998
-
4121
+ orig_dtype = logits.dtype
3999
4122
  B = logits.shape[0]
4123
+
4124
+ def _logprobs_chunk(lg: torch.Tensor, idx: torch.Tensor) -> torch.Tensor:
4125
+ if cast_to_fp32:
4126
+ lg = lg.float()
4127
+ max_lg = lg.amax(dim=-1, keepdim=True)
4128
+ shifted = lg - max_lg
4129
+ target = shifted.gather(dim=-1, index=idx.unsqueeze(-1)).squeeze(-1)
4130
+ log_z = torch.logsumexp(shifted, dim=-1)
4131
+ result = target - log_z
4132
+ return result.to(orig_dtype) if cast_to_fp32 else result
4133
+
4000
4134
  if B <= _chunk_rows:
4001
- return (
4002
- F.log_softmax(logits, dim=-1)
4003
- .gather(dim=-1, index=index.unsqueeze(-1))
4004
- .squeeze(-1)
4005
- )
4135
+ return _logprobs_chunk(logits, index)
4006
4136
 
4007
4137
  per_token_logps = []
4008
4138
  for start in range(0, B, _chunk_rows):
4009
4139
  end = min(start + _chunk_rows, B)
4010
- target_logits_chunk = (
4011
- logits[start:end]
4012
- .gather(dim=-1, index=index[start:end].unsqueeze(-1))
4013
- .squeeze(-1)
4140
+ per_token_logps.append(
4141
+ _logprobs_chunk(logits[start:end], index[start:end]),
4014
4142
  )
4015
- log_z_chunk = torch.logsumexp(logits[start:end], dim=-1)
4016
- per_token_logps_chunk = (target_logits_chunk - log_z_chunk).to(
4017
- logits.dtype
4018
- ) # Do we need to upcast to float 32 here??
4019
- per_token_logps.append(per_token_logps_chunk)
4020
4143
  return torch.cat(per_token_logps, dim=0)
4021
4144
 
4145
+ @staticmethod
4146
+ def _logprobs_from_hidden_fused(
4147
+ hidden: torch.Tensor,
4148
+ lm_head_weight: torch.Tensor,
4149
+ lm_head_bias: torch.Tensor | None,
4150
+ target_ids: torch.Tensor,
4151
+ temperature: float = 1.0,
4152
+ cast_to_fp32: bool = True,
4153
+ _chunk_rows: int = 1024,
4154
+ ) -> torch.Tensor:
4155
+ """Per-token target logprobs without materializing the full ``(B, T, V)``
4156
+ logits tensor.
4157
+
4158
+ Tiles flat over ``(B*T)`` with workspace bounded to ``(_chunk_rows, V)``
4159
+ per iteration. Counterpart of :meth:`_logprobs_from_logits` for
4160
+ callers that hold hidden states and the lm_head separately. **No-grad
4161
+ only** — gradients won't flow to ``lm_head_weight`` from this fn.
4162
+
4163
+ Numerical contract matches :meth:`_logprobs_from_logits` when fed
4164
+ equivalent inputs (``logits = (hidden @ Wᵀ + b) / T``): same
4165
+ ``cast_to_fp32`` semantics, same final-cast-back-to-input-dtype, same
4166
+ max-shift ``gather - logsumexp`` formulation. Default ``cast_to_fp32=True``
4167
+ keeps the two paths bit-comparable.
4168
+
4169
+ :param hidden: ``(B, T, H)`` last-hidden-state.
4170
+ :param lm_head_weight: ``(V, H)``.
4171
+ :param lm_head_bias: ``(V,)`` or ``None``.
4172
+ :param target_ids: ``(B, T)`` (caller does the ``[:, :-1]``/``[:, 1:]``
4173
+ shift before calling).
4174
+ :param temperature: scalar; logits divided by this before log_softmax
4175
+ (skipped when ``1.0``).
4176
+ :param cast_to_fp32: when True (default), run the per-chunk reduction
4177
+ in fp32 then cast back. Same semantics as
4178
+ :meth:`_logprobs_from_logits`.
4179
+ :param _chunk_rows: rows of the flattened ``(B*T)`` workspace per
4180
+ iteration; trades launch count vs ``_chunk_rows * V`` peak.
4181
+ :return: ``(B, T)`` per-token logprobs in ``hidden.dtype``.
4182
+ """
4183
+ orig_dtype = hidden.dtype
4184
+ B, T, H = hidden.shape
4185
+ flat_h = hidden.reshape(-1, H)
4186
+ flat_targets = target_ids.reshape(-1).to(torch.long)
4187
+ N = flat_h.shape[0]
4188
+ out = torch.empty(N, dtype=orig_dtype, device=hidden.device)
4189
+ W_t = lm_head_weight.t()
4190
+
4191
+ for s in range(0, N, _chunk_rows):
4192
+ e = min(s + _chunk_rows, N)
4193
+ chunk_logits = flat_h[s:e] @ W_t
4194
+ if lm_head_bias is not None:
4195
+ chunk_logits.add_(lm_head_bias)
4196
+ if temperature != 1.0:
4197
+ chunk_logits.div_(temperature)
4198
+ if cast_to_fp32:
4199
+ chunk_logits = chunk_logits.float()
4200
+ mx = chunk_logits.amax(dim=-1, keepdim=True)
4201
+ chunk_logits.sub_(mx)
4202
+ tgt = chunk_logits.gather(dim=-1, index=flat_targets[s:e, None]).squeeze(-1)
4203
+ log_z = torch.logsumexp(chunk_logits, dim=-1)
4204
+ del chunk_logits
4205
+ result = tgt - log_z
4206
+ out[s:e].copy_(result.to(orig_dtype) if cast_to_fp32 else result)
4207
+
4208
+ return out.reshape(B, T)
4209
+
4022
4210
  def _configure_batch_size_per_process(
4023
4211
  self,
4024
4212
  batch_size: int,
@@ -4651,26 +4839,64 @@ class LLMAlgorithm(EvolvableAlgorithm, ABC):
4651
4839
  if hasattr(self.actor.optimizer, "clip_grad"):
4652
4840
  self.actor.optimizer.clip_grad = self.max_grad_norm
4653
4841
 
4654
- def _get_lm_head(self):
4655
- """Locate the lm_head module, handling both raw and PEFT-wrapped models.
4842
+ def _get_lm_head_parent(self) -> tuple[Any, str]:
4843
+ """Locate the parent module owning ``lm_head`` (or ``embed_out``).
4656
4844
 
4657
- :return: The lm_head (or embed_out) linear layer.
4658
- :rtype: torch.nn.Module
4845
+ Walks through value-head, PEFT, and LoRA wrappers to the inner
4846
+ causal-LM that exposes the language-model head as an attribute.
4847
+ Returned so that callers can both read the head (``getattr(parent,
4848
+ attr)``) and replace it temporarily (``setattr(parent, attr, ...)``)
4849
+ — the latter is used by the no-grad fused-linear-logprob path.
4850
+
4851
+ :return: ``(parent_module, attr_name)``.
4659
4852
  :raises AttributeError: If no lm_head can be found.
4660
4853
  """
4661
4854
  model = self.actor
4855
+ if self.use_value_head and hasattr(model, "pretrained_model"):
4856
+ # Value-head wrapper (e.g. AutoModelForCausalLMWithValueHead) →
4857
+ # the PEFT/causal-LM inner model.
4858
+ model = model.pretrained_model
4662
4859
  if hasattr(model, "base_model"): # PeftModel → LoraModel
4663
4860
  model = model.base_model
4664
4861
  if hasattr(model, "model"): # LoraModel → CausalLM
4665
4862
  model = model.model
4666
4863
  for attr in ("lm_head", "embed_out"):
4667
4864
  if hasattr(model, attr):
4668
- return getattr(model, attr)
4669
- err_msg = f"""Cannot find lm_head in {type(self.actor).__name__}.
4670
- Set use_liger_loss=False.
4671
- """
4865
+ return model, attr
4866
+ err_msg = (
4867
+ f"Cannot find lm_head in {type(self.actor).__name__}. "
4868
+ "Set use_liger_loss=False and use_fused_linear_logprobs=False."
4869
+ )
4672
4870
  raise AttributeError(err_msg)
4673
4871
 
4872
+ def _get_lm_head(self):
4873
+ """Locate the lm_head module, handling value-head, PEFT and LoRA wrappers.
4874
+
4875
+ :return: The lm_head (or embed_out) linear layer.
4876
+ :rtype: torch.nn.Module
4877
+ :raises AttributeError: If no lm_head can be found.
4878
+ """
4879
+ parent, attr = self._get_lm_head_parent()
4880
+ return getattr(parent, attr)
4881
+
4882
+ @contextmanager
4883
+ def _patch_lm_head_to_identity(self):
4884
+ """Temporarily replace ``lm_head`` with ``nn.Identity``.
4885
+
4886
+ With the head identity-patched, the model's ``output.logits`` becomes
4887
+ the post-final-norm hidden state ``(B, T, H)`` instead of the full
4888
+ ``(B, T, V)`` logits — which is what the no-grad fused-linear-logprob
4889
+ kernel consumes directly. The original module is always restored,
4890
+ even if the wrapped block raises.
4891
+ """
4892
+ model, attr = self._get_lm_head_parent()
4893
+ original = getattr(model, attr)
4894
+ setattr(model, attr, torch.nn.Identity())
4895
+ try:
4896
+ yield original
4897
+ finally:
4898
+ setattr(model, attr, original)
4899
+
4674
4900
  def _get_unwrapped_actor(self) -> Any:
4675
4901
  """Return actor unwrapped from Accelerate and DummyEvolvable layers."""
4676
4902
  actor = (