diff --git a/tests/experimental/orchestrator/rl_program_test.py b/tests/experimental/orchestrator/rl_program_test.py index 91ba2a861..79a950d90 100644 --- a/tests/experimental/orchestrator/rl_program_test.py +++ b/tests/experimental/orchestrator/rl_program_test.py @@ -280,6 +280,7 @@ def test_initialization(self): self.assertEqual(program.step, 0) self.assertEqual(program.group_size, 2) self.assertEqual(program.mini_batch_size, 1) + self.assertEqual(program.full_batch_size, 1) self.assertIsNotNone(program.raw_q) self.assertIsNotNone(program.scored_q) @@ -757,30 +758,105 @@ def tracking_create_payloads(step_items, **kwargs): asyncio.run(_run()) def test_multi_group_mini_batch_gradient_accumulation(self): + class SingleMicrobatchAssembler: + group_size: int = 2 + groups_per_assembly_batch: int = 2 + + @property + def assembly_batch_size(self) -> int: + return self.groups_per_assembly_batch * self.group_size + + def __init__(self): + self.item_counts = [] + + def pack(self, items): + self.item_counts.append(len(items)) + return ["microbatch"] + async def _run(): self.mock_algo.mini_batch_size = 2 + assembler = SingleMicrobatchAssembler() _set_mock_poll_batches( self.mock_engine, _make_trajectory_group("prompt_0"), _make_trajectory_group("prompt_1"), ) - program = self._create_program(dataset=["p0", "p1"]) + program = self._create_program( + dataset=["p0", "p1"], assembler=assembler + ) await program.run_async(self.mock_engine) - self.assertEqual(self.mock_engine.train_step.call_count, 2) + self.assertEqual(assembler.item_counts, [4]) + self.assertEqual(self.mock_engine.train_step.call_count, 1) calls = self.mock_engine.train_step.call_args_list - # First group: accumulate_gradients=True, apply_optimizer=False self.assertTrue(calls[0].kwargs["accumulate_gradients"]) - self.assertFalse(calls[0].kwargs["apply_optimizer"]) - # Second group: accumulate_gradients=True, apply_optimizer=True - self.assertTrue(calls[1].kwargs["accumulate_gradients"]) - self.assertTrue(calls[1].kwargs["apply_optimizer"]) + self.assertTrue(calls[0].kwargs["apply_optimizer"]) self.assertEqual(program.last_step_result.num_rollouts, 4) - self.assertEqual(program.last_step_result.num_microbatches, 2) + self.assertEqual(program.last_step_result.num_microbatches, 1) + + asyncio.run(_run()) + + def test_full_batch_sync_boundary_spans_multiple_updates(self): + class TwoMicrobatchAssembler: + group_size: int = 2 + groups_per_assembly_batch: int = 2 + + @property + def assembly_batch_size(self) -> int: + return self.groups_per_assembly_batch * self.group_size + + def __init__(self): + self.item_counts = [] + + def pack(self, items): + self.item_counts.append(len(items)) + return ["microbatch_0", "microbatch_1"] + + async def _run(): + self.mock_algo.mini_batch_size = 2 + assembler = TwoMicrobatchAssembler() + _set_mock_poll_batches( + self.mock_engine, + _make_trajectory_group("prompt_0"), + _make_trajectory_group("prompt_1"), + _make_trajectory_group("prompt_2"), + _make_trajectory_group("prompt_3"), + ) + program = self._create_program( + dataset=["p0", "p1", "p2", "p3"], + assembler=assembler, + batch_size=4, + sync_weights=True, + ) + + await program.run_async(self.mock_engine) + + self.assertEqual(assembler.item_counts, [4, 4]) + self.assertEqual(self.mock_engine.train_step.call_count, 4) + self.assertEqual( + [ + call.kwargs["apply_optimizer"] + for call in self.mock_engine.train_step.call_args_list + ], + [False, True, False, True], + ) + self.assertEqual(self.mock_engine.save_checkpoint.call_count, 2) + self.mock_engine.sync_weights.assert_called_once_with( + role=datatypes.Role.ACTOR + ) + self.assertEqual(program.step, 1) + self.assertEqual(program.last_step_result.num_rollouts, 8) + self.assertEqual(program.last_step_result.num_microbatches, 4) + self.assertEqual(program.last_step_result.policy_version, 1) asyncio.run(_run()) + def test_full_batch_size_must_be_divisible_by_mini_batch_size(self): + self.mock_algo.mini_batch_size = 2 + with self.assertRaisesRegex(ValueError, "batch_size must be divisible"): + self._create_program(batch_size=3) + def test_reference_kl_logprobs_scoring_in_train_stage(self): async def _run(): self.mock_algo.requires_reference_kl = True diff --git a/tunix/experimental/examples/math_gsm8k_dist/k8s_launcher.sh b/tunix/experimental/examples/math_gsm8k_dist/k8s_launcher.sh index a1c39a82f..524e1f376 100755 --- a/tunix/experimental/examples/math_gsm8k_dist/k8s_launcher.sh +++ b/tunix/experimental/examples/math_gsm8k_dist/k8s_launcher.sh @@ -34,7 +34,7 @@ export TRAIN_MICRO_BATCH_SIZE=${TRAIN_MICRO_BATCH_SIZE:-1} # Set to tunix to run Tunix's PeftTrainer, and maxtext to run MaxText's MaxTextTrainingEngine export TRAINER_BACKEND=${TRAINER_BACKEND:-tunix} -export MINI_BATCH_SIZE=${MINI_BATCH_SIZE:-$((BATCH_SIZE * NUM_GENERATIONS))} +export MINI_BATCH_SIZE=${MINI_BATCH_SIZE:-1} export EVAL_EVERY_N_STEPS=${EVAL_EVERY_N_STEPS:-1000000} export LORA_RANK=${LORA_RANK:-16} export LORA_ALPHA=${LORA_ALPHA:-16.0} @@ -108,6 +108,7 @@ start_orchestrator() { --model_id=${MODEL_ID} \ --tokenizer_path=${TOKENIZER_PATH} \ --batch_size=${BATCH_SIZE} \ + --mini_batch_size=${MINI_BATCH_SIZE} \ --num_generations=${NUM_GENERATIONS} \ --max_steps=${MAX_STEPS} \ --max_prompt_length=${MAX_PROMPT_LENGTH} \ @@ -162,6 +163,7 @@ start_trainer() { --max_prompt_length=${MAX_PROMPT_LENGTH} \ --max_response_length=${MAX_RESPONSE_LENGTH} \ --mini_batch_size=${MINI_BATCH_SIZE} \ + --num_generations=${NUM_GENERATIONS} \ --train_micro_batch_size=${TRAIN_MICRO_BATCH_SIZE} \ --eval_every_n_steps=${EVAL_EVERY_N_STEPS} \ --lora_rank=${LORA_RANK} \ @@ -251,6 +253,18 @@ if [[ -z "$TUNIX_IMAGE" ]]; then exit 1 fi +if (( BATCH_SIZE % MINI_BATCH_SIZE != 0 )); then + echo "Error: BATCH_SIZE must be divisible by MINI_BATCH_SIZE." + echo " BATCH_SIZE=$BATCH_SIZE MINI_BATCH_SIZE=$MINI_BATCH_SIZE" + exit 1 +fi + +if (( (MINI_BATCH_SIZE * NUM_GENERATIONS) % TRAIN_MICRO_BATCH_SIZE != 0 )); then + echo "Error: MINI_BATCH_SIZE * NUM_GENERATIONS must be divisible by TRAIN_MICRO_BATCH_SIZE." + echo " MINI_BATCH_SIZE=$MINI_BATCH_SIZE NUM_GENERATIONS=$NUM_GENERATIONS TRAIN_MICRO_BATCH_SIZE=$TRAIN_MICRO_BATCH_SIZE" + exit 1 +fi + if [[ "$COMMAND" == "start" ]]; then stop_orchestrator stop_trainer diff --git a/tunix/experimental/examples/math_gsm8k_dist/launcher.sh b/tunix/experimental/examples/math_gsm8k_dist/launcher.sh index 5d244b033..37b78ec8e 100755 --- a/tunix/experimental/examples/math_gsm8k_dist/launcher.sh +++ b/tunix/experimental/examples/math_gsm8k_dist/launcher.sh @@ -308,6 +308,18 @@ PY done } +if (( BATCH_SIZE % MINI_BATCH_SIZE != 0 )); then + echo "Error: BATCH_SIZE must be divisible by MINI_BATCH_SIZE." + echo " BATCH_SIZE=$BATCH_SIZE MINI_BATCH_SIZE=$MINI_BATCH_SIZE" + exit 1 +fi + +if (( (MINI_BATCH_SIZE * NUM_GENERATIONS) % TRAIN_MICRO_BATCH_SIZE != 0 )); then + echo "Error: MINI_BATCH_SIZE * NUM_GENERATIONS must be divisible by TRAIN_MICRO_BATCH_SIZE." + echo " MINI_BATCH_SIZE=$MINI_BATCH_SIZE NUM_GENERATIONS=$NUM_GENERATIONS TRAIN_MICRO_BATCH_SIZE=$TRAIN_MICRO_BATCH_SIZE" + exit 1 +fi + echo "==================================================" echo "Starting distributed GSM8K GRPO chain demo locally" echo " rollout engine: vLLM" @@ -315,14 +327,14 @@ echo " model dir: $MODEL_DIR" echo " tokenizer path: $TOKENIZER_PATH" echo " python: $PYTHON_BIN" echo " trajectories: $((BATCH_SIZE * NUM_GENERATIONS)) per step" -echo " batch size: $BATCH_SIZE" +echo " batch size: $BATCH_SIZE prompt groups/full step" echo " generations: $NUM_GENERATIONS" echo " max steps: $MAX_STEPS" echo " eval interval: $EVAL_EVERY_N_STEPS" echo " prompt length: $MAX_PROMPT_LENGTH" echo " response len: $MAX_RESPONSE_LENGTH" echo " train micro: $TRAIN_MICRO_BATCH_SIZE" -echo " mini batch: $MINI_BATCH_SIZE" +echo " mini batch: $MINI_BATCH_SIZE prompt groups/update ($((MINI_BATCH_SIZE * NUM_GENERATIONS)) trajectories)" echo " beta: $BETA" echo " epsilon: $EPSILON" echo " reward mode: $REWARD_MODE" @@ -401,6 +413,7 @@ echo "Launching trainer node on TPU chips $TRAINER_TPU_CHIPS..." --max_prompt_length="$MAX_PROMPT_LENGTH" --max_response_length="$MAX_RESPONSE_LENGTH" --mini_batch_size="$MINI_BATCH_SIZE" + --num_generations="$NUM_GENERATIONS" --train_micro_batch_size="$TRAIN_MICRO_BATCH_SIZE" --eval_every_n_steps="$EVAL_EVERY_N_STEPS" --lora_rank="$LORA_RANK" @@ -625,6 +638,7 @@ echo "Launching CPU orchestrator..." --model_id="$MODEL_ID" --tokenizer_path="$TOKENIZER_PATH" --batch_size="$BATCH_SIZE" + --mini_batch_size="$MINI_BATCH_SIZE" --num_generations="$NUM_GENERATIONS" --max_steps="$MAX_STEPS" --max_prompt_length="$MAX_PROMPT_LENGTH" diff --git a/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py b/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py index 7d93c0a25..504a8cd28 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py @@ -75,7 +75,13 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: "--batch_size", type=int, default=4, - help="Number of prompt groups per step.", + help="Number of prompt groups per full/global step.", + ) + parser.add_argument( + "--mini_batch_size", + type=int, + default=2, + help="Number of prompt groups per optimizer update.", ) parser.add_argument("--num_generations", type=int, default=8) parser.add_argument("--max_steps", type=int, default=1) @@ -266,8 +272,7 @@ def _grpo_model_input( def _build_algo(args: argparse.Namespace) -> algorithm_adapter.GRPOAdapter: algo = algorithm_adapter.GRPOAdapter( group_size=args.num_generations, - # StandardRLProgram consumes this many prompt groups per trainer update. - mini_batch_size=args.batch_size, + mini_batch_size=args.mini_batch_size, max_packed_len=args.max_prompt_length + args.max_response_length, clip_epsilon=args.epsilon, beta_kl=args.beta, @@ -437,21 +442,41 @@ def main(argv: list[str], context: Any = None) -> None: raise ValueError("num_generations must be greater than 1 for GRPO.") if args.batch_size <= 0: raise ValueError("batch_size must be positive.") + if args.mini_batch_size <= 0: + raise ValueError("mini_batch_size must be positive.") + if args.batch_size % args.mini_batch_size != 0: + raise ValueError( + "batch_size must be divisible by mini_batch_size; got " + f"batch_size={args.batch_size}, " + f"mini_batch_size={args.mini_batch_size}." + ) 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}." + ) if args.max_staleness < 0: raise ValueError("offpolicy/max_staleness must be non-negative.") logging.info("=== Starting Distributed GSM8K GRPO Orchestrator ===") logging.info( - "Configuration: model_id=%s, batch_size=%d (prompt groups), " - "num_generations=%d (%d rollouts/step), max_steps=%d, " + "Configuration: model_id=%s, batch_size=%d prompt groups/full step, " + "mini_batch_size=%d prompt groups/update, num_generations=%d " + "(%d rollouts/full step, %d rollouts/update), max_steps=%d, " "train_micro_batch_size=%d, beta=%.4f, epsilon=%.2f, reward_mode=%s, " "max_staleness=%d, weight_sync_mode=%s.", args.model_id, args.batch_size, + args.mini_batch_size, args.num_generations, args.batch_size * args.num_generations, + args.mini_batch_size * args.num_generations, args.max_steps, args.train_micro_batch_size, args.beta, @@ -583,6 +608,7 @@ def accept_worker(hostname: str, _: int, metadata: bytes) -> None: dataset=_iter_prompt_items(args), max_steps=args.max_steps, reward_fns=reward_fns, + batch_size=args.batch_size, assembler=batch_assembly.PaddedBatchAssembler( batch_size=args.train_micro_batch_size, max_prompt_length=args.max_prompt_length, diff --git a/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py b/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py index e9760ade9..126d56b63 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py @@ -20,7 +20,6 @@ import asyncio import contextlib import logging -import math import os from pathlib import Path import pickle @@ -70,8 +69,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=8, + help="Number of trajectories generated for each 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) @@ -337,8 +352,21 @@ 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) + 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.") + 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; " + f"got mini_batch_size={args.mini_batch_size}, " + f"num_generations={args.num_generations}, " + f"train_micro_batch_size={args.train_micro_batch_size}." + ) + grad_accumulation_steps = ( + update_trajectories // args.train_micro_batch_size ) checkpointing_options = ocp.CheckpointManagerOptions( save_interval_steps=args.checkpoint_save_interval_steps, @@ -354,8 +382,14 @@ def _create_tunix_trainer_factory(args) -> Any: checkpoint_root_directory=args.checkpoint_root_directory, ) 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(): diff --git a/tunix/experimental/orchestrator/rl_program.py b/tunix/experimental/orchestrator/rl_program.py index df7f95502..91ef03610 100644 --- a/tunix/experimental/orchestrator/rl_program.py +++ b/tunix/experimental/orchestrator/rl_program.py @@ -100,6 +100,7 @@ def __init__( reward_fns: Sequence[Callable[..., Any]] | None = None, assembler: batch_assembly.BatchAssembler | None = None, group_size: int = 8, + batch_size: int | None = None, mini_batch_size: int = 4, max_staleness: int = 0, sync_weights: bool = True, @@ -121,6 +122,18 @@ def __init__( self.mini_batch_size = getattr(algo, "mini_batch_size", mini_batch_size) if self.mini_batch_size <= 0 or self.group_size <= 0: raise ValueError("mini_batch_size and group_size must be positive.") + self.full_batch_size = ( + self.mini_batch_size if batch_size is None else batch_size + ) + self.batch_size = self.full_batch_size + if self.full_batch_size <= 0: + raise ValueError("batch_size must be positive.") + if self.full_batch_size % self.mini_batch_size != 0: + raise ValueError( + "batch_size must be divisible by mini_batch_size; got " + f"batch_size={self.full_batch_size}, " + f"mini_batch_size={self.mini_batch_size}." + ) self.assembler = assembler or batch_assembly.SequencePackedBatchAssembler( group_size=self.group_size, max_packed_len=getattr(algo, "max_packed_len", 8192), @@ -603,94 +616,109 @@ async def train_stage(self) -> None: groups_per_assembly_batch = self.assembler.groups_per_assembly_batch - while groups_consumed < self.mini_batch_size: - groups_to_fetch = min( - groups_per_assembly_batch, self.mini_batch_size - groups_consumed - ) - scored_items = await self.scored_q.get_batch(num_groups=groups_to_fetch) - if not scored_items: - break + num_mini_batches = self.full_batch_size // self.mini_batch_size + for _ in range(num_mini_batches): + mini_batch_groups_consumed = 0 + while mini_batch_groups_consumed < self.mini_batch_size: + groups_to_fetch = min( + groups_per_assembly_batch, + self.mini_batch_size - mini_batch_groups_consumed, + ) + scored_items = await self.scored_q.get_batch( + num_groups=groups_to_fetch + ) + if not scored_items: + break - if groups_consumed == 0 and self.on_step_begin: - self.on_step_begin(current_step) - - groups_consumed += groups_to_fetch - uncommitted_groups.append(scored_items) - all_step_items.extend(scored_items) - num_rollouts += len(scored_items) - for item in scored_items: - step_rewards.append(float(getattr(item.traj, "reward", 0.0))) - - prompt_ids = list( - dict.fromkeys( - item.prompt_id - for item in scored_items - if getattr(item, "prompt_id", None) - ) - ) - payloads = [getattr(item, "payload", None) for item in scored_items] - # TODO(tunix-dev): Implement streaming microbatch assembly to overlap - # packing with trainer execution. - microbatches = self.assembler.pack(payloads) # pyrefly: ignore[bad-argument-type] - logging.info( - "Packed %d prompt groups into %d microbatches (total_rollouts=%d)." - " All prompt groups ids packed: %s", - len(prompt_ids), - len(microbatches), - len(payloads), - logging_utils.summarize_list(prompt_ids), - ) - if getattr(self.algo, "requires_reference_kl", False): - scored_microbatches = [] - for batch in microbatches: - if not isinstance(batch, datatypes.RLTrainerPayload): - raise TypeError( - "Reference KL requires an assembler that returns " - "datatypes.RLTrainerPayload microbatches; got " - f"{type(batch).__name__}." + if groups_consumed == 0 and self.on_step_begin: + self.on_step_begin(current_step) + + groups_fetched = len(scored_items) // self.group_size + mini_batch_groups_consumed += groups_fetched + groups_consumed += groups_fetched + uncommitted_groups.append(scored_items) + all_step_items.extend(scored_items) + num_rollouts += len(scored_items) + for item in scored_items: + step_rewards.append(float(getattr(item.traj, "reward", 0.0))) + + prompt_ids = list( + dict.fromkeys( + item.prompt_id + for item in scored_items + if getattr(item, "prompt_id", None) ) - ref_logps = await self.engine.per_token_logps( - datatypes.Role.REFERENCE, items=batch - ) - scored_microbatches.append( - batch_assembly.with_ref_per_token_logps(batch, ref_logps) - ) - microbatches = scored_microbatches - - num_microbatches += len(microbatches) - is_final_group = groups_consumed >= self.mini_batch_size - for batch_idx, batch in enumerate(microbatches): - is_final_batch = is_final_group and batch_idx == len(microbatches) - 1 - step_result = await self.engine.train_step( - batch, - role=datatypes.Role.ACTOR, - accumulate_gradients=True, - apply_optimizer=is_final_batch, ) - if is_final_batch: - # TODO(tunix-dev): Current checkpoint and metrics logic only works - # for fully on-policy. We need to come up with a solution for - # semi-off-policy where a single full batch has multiple mini - # batches. - trainer_metrics = await self.engine.get_metrics( - role=datatypes.Role.ACTOR + payloads = [getattr(item, "payload", None) for item in scored_items] + # TODO(tunix-dev): Implement streaming microbatch assembly to overlap + # packing with trainer execution. + microbatches = self.assembler.pack(payloads) # pyrefly: ignore[bad-argument-type] + logging.info( + "Packed %d prompt groups into %d microbatches " + "(total_rollouts=%d). All prompt groups ids packed: %s", + len(prompt_ids), + len(microbatches), + len(payloads), + logging_utils.summarize_list(prompt_ids), + ) + if getattr(self.algo, "requires_reference_kl", False): + scored_microbatches = [] + for batch in microbatches: + if not isinstance(batch, datatypes.RLTrainerPayload): + raise TypeError( + "Reference KL requires an assembler that returns " + "datatypes.RLTrainerPayload microbatches; got " + f"{type(batch).__name__}." + ) + ref_logps = await self.engine.per_token_logps( + datatypes.Role.REFERENCE, items=batch + ) + scored_microbatches.append( + batch_assembly.with_ref_per_token_logps(batch, ref_logps) + ) + microbatches = scored_microbatches + + num_microbatches += len(microbatches) + is_final_assembly_batch = ( + mini_batch_groups_consumed >= self.mini_batch_size + ) + for batch_idx, batch in enumerate(microbatches): + is_final_batch = ( + is_final_assembly_batch + and batch_idx == len(microbatches) - 1 ) - # TODO(tunix-dev): Configurable checkpointing frequency. Today we - # checkpoint at the same frequency as the weight update. - # TODO(tunix-dev): For now any failures in save_checkpoint will - # abort the entire program. Make it configurable on whether to fail - # or continue. - await self.engine.save_checkpoint( + step_result = await self.engine.train_step( + batch, role=datatypes.Role.ACTOR, - metadata={ - "step": self.step + 1, - "policy_version": self.policy_version, - "num_rollouts": num_rollouts, - "num_microbatches": num_microbatches, - }, + accumulate_gradients=True, + apply_optimizer=is_final_batch, ) + if is_final_batch: + # TODO(tunix-dev): Current checkpoint and metrics logic only works + # for fully on-policy. We need to come up with a solution for + # semi-off-policy where a single full batch has multiple mini + # batches. + trainer_metrics = await self.engine.get_metrics( + role=datatypes.Role.ACTOR + ) + # TODO(tunix-dev): Configurable checkpointing frequency. Today we + # checkpoint at the same frequency as the weight update. + # TODO(tunix-dev): For now any failures in save_checkpoint will + # abort the entire program. Make it configurable on whether to + # fail or continue. + await self.engine.save_checkpoint( + role=datatypes.Role.ACTOR, + metadata={ + "step": self.step + 1, + "policy_version": self.policy_version, + "num_rollouts": num_rollouts, + "num_microbatches": num_microbatches, + }, + ) + if mini_batch_groups_consumed != self.mini_batch_size: + break - if not scored_items: + if groups_consumed != self.full_batch_size: # TODO: We currently silently drop in-progress partial microbatch accumulators if # the dataset ends early. We may need to force-apply gradients here instead. logging.info( @@ -770,7 +798,7 @@ async def run_async( policy_version=self.policy_version, ) - max_groups_ahead = self.mini_batch_size * (self.max_staleness + 1) + max_groups_ahead = self.full_batch_size * (self.max_staleness + 1) self._dispatch_capacity = asyncio.Semaphore(max_groups_ahead) train_task = asyncio.create_task(self.train_stage())