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.
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/PKG-INFO +1 -1
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/cispo.py +1 -4
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/base.py +322 -96
- agilerl-2.7.0.dev3/agilerl/algorithms/core/llm_ops/__init__.py +36 -0
- {agilerl-2.7.0.dev2/agilerl/algorithms/core → agilerl-2.7.0.dev3/agilerl/algorithms/core/llm_ops}/fused_lora.py +5 -0
- agilerl-2.7.0.dev3/agilerl/algorithms/core/llm_ops/fused_loss.py +484 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/dpo.py +8 -4
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/grpo.py +91 -56
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/gspo.py +1 -4
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ppo_llm.py +285 -59
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/reinforce_llm.py +170 -35
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/sft.py +4 -3
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/sft.py +1 -1
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/token_observation.py +14 -2
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/protocols.py +8 -3
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_llm.py +15 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/llm_utils.py +61 -113
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/ppo_value_head.py +0 -3
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/utils.py +26 -5
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_llm_multiturn.py +57 -12
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/cispo.yaml +17 -8
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/grpo.yaml +12 -11
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/grpo_multiturn.yaml +17 -6
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/gspo.yaml +15 -6
- agilerl-2.7.0.dev3/configs/training/llm_finetuning/ppo_llm.yaml +44 -0
- agilerl-2.7.0.dev3/configs/training/llm_finetuning/reinforce_llm.yaml +40 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_stage_2.py +1 -1
- agilerl-2.7.0.dev3/docs/llm_finetuning/fused_logprobs.rst +109 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/llm_finetuning/index.rst +7 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/pyproject.toml +1 -1
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_core_base.py +447 -5
- {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
- agilerl-2.7.0.dev3/tests/test_algorithms/test_llm_ops/test_fused_loss.py +775 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_dpo.py +8 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_grpo.py +114 -376
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_ppo_llm.py +201 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_reinforce_llm.py +116 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_sft.py +4 -4
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_protocols.py +21 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_llm_utils.py +9 -63
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_ppo_value_head.py +0 -8
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_utils.py +103 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_llm_envs.py +9 -7
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/uv.lock +1 -1
- agilerl-2.7.0.dev2/configs/training/llm_finetuning/ppo_llm.yaml +0 -47
- agilerl-2.7.0.dev2/configs/training/llm_finetuning/reinforce_llm.yaml +0 -42
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/ISSUE_TEMPLATE/bug_report.md +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/ISSUE_TEMPLATE/feature_request.md +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/PULL_REQUEST_TEMPLATE/pull_request_template.md +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/badges/arena-github-badge.svg +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/codeql/install_codeql.sh +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/codeql/run_codeql.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/dependabot.yml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/codeql.yml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/linux-tests.yml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/macos-tests.yml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.github/workflows/windows-tests.yml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.gitignore +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.pre-commit-config.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/.readthedocs.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/CITATION.cff +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/CODE_OF_CONDUCT.md +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/CONTRIBUTING.md +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/LICENSE +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/README.md +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/bc_lm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/optimizer_wrapper.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/core/registry.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/cqn.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ddpg.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/dqn.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/dqn_rainbow.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ilql.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ippo.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/maddpg.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/matd3.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/neural_ts_bandit.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/neural_ucb_bandit.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/ppo.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/algorithms/td3.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/data.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/multi_agent_replay_buffer.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/replay_buffer.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/rollout_buffer.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/sampler.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/components/segment_tree.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/language_environment.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/rl_data.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/tokenizer.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/data/torch_datasets.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/hpo/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/hpo/mutation.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/hpo/tournament.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/base.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/preference.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/reasoning.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/search.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/llm_envs/sync_vec_env.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/base.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/bert.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/cnn.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/configs.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/custom_components.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/dummy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/gpt.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/lstm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/mlp.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/multi_input.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/resnet.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/modules/simba.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/actors.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/base.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/custom_modules.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/distributions.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/q_networks.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/networks/value_networks.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/rollouts/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/rollouts/on_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_bandits.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_multi_agent_off_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_multi_agent_on_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_off_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_offline.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/training/train_on_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/typing.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/algo_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/cache.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/evolvable_networks.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/ilql_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/log_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/minari_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/probe_envs.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/probe_envs_llm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/probe_envs_ma.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/sampling_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/utils/torch_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/vector/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/vector/pz_async_vec_env.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/vector/pz_vec_env.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/agent.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/learning.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/llm_envs.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/make_evolvable.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/pettingzoo_wrappers.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/agilerl/wrappers/utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_bandits.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_llm_preference.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_llm_reasoning.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_multi_agent_off_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_multi_agent_on_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_off_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_off_policy_distributed.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_offline.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_offline_distributed.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_on_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_rainbow.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_recurrent.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_resnet.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_sft.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/benchmarking_simba.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/configs/ds_config.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/make_evolvable_benchmarking.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/benchmarking/networks.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/accelerate/accelerate.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/accelerate/grpo_accelerate_config.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/bandit/neural_ts.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/bandit/neural_ucb.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/cqn.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ddpg/ddpg.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ddpg/ddpg_lstm.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ddpg/ddpg_simba.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/dqn/dqn.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/dqn/dqn_lstm.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/dqn/dqn_rainbow.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/llm_finetuning/dpo.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/ippo.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/ippo_pong.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/maddpg.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_agent/matd3.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/multi_input.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ppo/ppo.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ppo/ppo_image.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/ppo/ppo_recurrent.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/sft.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/configs/training/td3.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/cartpole/cartpole_random_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/cartpole/cartpole_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/pendulum/pendulum_random_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/data/pendulum/pendulum_v1.1.0.h5 +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/bandits/demo_bandit.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/config_load.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/grpo_constant_target.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/grpo_grid_navigation.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_conditional_target.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_constant_target.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_grid_navigation.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_multi_input.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/configs/ppo_value_head.yaml +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_stage_1.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_stage_3.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_llm_training_matrix.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/debugging_value.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/llm_debug_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/debugging/tiny_model.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/llm/demo_llm_finetuning.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/multi_agent/demo_multi_agent.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_custom_network.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_off_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_off_policy_distributed.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_offline.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_offline_distributed.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy_rnn_cartpole.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy_rnn_memory.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/demo_on_policy_rnn_minigrid.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_cartpole.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_lunar_lander.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_lunar_lander_rnn.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/demos/single_agent/performance_flamegraph_rnn_memory.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/Makefile +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/arena-github-badge.svg +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/css/custom.css +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/favicon.ico +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/js/expand_sidebar.js +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/logo_teal.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/logo_white.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/module.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/network.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/iris-thumbnail.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/pendigits-thumbnail.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/rainbow_performance.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/_static/thumbnails/simba_thumbnail.png +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/base.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/cispo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/cql.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ddpg.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/dpo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/dqn.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/dqn_rainbow.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/grpo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/gspo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ilql.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ippo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/llmppo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/llmreinforce.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/maddpg.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/matd3.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/neural_ts.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/neural_ucb.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/ppo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/registry.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/sft.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/td3.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/algorithms/wrappers.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/data.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/multi_agent_replay_buffer.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/replay_buffer.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/rollout_buffer.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/sampler.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/components/segment_tree.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/hpo/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/hpo/mutation.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/hpo/tournament.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/base.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/bert.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/cnn.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/custom_activation.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/dummy.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/gpt.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/lstm.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/mlp.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/multi_input.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/resnet.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/modules/simba.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/actors.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/base.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/q_networks.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/networks/value_networks.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/rollouts/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/rollouts/on_policy.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/train.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/algo_utils.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/cache.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/evolvable_networks.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/ilql_utils.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/llm_utils.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/log_utils.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/minari_utils.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/probe_envs.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/torch_utils.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/utils/utils.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/vector/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/vector/petting_zoo_async_vector_env.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/vector/petting_zoo_vector_env.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/agent.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/learning.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/llm_envs.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/make_evolvable.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/api/wrappers/pettingzoo.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/bandits/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/conf.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/custom_algorithms/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/debugging_rl/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/distributed_training/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/evo_hyperparam_opt/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/evolvable_networks/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/get_started/agilerl2changes.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/get_started/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/llm_finetuning/llm_checkpoints.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/make.bat +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/multi_agent_training/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/off_policy/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/offline_training/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/on_policy/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/pomdp/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/releases/index.rst +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/docs/requirements.txt +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/sitecustomize.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/build_minari_fixture.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/build_tiny_llm_fixture.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/door/human-v2/data/main_data.hdf5 +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/door/human-v2/data/metadata.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/door/namespace_metadata.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/minari_cache/D4RL/namespace_metadata.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/added_tokens.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/chat_template.jinja +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/config.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/generation_config.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/model.safetensors +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/special_tokens_map.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/tokenizer.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/assets/tiny_llm/tokenizer_config.json +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/conftest.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/helper_functions.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/pz_vector_test_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/subprocess_runner.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/conftest.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bandits/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bandits/test_neural_ts.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bandits/test_neural_ucb.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_base.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_bc_lm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/conftest.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_llm_checkpoint.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_llms/test_vllm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/conftest.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/test_ippo.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/test_maddpg.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_multi_agent/test_matd3.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_optimizer_wrapper.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_registry.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_cqn.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_ddpg.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_dqn.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_dqn_rainbow.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_ilql.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_ppo.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_algorithms/test_single_agent/test_td3.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_multi_agent_replay_buffer.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_replay_buffer.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_replay_data.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_rollout_buffer.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_sampler.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_components/test_segment_tree.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_data.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_hpo/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_hpo/test_mutation.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_hpo/test_tournament.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_init.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_base.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_bert.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_cnn.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_configs.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_custom_activation.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_dummy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_gpt.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_lstm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_mlp.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_multi_input.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_resnet.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_modules/test_simba.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_actors.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_base.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_distributions.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_q_networks.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_networks/test_value_functions.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_rollouts/test_on_policy.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_train/test_train.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_train/test_train_llm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_algo_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_cache.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_ilql_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_log_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_minari_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_probe_envs.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_probe_envs_llm.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_probe_envs_ma.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_sampling_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_torch_utils.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_utils/test_utils_evolvable.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_vector/test_vector.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/__init__.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_agent.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_autoreset.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_bandit_env.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_make_evolvable.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_multiturn_wrappers.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_ppo_test_method.py +0 -0
- {agilerl-2.7.0.dev2 → agilerl-2.7.0.dev3}/tests/test_wrappers/test_skills.py +0 -0
- {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.
|
|
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
|
-
|
|
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
|
-
|
|
3440
|
-
|
|
3441
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
3500
|
+
first = output.logits
|
|
3461
3501
|
value = None
|
|
3462
|
-
|
|
3463
3502
|
del output
|
|
3464
|
-
logits = logits / self.temperature
|
|
3465
3503
|
|
|
3466
|
-
|
|
3467
|
-
LLMAlgorithm.
|
|
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
|
-
|
|
3473
|
-
all_values.append(value[:, :-1])
|
|
3527
|
+
return chunk_lp, chunk_v
|
|
3474
3528
|
|
|
3475
|
-
|
|
3476
|
-
|
|
3477
|
-
|
|
3478
|
-
|
|
3479
|
-
|
|
3480
|
-
|
|
3481
|
-
|
|
3482
|
-
|
|
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
|
-
|
|
3657
|
-
|
|
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(
|
|
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
|
|
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
|
-
|
|
3826
|
-
|
|
3827
|
-
|
|
3828
|
-
|
|
3829
|
-
]
|
|
3830
|
-
|
|
3831
|
-
|
|
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
|
-
|
|
3926
|
-
|
|
3927
|
-
|
|
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
|
-
)
|
|
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
|
-
|
|
3964
|
-
|
|
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
|
|
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
|
|
3985
|
-
``(_chunk_rows, seq_len, vocab_size)`` rather than the full batch,
|
|
3986
|
-
|
|
3987
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
|
4655
|
-
"""Locate the
|
|
4842
|
+
def _get_lm_head_parent(self) -> tuple[Any, str]:
|
|
4843
|
+
"""Locate the parent module owning ``lm_head`` (or ``embed_out``).
|
|
4656
4844
|
|
|
4657
|
-
|
|
4658
|
-
|
|
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
|
|
4669
|
-
err_msg =
|
|
4670
|
-
|
|
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 = (
|