From c949a1b259bcd2457e5fffb75f6bfa216779057a Mon Sep 17 00:00:00 2001 From: charlesxu91 Date: Sat, 15 Aug 2026 22:25:36 +0800 Subject: [PATCH 1/2] feat: configure optimizer backend controls --- tests/common/config_test.py | 27 +++++++++++++++++++++++++-- trinity/common/config.py | 2 ++ trinity/trainer/verl/config.py | 10 +++++++++- 3 files changed, 36 insertions(+), 3 deletions(-) diff --git a/tests/common/config_test.py b/tests/common/config_test.py index 3ed9c797b7e..e72690d9471 100644 --- a/tests/common/config_test.py +++ b/tests/common/config_test.py @@ -11,11 +11,11 @@ import torch from tests.tools import get_template_config, get_unittest_dataset_config -from trinity.common.config import InferenceModelConfig, load_config +from trinity.common.config import InferenceModelConfig, OptimizerConfig, load_config from trinity.common.constants import SyncMethod from trinity.common.models.model import InferenceModel from trinity.trainer.trainer import is_verl_legacy -from trinity.trainer.verl.config import build_verl_config +from trinity.trainer.verl.config import _build_optimizer_config, build_verl_config CHECKPOINT_ROOT_DIR = os.path.join(os.path.dirname(__file__), "temp_checkpoint_dir") @@ -369,6 +369,29 @@ def test_optimizer_config_propagation(self): self.assertEqual(verl_config.critic.optim.weight_decay, 0.01) self.assertEqual(verl_config.critic.optim.clip_grad, 1.0) + def test_fsdp_optimizer_backend_options_are_propagated(self): + cases = [ + (OptimizerConfig(), None), + (OptimizerConfig(fused=True), {"fused": True}), + (OptimizerConfig(foreach=False), {"foreach": False}), + ( + OptimizerConfig(fused=True, foreach=False), + {"fused": True, "foreach": False}, + ), + ] + + for optimizer, expected in cases: + with self.subTest(optimizer=optimizer): + result = _build_optimizer_config(optimizer, "fsdp2", 100) + self.assertEqual(result["override_optimizer_config"], expected) + + def test_megatron_ignores_torch_optimizer_backend_options(self): + optimizer = OptimizerConfig(fused=True, foreach=False) + + result = _build_optimizer_config(optimizer, "megatron", 100) + + self.assertIsNone(result["override_optimizer_config"]) + def test_chat_template_path(self): config = get_template_config() config.model.chat_template_path = "tests/template/custom_chat_template.j2" diff --git a/trinity/common/config.py b/trinity/common/config.py index 710d9d242d2..d3583ffde6d 100644 --- a/trinity/common/config.py +++ b/trinity/common/config.py @@ -96,6 +96,8 @@ class OptimizerConfig: warmup_style: Optional[str] = None # deprecated ! lr_scheduler_type: str = "constant" optimizer_type: str = "adam" + fused: Optional[bool] = None + foreach: Optional[bool] = None betas: List[float] = field(default_factory=lambda: [0.9, 0.999]) weight_decay: float = 0.01 clip_grad: float = 1.0 diff --git a/trinity/trainer/verl/config.py b/trinity/trainer/verl/config.py index da155b8dbf4..f0dc50eb67c 100644 --- a/trinity/trainer/verl/config.py +++ b/trinity/trainer/verl/config.py @@ -586,7 +586,15 @@ def _build_optimizer_config( optim["min_lr_ratio"] = trinity_optim.min_lr_ratio optim["lr_scheduler_type"] = trinity_optim.lr_scheduler_type optim["num_cycles"] = 0.5 - optim["override_optimizer_config"] = None + optimizer_overrides = { + key: value + for key, value in { + "fused": trinity_optim.fused, + "foreach": trinity_optim.foreach, + }.items() + if value is not None + } + optim["override_optimizer_config"] = optimizer_overrides or None optim["zero_indexed_step"] = True else: # Megatron uses McoreOptimizerConfig From 295d6b81fa1e6e0c9492472880921e6742e4a89d Mon Sep 17 00:00:00 2001 From: charlesxu91 Date: Sat, 15 Aug 2026 22:27:04 +0800 Subject: [PATCH 2/2] docs: explain optimizer backend controls --- docs/sphinx_doc/source/tutorial/trinity_configs.md | 4 ++++ docs/sphinx_doc/source_zh/tutorial/trinity_configs.md | 4 ++++ 2 files changed, 8 insertions(+) diff --git a/docs/sphinx_doc/source/tutorial/trinity_configs.md b/docs/sphinx_doc/source/tutorial/trinity_configs.md index e9793b0933a..2e99f1ccedc 100644 --- a/docs/sphinx_doc/source/tutorial/trinity_configs.md +++ b/docs/sphinx_doc/source/tutorial/trinity_configs.md @@ -102,6 +102,8 @@ algorithm: optimizer: lr: 1e-6 lr_scheduler_type: "constant" + fused: true + foreach: false # The following parameters are optional # If not specified, they will automatically be set based on the `algorithm_type` sample_strategy: "default" @@ -117,6 +119,8 @@ algorithm: - `lr`: Learning rate for actor. - `warmup_style`: Deprecated, use `lr_scheduler_type` instead. We will remove this field in future versions. - `lr_scheduler_type`: Learning rate scheduler type for actor model. Default is `constant`. Supported types: `constant`, `cosine`. + - `fused`: Optional PyTorch optimizer setting for FSDP/FSDP2 trainers. When unset, the backend default is used. This setting is ignored by Megatron trainers. + - `foreach`: Optional PyTorch optimizer setting for FSDP/FSDP2 trainers. When unset, the backend default is used. This setting is ignored by Megatron trainers. - `sample_strategy`: The sampling strategy used for loading experiences from experience buffer. Supported types: `default`, `staleness_control`, `mix`. - `advantage_fn`: The advantage function used for computing advantages. - `kl_penalty_fn`: The KL penalty function used for computing KL penalty applied in reward. diff --git a/docs/sphinx_doc/source_zh/tutorial/trinity_configs.md b/docs/sphinx_doc/source_zh/tutorial/trinity_configs.md index 681f9e4f03a..9bea432d801 100644 --- a/docs/sphinx_doc/source_zh/tutorial/trinity_configs.md +++ b/docs/sphinx_doc/source_zh/tutorial/trinity_configs.md @@ -102,6 +102,8 @@ algorithm: optimizer: lr: 1e-6 lr_scheduler_type: constant + fused: true + foreach: false # 以下参数为可选 # 若未指定,将根据 `algorithm_type` 自动设置 sample_strategy: "default" @@ -117,6 +119,8 @@ algorithm: - `lr`: 优化器的学习率。 - `warmup_style`:已弃用,请改用 `lr_scheduler_type`。该域将会在未来版本中移除。 - `lr_scheduler_type`:Actor 模型的学习率调度器类型。默认值为 `constant`。支持类型:`constant`、`cosine`。 + - `fused`:FSDP/FSDP2 trainer 的可选 PyTorch 优化器设置。未设置时使用后端默认值;Megatron trainer 会忽略该设置。 + - `foreach`:FSDP/FSDP2 trainer 的可选 PyTorch 优化器设置。未设置时使用后端默认值;Megatron trainer 会忽略该设置。 - `sample_strategy`: 从 experience buffer 加载 experience 时使用的采样策略。支持类型:`default`、`staleness_control`、`mix`。 - `advantage_fn`: 用于计算优势值的函数。 - `kl_penalty_fn`: 用于在奖励中计算 KL 惩罚的函数。