Skip to content
Merged
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
91 changes: 91 additions & 0 deletions tests/experimental/examples/common/run_trainer_node_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand All @@ -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)
Expand All @@ -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)
Expand Down
26 changes: 21 additions & 5 deletions tunix/experimental/examples/common/run_rollout_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)


Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
63 changes: 55 additions & 8 deletions tunix/experimental/examples/common/run_trainer_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)

Expand All @@ -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,
Expand All @@ -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,
)
Expand Down Expand Up @@ -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)
Expand Down
Loading