Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
92 changes: 84 additions & 8 deletions tests/experimental/orchestrator/rl_program_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand Down
16 changes: 15 additions & 1 deletion tunix/experimental/examples/math_gsm8k_dist/k8s_launcher.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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} \
Expand Down Expand Up @@ -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} \
Expand Down Expand Up @@ -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
Expand Down
18 changes: 16 additions & 2 deletions tunix/experimental/examples/math_gsm8k_dist/launcher.sh
Original file line number Diff line number Diff line change
Expand Up @@ -308,21 +308,33 @@ 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"
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"
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
46 changes: 40 additions & 6 deletions tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,6 @@
import asyncio
import contextlib
import logging
import math
import os
from pathlib import Path
import pickle
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand All @@ -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():
Expand Down
Loading
Loading