From 9d360298d836c0f02a9347125a31ca8381662803 Mon Sep 17 00:00:00 2001 From: Haoyu Gao Date: Thu, 10 Sep 2026 15:00:56 -0700 Subject: [PATCH] Expose DeepSWE controls in distributed workers. Support --enable_thinking for rollout nodes, and configure prompt-level mini-batch gradient accumulation and AdamW optimizer parameters in trainer nodes. PiperOrigin-RevId: 979416686 --- .../examples/common/run_trainer_node_test.py | 91 +++++++++++++++++++ .../examples/common/run_rollout_node.py | 26 +++++- .../examples/common/run_trainer_node.py | 63 +++++++++++-- 3 files changed, 167 insertions(+), 13 deletions(-) diff --git a/tests/experimental/examples/common/run_trainer_node_test.py b/tests/experimental/examples/common/run_trainer_node_test.py index 3da3e612a..c704c7c75 100644 --- a/tests/experimental/examples/common/run_trainer_node_test.py +++ b/tests/experimental/examples/common/run_trainer_node_test.py @@ -440,6 +440,12 @@ def test_parse_args_defaults_and_custom(self): self.assertEqual(args.mesh_tp, 1) self.assertEqual(args.checkpoint_save_interval_steps, 1) self.assertEqual(args.checkpoint_max_to_keep, 10) + self.assertEqual(args.mini_batch_size, 1) + self.assertEqual(args.num_generations, 1) + self.assertEqual(args.adam_b1, 0.9) + self.assertEqual(args.adam_b2, 0.999) + self.assertEqual(args.weight_decay, 0.0) + self.assertIsNone(args.max_grad_norm) self.assertFalse(args.use_lora) custom_argv = [ @@ -464,6 +470,18 @@ def test_parse_args_defaults_and_custom(self): "32", "--lora_alpha", "64.0", + "--adam_b2", + "0.99", + "--weight_decay", + "0.01", + "--max_grad_norm", + "1.0", + "--mini_batch_size", + "2", + "--num_generations", + "8", + "--train_micro_batch_size", + "4", ] args_custom = run_trainer_node._parse_args(custom_argv) self.assertEqual(args_custom.port, 20050) @@ -477,6 +495,79 @@ def test_parse_args_defaults_and_custom(self): self.assertTrue(args_custom.use_lora) self.assertEqual(args_custom.lora_rank, 32) self.assertEqual(args_custom.lora_alpha, 64.0) + self.assertEqual(args_custom.adam_b2, 0.99) + self.assertEqual(args_custom.weight_decay, 0.01) + self.assertEqual(args_custom.max_grad_norm, 1.0) + self.assertEqual(args_custom.mini_batch_size, 2) + self.assertEqual(args_custom.num_generations, 8) + self.assertEqual(args_custom.train_micro_batch_size, 4) + + def test_gradient_accumulation_uses_prompt_level_mini_batch(self): + # pylint: disable=protected-access + args = run_trainer_node._parse_args( + [ + "--mini_batch_size=2", + "--num_generations=8", + "--train_micro_batch_size=4", + ] + ) + + steps = run_trainer_node._gradient_accumulation_steps(args) + # pylint: enable=protected-access + + self.assertEqual(steps, 4) + + def test_gradient_accumulation_requires_exact_divisibility(self): + # pylint: disable=protected-access + args = run_trainer_node._parse_args( + [ + "--mini_batch_size=2", + "--num_generations=3", + "--train_micro_batch_size=4", + ] + ) + + with self.assertRaisesRegex(ValueError, "must be divisible"): + run_trainer_node._gradient_accumulation_steps(args) + # pylint: enable=protected-access + + def test_build_optimizer_applies_adamw_and_gradient_clipping(self): + args = run_trainer_node._parse_args( # pylint: disable=protected-access + [ + "--learning_rate=1e-6", + "--adam_b1=0.9", + "--adam_b2=0.99", + "--weight_decay=0.01", + "--max_grad_norm=1.0", + ] + ) + adamw = object() + clipped = object() + chained = object() + with mock.patch.object( + run_trainer_node.optax, "adamw", return_value=adamw + ) as mock_adamw: + with mock.patch.object( + run_trainer_node.optax, + "clip_by_global_norm", + return_value=clipped, + ) as mock_clip: + with mock.patch.object( + run_trainer_node.optax, "chain", return_value=chained + ) as mock_chain: + # pylint: disable=protected-access + optimizer = run_trainer_node._build_optimizer(args) + # pylint: enable=protected-access + + self.assertIs(optimizer, chained) + mock_adamw.assert_called_once_with( + learning_rate=1e-6, + b1=0.9, + b2=0.99, + weight_decay=0.01, + ) + mock_clip.assert_called_once_with(1.0) + mock_chain.assert_called_once_with(clipped, adamw) def test_create_mesh_validates_device_count(self): args = mock.MagicMock(mesh_fsdp=2, mesh_tp=2) diff --git a/tunix/experimental/examples/common/run_rollout_node.py b/tunix/experimental/examples/common/run_rollout_node.py index d35100e19..b7b3ae9f9 100644 --- a/tunix/experimental/examples/common/run_rollout_node.py +++ b/tunix/experimental/examples/common/run_rollout_node.py @@ -59,14 +59,16 @@ def _import_vllm_sampler(): return vllm_sampler -def _chat_parser_for(model_id: str, tokenizer): +def _chat_parser_for( + model_id: str, tokenizer: Any, *, enable_thinking: bool = False +): """Selects the chat template parser by model family.""" name = model_id.lower() for family, parser_cls in CHAT_PARSERS.items(): if family in name: - return parser_cls(tokenizer, enable_thinking=False) + return parser_cls(tokenizer, enable_thinking=enable_thinking) return chat_parser_lib.DefaultChatTemplateParser( - tokenizer, enable_thinking=False + tokenizer, enable_thinking=enable_thinking ) @@ -152,6 +154,12 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: default=int(os.getenv("ROLLOUT_MAX_CONCURRENCY", "64")), help="Maximum concurrent trajectory collections inside this worker.", ) + parser.add_argument( + "--enable_thinking", + action=argparse.BooleanOptionalAction, + default=False, + help="Enable the model family's thinking chat-template mode.", + ) parser.add_argument( "--weight_sync_mode", @@ -239,7 +247,11 @@ def _create_vanilla_worker(args, tokenizer): ) rollout_tokenizer = tokenizer_adapter_lib.TokenizerAdapter(tokenizer) - chat_parser = _chat_parser_for(args.model_id or args.model_name, tokenizer) + chat_parser = _chat_parser_for( + args.model_id or args.model_name, + tokenizer, + enable_thinking=args.enable_thinking, + ) return rollout_worker.RolloutWorker( worker_id=args.worker_id, config=config, @@ -267,7 +279,11 @@ def _create_vllm_worker(args, tokenizer): ) rollout_tokenizer = tokenizer_adapter_lib.TokenizerAdapter(tokenizer) - chat_parser = _chat_parser_for(args.model_id or args.model_name, tokenizer) + chat_parser = _chat_parser_for( + args.model_id or args.model_name, + tokenizer, + enable_thinking=args.enable_thinking, + ) logging.info("Creating RolloutWorker wrapper...") return rollout_worker.RolloutWorker( worker_id=args.worker_id, diff --git a/tunix/experimental/examples/common/run_trainer_node.py b/tunix/experimental/examples/common/run_trainer_node.py index 086233691..107714dbb 100644 --- a/tunix/experimental/examples/common/run_trainer_node.py +++ b/tunix/experimental/examples/common/run_trainer_node.py @@ -50,7 +50,7 @@ ) -def _build_actor_optimizer(args): +def _build_optimizer(args): """Builds the actor optimizer from CLI flags. Defaults reproduce the previous bare optax.adamw (no clipping, optax defaults @@ -91,8 +91,24 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: parser.add_argument("--mesh_expert", type=int, default=1) parser.add_argument("--max_prompt_length", type=int, default=512) parser.add_argument("--max_response_length", type=int, default=128) - parser.add_argument("--mini_batch_size", type=int, default=1) - parser.add_argument("--train_micro_batch_size", type=int, default=1) + parser.add_argument( + "--mini_batch_size", + type=int, + default=1, + help="Number of prompt groups per optimizer update.", + ) + parser.add_argument( + "--num_generations", + type=int, + default=1, + help="Number of rollout trajectories generated per prompt group.", + ) + parser.add_argument( + "--train_micro_batch_size", + type=int, + default=1, + help="Number of trajectories per forward/backward microbatch.", + ) parser.add_argument("--compute_logps_micro_batch_size", type=int, default=1) parser.add_argument("--compute_logps_chunk_size", type=int, default=0) parser.add_argument("--eval_every_n_steps", type=int, default=1000000) @@ -367,9 +383,31 @@ def _factory(): return _factory +def _gradient_accumulation_steps(args: argparse.Namespace) -> int: + if args.mini_batch_size <= 0: + raise ValueError("--mini_batch_size must be positive.") + if args.num_generations <= 0: + raise ValueError("--num_generations must be positive.") + if args.train_micro_batch_size <= 0: + raise ValueError("--train_micro_batch_size must be positive.") + update_trajectories = args.mini_batch_size * args.num_generations + if update_trajectories % args.train_micro_batch_size != 0: + raise ValueError( + "--mini_batch_size * --num_generations must be divisible by " + "--train_micro_batch_size; got " + f"mini_batch_size={args.mini_batch_size}, " + f"num_generations={args.num_generations}, " + f"train_micro_batch_size={args.train_micro_batch_size}." + ) + return update_trajectories // args.train_micro_batch_size + + def _create_tunix_trainer_factory(args) -> Any: """Creates the trainer factory function for Tunix's PeftTrainer.""" logging.info("Trainer backend: Tunix's PeftTrainer.") + grad_accumulation_steps = _gradient_accumulation_steps(args) + update_trajectories = args.mini_batch_size * args.num_generations + args.model_dir = _ensure_model_dir_for_trainer(args.model_dir, args.model_id) logging.info("Prepared trainer safetensors directory: %s", args.model_dir) @@ -381,9 +419,6 @@ def _create_tunix_trainer_factory(args) -> Any: actor_model = _load_actor_model(args, mesh, lora=args.use_lora) logging.info("Building PeftTrainer v2 config...") - grad_accumulation_steps = max( - 1, math.ceil(args.mini_batch_size / args.train_micro_batch_size) - ) checkpointing_options = ocp.CheckpointManagerOptions( save_interval_steps=args.checkpoint_save_interval_steps, max_to_keep=args.checkpoint_max_to_keep, @@ -402,15 +437,21 @@ def _create_tunix_trainer_factory(args) -> Any: resume_from_checkpoint_on_init=False, ) logging.info( - "PeftTrainer v2 gradient_accumulation_steps=%d.", + "PeftTrainer v2 gradient_accumulation_steps=%d " + "(mini_batch_size=%d prompt groups, num_generations=%d, " + "update_trajectories=%d, train_micro_batch_size=%d).", grad_accumulation_steps, + args.mini_batch_size, + args.num_generations, + update_trajectories, + args.train_micro_batch_size, ) def _factory(): with mesh: trainer = peft_trainer_v2.PeftTrainer( actor_model, - _build_actor_optimizer(args), + _build_optimizer(args), training_config, sampler_type=args.sampler_type, ) @@ -450,6 +491,12 @@ def main(argv: list[str], context: Any = None) -> None: if args.train_micro_batch_size <= 0: raise ValueError("--train_micro_batch_size must be positive.") + if args.mini_batch_size <= 0: + raise ValueError("--mini_batch_size must be positive.") + if args.num_generations <= 0: + raise ValueError("--num_generations must be positive.") + if args.max_grad_norm is not None and args.max_grad_norm <= 0: + raise ValueError("--max_grad_norm must be positive when specified.") logging.info("Creating generic TrainerWorker and gRPC server...") trainer_factory = _create_trainer_factory(args)