diff --git a/tunix/experimental/examples/math_gsm8k_dist/launcher.sh b/tunix/experimental/examples/math_gsm8k_dist/launcher.sh index 5d244b033..46138e8f0 100755 --- a/tunix/experimental/examples/math_gsm8k_dist/launcher.sh +++ b/tunix/experimental/examples/math_gsm8k_dist/launcher.sh @@ -397,6 +397,7 @@ echo "Launching trainer node on TPU chips $TRAINER_TPU_CHIPS..." --model_id="$MODEL_ID" --model_dir="$MODEL_DIR" --model_name="$MODEL_NAME" + --sampler="$SAMPLER" --tokenizer_path="$TOKENIZER_PATH" --max_prompt_length="$MAX_PROMPT_LENGTH" --max_response_length="$MAX_RESPONSE_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..bb69e67f5 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py @@ -70,6 +70,7 @@ 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("--sampler", type=str, default="vanilla") parser.add_argument("--mini_batch_size", type=int, default=1) parser.add_argument("--train_micro_batch_size", type=int, default=1) parser.add_argument("--compute_logps_micro_batch_size", type=int, default=1) @@ -364,6 +365,7 @@ def _factory(): actor_model, optax.adamw(learning_rate=args.learning_rate), training_config, + sampler_type=args.sampler, ) return _MeshBoundTrainer(trainer, mesh)