From 6336d4cee75025f9013d0ef15d7f502791d4eb48 Mon Sep 17 00:00:00 2001 From: Leilei Wu Date: Thu, 24 Sep 2026 19:20:19 +0000 Subject: [PATCH] Implement MLPerf v6.1 RCP logging and deferred offline evaluation for DeepSWE --- tests/examples/deepswe_mllog_utils_test.py | 158 +++++++++++- .../deepswe_dist/eval_deepswe_test.py | 30 +++ .../orchestrator/rl_program_test.py | 42 ++++ .../worker/trainer_worker_test.py | 18 ++ tests/utils/maxtext_utils_test.py | 2 + .../examples/deepswe_dist/eval_deepswe.py | 127 +++++++++- .../examples/deepswe_dist/k8s_launcher.sh | 14 ++ .../examples/deepswe_dist/run_deepswe_dist.py | 74 +++++- tunix/experimental/examples/recipes/README.md | 1 + .../examples/recipes/mlperf_35b_128_v5p.sh | 8 + .../examples/recipes/mlperf_35b_eval.sh | 78 ++++++ .../examples/recipes/mlperf_base.sh | 121 +++++++++- tunix/experimental/orchestrator/rl_program.py | 26 +- tunix/experimental/train/abstract_trainer.py | 5 + tunix/experimental/train/peft_trainer_v2.py | 5 + tunix/experimental/worker/trainer_worker.py | 12 +- tunix/utils/maxtext_utils.py | 1 + tunix/utils/mllog_utils.py | 225 +++++++++++++++++- 18 files changed, 925 insertions(+), 22 deletions(-) create mode 100755 tunix/experimental/examples/recipes/mlperf_35b_eval.sh diff --git a/tests/examples/deepswe_mllog_utils_test.py b/tests/examples/deepswe_mllog_utils_test.py index 0e0fe2f104..a9971678f8 100644 --- a/tests/examples/deepswe_mllog_utils_test.py +++ b/tests/examples/deepswe_mllog_utils_test.py @@ -22,6 +22,7 @@ import shutil import tempfile import types +from typing import Any from unittest import mock import numpy as np @@ -30,6 +31,15 @@ from tunix.utils import mllog_utils +def _read_mllog_events(path: str) -> list[dict[str, Any]]: + with open(path, "r", encoding="utf-8") as f: + return [ + json.loads(line.split(":::MLLOG ", 1)[1]) + for line in f + if ":::MLLOG " in line + ] + + @absltest.skipIf(mllog_utils.mllogger is None, "mlperf_logging is not installed") class MllogUtilsTest(absltest.TestCase): @@ -64,6 +74,7 @@ def test_file_logging_with_metric_logger_dir(self): tpu_topology="v5p-64", rollout_engine="vllm", target_accuracy=0.69, + model_id="", ) mllog_utils.init_start(args) @@ -221,6 +232,7 @@ def test_end_to_end_mlperf_logging_with_train_configs(self): seed=1, learning_rate=1e-6, eval_every_n_steps=5, + model_id="", ) mock_train_dataset = [None] * 5480 @@ -602,7 +614,151 @@ def test_train_stop(self): content = f.read() self.assertIn('"key": "block_stop"', content) - self.assertIn('"key": "run_stop"', content) + # run_stop is emitted by the offline evaluator, not by train_stop. + self.assertNotIn('"key": "run_stop"', content) + + def test_compute_val_start_step(self): + self.assertEqual(mllog_utils.compute_val_start_step(256), 18) + self.assertEqual(mllog_utils.compute_val_start_step(512), 10) + self.assertEqual(mllog_utils.compute_val_start_step(1024), 7) + self.assertEqual(mllog_utils.compute_val_start_step(256, 1), 1) + self.assertEqual(mllog_utils.compute_val_start_step(256, 0), 18) + with self.assertRaises(ValueError): + mllog_utils.compute_val_start_step(0) + + def test_append_checkpoint_manifest_upserts_sorted_records(self): + manifest_path = os.path.join( + self.test_dir, "mllog", "eval_checkpoints.jsonl" + ) + for step, ts_ms in ((19, 2000), (18, 1000), (18, 1500)): + mllog_utils.append_checkpoint_manifest( + manifest_path, + { + "step": step, + "checkpoint_path": f"gs://ckpt/{step}/model_params", + "timestamp_ms": ts_ms, + "samples_count": step * 256, + "val_start_at": 18, + }, + ) + with open(manifest_path, "r", encoding="utf-8") as f: + records = [json.loads(line) for line in f] + self.assertEqual([r["step"] for r in records], [18, 19]) + self.assertEqual([r["timestamp_ms"] for r in records], [1500, 2000]) + + def test_configure_logger_downloads_existing_gcs_log_before_config(self): + calls = mock.MagicMock() + fake_mllogger = mock.MagicMock() + fake_mllogger.logger.handlers = [] + with ( + mock.patch.object(mllog_utils, "mllog", calls.mllog), + mock.patch.object(mllog_utils, "mllogger", fake_mllogger), + mock.patch.object(mllog_utils, "_is_master_process", return_value=True), + mock.patch.object(mllog_utils, "_gcs_target_path", None), + mock.patch.object(mllog_utils, "_local_log_path", None), + mock.patch.object( + mllog_utils, "_download_from_gcs_if_exists", calls.download + ), + ): + mllog_utils.configure_logger(metric_logger_dir="gs://b/mllog", seed=42) + local_path = mllog_utils._local_log_path # pylint: disable=protected-access + + self.assertEqual( + [c[0] for c in calls.mock_calls], ["download", "mllog.config"] + ) + calls.download.assert_called_once_with("gs://b/mllog/seed_42.out", local_path) + self.assertEqual( + os.path.abspath(calls.mllog.config.call_args.kwargs["filename"]), + local_path, + ) + + def test_offline_eval_rcp_sequence_converged(self): + log_dir = self.test_dir + args = types.SimpleNamespace(batch_size=16, num_generations=16) + mllog_utils.configure_logger(metric_logger_dir=log_dir, seed=42) + mllog_utils.train_stop(args, step=19, time_ms=2000) + self.assertFalse( + mllog_utils.log_offline_eval_step( + step=18, + samples_count=4608, + eval_accuracy=0.65, + checkpoint_timestamp_ms=1000, + ) + ) + self.assertTrue( + mllog_utils.log_offline_eval_step( + step=19, + samples_count=4864, + eval_accuracy=0.70, + checkpoint_timestamp_ms=2000, + is_last_checkpoint=True, + ) + ) + + events = _read_mllog_events(os.path.join(log_dir, "seed_42.out")) + self.assertEqual( + [e["key"] for e in events], + ["block_stop"] + ["eval_start", "eval_accuracy", "eval_stop"] * 2 + + ["run_stop"], + ) + self.assertEqual(events[0]["time_ms"], 2000) + self.assertEqual(events[0]["metadata"]["step"], 19) + self.assertEqual([events[2]["value"], events[5]["value"]], [0.65, 0.70]) + run_stop = events[-1] + self.assertEqual(run_stop["time_ms"], 2000) + self.assertEqual(run_stop["metadata"]["status"], "success") + self.assertEqual(run_stop["metadata"]["samples_count"], 4864) + + def test_offline_eval_rcp_sequence_not_converged(self): + log_dir = self.test_dir + mllog_utils.configure_logger(metric_logger_dir=log_dir, seed=42) + for step, ts_ms in ((18, 1000), (19, 2000), (20, 3000)): + self.assertFalse( + mllog_utils.log_offline_eval_step( + step=step, + samples_count=step * 256, + eval_accuracy=0.5, + checkpoint_timestamp_ms=ts_ms, + is_last_checkpoint=step == 20, + ) + ) + + events = _read_mllog_events(os.path.join(log_dir, "seed_42.out")) + run_stops = [e for e in events if e["key"] == "run_stop"] + self.assertLen(run_stops, 1) + self.assertEqual(events[-1]["key"], "run_stop") + self.assertEqual(run_stops[0]["time_ms"], 3000) + self.assertEqual(run_stops[0]["metadata"]["status"], "aborted") + self.assertEqual(run_stops[0]["metadata"]["samples_count"], 5120) + + def test_mlperf_6_1_0_init_print_disclosures(self): + fake_mllogger = mock.MagicMock() + with ( + mock.patch.object(mllog_utils, "mllogger", fake_mllogger), + mock.patch.object(mllog_utils, "_is_master_process", return_value=True), + mock.patch.object(mllog_utils, "_flush_to_gcs_if_needed"), + ): + args = types.SimpleNamespace( + batch_size=16, + num_generations=16, + max_prompt_length=4096, + max_response_length=61440, + model_id="", + ) + mllog_utils.init_print(args) + emitted = { + c.kwargs["key"]: c.kwargs["value"] + for c in fake_mllogger.event.call_args_list + } + self.assertEqual(emitted["eval_samples"], 251) + self.assertEqual(emitted["max_sequence_length"], 65536) + for key in ( + "lowest_numerical_precision_in_linear", + "lowest_numerical_precision_in_attn", + "lowest_numerical_precision_in_comm", + ): + self.assertEqual(emitted[key], "bfloat16") + self.assertEqual(emitted["config_filename"], "qwen35_397b_grpo") if __name__ == "__main__": diff --git a/tests/experimental/examples/deepswe_dist/eval_deepswe_test.py b/tests/experimental/examples/deepswe_dist/eval_deepswe_test.py index e68927dfb3..63c20e0f16 100644 --- a/tests/experimental/examples/deepswe_dist/eval_deepswe_test.py +++ b/tests/experimental/examples/deepswe_dist/eval_deepswe_test.py @@ -168,6 +168,21 @@ def test_missing_and_failed_attempts_are_not_dropped(self): with self.assertRaisesRegex(ValueError, "duplicate"): eval_lib.summarize(rows + [rows[0]], ["a", "b"], 4) + def test_pass_at_4_is_reported_with_more_attempts(self): + rows = [ + dict( + instance_id="a", + attempt=i, + reward=float(i == 5), + resolved=i == 5, + status="SUCCEEDED", + ) + for i in range(8) + ] + summary = eval_lib.summarize(rows, ["a"], 8) + # Unbiased pass@k = 1 - C(n - c, k) / C(n, k) with n=8, c=1. + self.assertEqual(summary["pass_at_k"], {"1": 1 / 8, "4": 1 / 2, "8": 1.0}) + def test_reward_comes_from_trajectory_not_completion_status(self): response = types.SimpleNamespace( error=None, @@ -182,6 +197,21 @@ def test_reward_comes_from_trajectory_not_completion_status(self): response.error = "infrastructure failure" self.assertFalse(eval_lib.compact_result(response)["resolved"]) + def test_rcp_logging_fails_fast_without_mlperf_logging(self): + from tunix.utils import mllog_utils # pylint: disable=g-import-not-at-top + + log_file = "gs://bucket/mllog/seed_42.out" + a = self.args("--rcp_logging=true", "--metric_logger_dir", log_file) + with mock.patch.object(mllog_utils, "configure_logger") as configure: + with mock.patch.object(mllog_utils, "mllogger", None): + with self.assertRaisesRegex(RuntimeError, "mlperf_logging"): + eval_lib.setup_rcp_logging(a) + configure.assert_not_called() + with mock.patch.object(mllog_utils, "mllogger", object()): + eval_lib.setup_rcp_logging(a) + eval_lib.setup_rcp_logging(self.args("--rcp_logging=false")) + configure.assert_called_once_with(metric_logger_dir=log_file, seed=42) + def test_individual_results_are_persisted(self): with tempfile.TemporaryDirectory() as directory: writer = eval_lib.ResultWriter(directory) diff --git a/tests/experimental/orchestrator/rl_program_test.py b/tests/experimental/orchestrator/rl_program_test.py index 929c8ae183..034240b2a5 100644 --- a/tests/experimental/orchestrator/rl_program_test.py +++ b/tests/experimental/orchestrator/rl_program_test.py @@ -1480,6 +1480,48 @@ async def _run(): asyncio.run(_run()) + def test_train_stage_on_checkpoint_saved_gets_path_from_response(self): + async def _run(): + self.mock_algo.num_generations = 1 + self.mock_algo.mini_batch_size = 1 + self.mock_engine.save_checkpoint.return_value = datatypes.Response( + metadata={ + "checkpoint_saved": True, + "checkpoint_path": "gs://ckpt/1/model_params", + } + ) + saved = [] + program = self._create_program( + batch_size=1, on_checkpoint_saved=saved.append + ) + program.engine = self.mock_engine + + payload = datatypes.RLTrainerPayload( + prompt_ids=np.array([1, 2], dtype=np.int32), + prompt_mask=np.array([1.0, 1.0], dtype=np.float32), + completion_ids=np.array([3, 4], dtype=np.int32), + completion_mask=np.array([1.0, 1.0], dtype=np.float32), + advantages=np.array([1.0, 1.0], dtype=np.float32), + ) + item = datatypes.TrajectoryItem( + group_index=0, + prompt_id="prompt_0", + start_step=0, + traj={"trajectory_reward": 1.0}, + ) + item.payload = payload + await program.scored_q.put(item) + await program.scored_q.close() + + await program.train_stage() + + self.assertLen(saved, 1) + self.assertEqual(saved[0]["step"], 1) + self.assertEqual(saved[0]["checkpoint_path"], "gs://ckpt/1/model_params") + self.assertEqual(saved[0]["timestamp_ms"], program.last_step_timestamp_ms) + + asyncio.run(_run()) + def test_train_stage_sequence_packed_final_batch_broken_down_into_multiple_microbatches( self, ): diff --git a/tests/experimental/worker/trainer_worker_test.py b/tests/experimental/worker/trainer_worker_test.py index 31f88fe21d..98e458fa86 100644 --- a/tests/experimental/worker/trainer_worker_test.py +++ b/tests/experimental/worker/trainer_worker_test.py @@ -44,6 +44,11 @@ def __init__(self): self.step_count = 10 self.target_state = None self.gen_model_input_fn = None + self._checkpoint_dir = None + + @property + def checkpoint_dir(self) -> str | None: + return self._checkpoint_dir def compile(self, dummy_data=None): pass @@ -152,6 +157,19 @@ def test_update_returns_step_count(self): step = self.worker.update() self.assertEqual(step, 11) + def test_save_checkpoint_returns_checkpoint_path_from_checkpoint_dir(self): + resp_empty = self.worker.save_checkpoint(metadata={"step": 5}) + self.assertTrue(resp_empty.metadata["checkpoint_saved"]) + self.assertEqual(resp_empty.metadata["checkpoint_path"], "") + + self.fake_trainer._checkpoint_dir = "gs://bucket/checkpoints" + resp = self.worker.save_checkpoint(metadata={"step": 5}) + self.assertTrue(resp.metadata["checkpoint_saved"]) + self.assertEqual( + resp.metadata["checkpoint_path"], + "gs://bucket/checkpoints/5/model_params", + ) + def test_set_target_state_configures_trainer(self): target_state = {"params": np.zeros((4, 4))} resp = self.worker.set_target_state(target_state=target_state) diff --git a/tests/utils/maxtext_utils_test.py b/tests/utils/maxtext_utils_test.py index cae894a112..218bfc1d56 100644 --- a/tests/utils/maxtext_utils_test.py +++ b/tests/utils/maxtext_utils_test.py @@ -913,6 +913,7 @@ def _init_state(self): mock_cfg = mock.MagicMock() mock_cfg.weight_dtype = "bfloat16" mock_cfg.float32_gate_logits = True + mock_cfg.checkpoint_dir = "/tmp/ckpts" mock_mesh = mock.MagicMock() with mock.patch.object( @@ -927,6 +928,7 @@ def _init_state(self): self.assertIsNot(captured["optimizer_cls_during_init"], FakeOptimizer) self.assertEqual(FakeNNX.Optimizer, FakeOptimizer) self.assertIsNone(engine._weight_converter._direct.target_dtype) + self.assertEqual(engine.checkpoint_dir, "/tmp/ckpts") if __name__ == "__main__": diff --git a/tunix/experimental/examples/deepswe_dist/eval_deepswe.py b/tunix/experimental/examples/deepswe_dist/eval_deepswe.py index 0c2cf1c984..645f428c86 100644 --- a/tunix/experimental/examples/deepswe_dist/eval_deepswe.py +++ b/tunix/experimental/examples/deepswe_dist/eval_deepswe.py @@ -128,6 +128,59 @@ def parse_args(argv=None): ) p.add_argument("--max_warmpool_size", type=int, default=1) p.add_argument("--output_dir", default="eval_results") + p.add_argument( + "--rcp_logging", + type=boolean, + nargs="?", + const=True, + default=os.environ.get("RCP_LOGGING", "0").lower() in ("1", "true"), + help="Enable MLPerf RCP (mllog) compliance logging.", + ) + p.add_argument( + "--metric_logger_dir", + default=os.environ.get("METRIC_LOGGER_DIR", ""), + help="Directory or GCS URI for MLPerf RCP output (seed_.out).", + ) + p.add_argument( + "--target_accuracy", + type=float, + default=float(os.environ.get("TARGET_ACCURACY", "0.69")), + help="Target evaluation accuracy for MLPerf RCP compliance logging.", + ) + p.add_argument( + "--checkpoint_step", + type=int, + default=int(os.environ.get("CHECKPOINT_STEP", "0")), + help="Optimizer step corresponding to the evaluated checkpoint.", + ) + p.add_argument( + "--checkpoint_timestamp_ms", + type=int, + default=( + int(os.environ["CHECKPOINT_TIMESTAMP_MS"]) + if os.environ.get("CHECKPOINT_TIMESTAMP_MS", "").strip() + else None + ), + help=( + "Training epoch timestamp (ms) when the checkpoint weights were " + "updated, used to backdate run_stop." + ), + ) + p.add_argument( + "--samples_count", + type=int, + default=int(os.environ.get("SAMPLES_COUNT", "0")), + help="Cumulative training samples at checkpoint_step.", + ) + p.add_argument( + "--is_last_checkpoint", + type=boolean, + nargs="?", + const=True, + default=os.environ.get("IS_LAST_CHECKPOINT", "0").lower() + in ("1", "true"), + help="Whether this checkpoint is the final checkpoint in the manifest.", + ) a = p.parse_args(argv) for name in ( "mesh_fsdp", @@ -268,7 +321,10 @@ def summarize(rows, instance_ids, attempts): total = len(instance_ids) * attempts solved = sum(row["resolved"] for row in rows) pass_at_k = {} - for k in sorted({1, attempts}): + ks = {1, attempts} + if attempts >= 4: + ks.add(4) + for k in sorted(ks): values = [] for group in grouped.values(): c = sum(row["resolved"] for row in group) @@ -421,8 +477,31 @@ def load_entries(a): return entries +def setup_rcp_logging(a): + """Points mllog at the RCP log; fails fast if mlperf_logging is missing.""" + if not a.rcp_logging: + return + from tunix.utils import mllog_utils + + # Without mlperf_logging every mllog call is a silent no-op, so fail before + # the eval runs instead of dropping its eval_* / run_stop events. + if mllog_utils.mllogger is None: + raise RuntimeError( + "--rcp_logging requires the mlperf_logging package, which is not" + " installed." + ) + if a.metric_logger_dir: + mllog_utils.configure_logger( + metric_logger_dir=a.metric_logger_dir, + seed=a.seed, + ) + + async def run_controller(a): from tunix.experimental.worker import remote_execution + from tunix.utils import mllog_utils + + setup_rcp_logging(a) entries = load_entries(a) run_id = ( @@ -458,6 +537,8 @@ async def run_controller(a): failure = None fleet = None entry_stream = entries + t_eval_start = None + eval_start_time_ms = None def record(row): writer.record(row) @@ -479,6 +560,8 @@ async def ready(handle): raise ValueError(f"Worker/controller model settings differ: {profile}") await asyncio.gather(*(ready(h) for h in handles)) + t_eval_start = time.monotonic() + eval_start_time_ms = time.time_ns() // 1_000_000 if a.use_agent_sandbox: from examples.deepswe import sandbox_utils # pylint: disable=import-outside-toplevel @@ -535,12 +618,54 @@ async def ready(handle): from examples.deepswe import sandbox_utils # pylint: disable=import-outside-toplevel await asyncio.to_thread(sandbox_utils.teardown_global_fleet) + validation_time = ( + time.monotonic() - t_eval_start if t_eval_start is not None else None + ) summary = summarize( all_rows, [str(e["instance_id"]) for e in entries], a.num_rollouts_per_instance, ) summary["fatal_error"] = failure + expected_attempts = int(summary["expected_attempts"]) + error_attempts = int(summary["error_attempts"]) + eval_ok = ( + failure is None + and bool(summary["complete"]) + and (expected_attempts > 0 and error_attempts < expected_attempts) + ) + pass_at_k = summary["pass_at_k"] + eval_accuracy = float( + pass_at_k["4"] + if "4" in pass_at_k + else pass_at_k[str(a.num_rollouts_per_instance)] + ) + target_acc = float(a.target_accuracy) + target_reached = bool(eval_ok and eval_accuracy >= target_acc) + rcp_logged = False + if a.rcp_logging and eval_ok: + mllog_utils.start_eval( + step=int(a.checkpoint_step), + samples_count=int(a.samples_count), + time_ms=eval_start_time_ms, + ) + target_reached = mllog_utils.log_offline_eval_step( + step=int(a.checkpoint_step), + samples_count=int(a.samples_count), + eval_accuracy=eval_accuracy, + target_accuracy=target_acc, + checkpoint_timestamp_ms=a.checkpoint_timestamp_ms, + is_last_checkpoint=bool(a.is_last_checkpoint), + validation_time=validation_time, + emit_start_eval=False, + ) + rcp_logged = True + summary["rcp_logged"] = rcp_logged + summary["target_accuracy"] = target_acc + summary["target_reached"] = bool(target_reached) + summary["checkpoint_step"] = int(a.checkpoint_step) + summary["checkpoint_timestamp_ms"] = a.checkpoint_timestamp_ms + summary["samples_count"] = int(a.samples_count) try: writer.write("summary.json", summary) finally: diff --git a/tunix/experimental/examples/deepswe_dist/k8s_launcher.sh b/tunix/experimental/examples/deepswe_dist/k8s_launcher.sh index f5678681df..0bed0386be 100755 --- a/tunix/experimental/examples/deepswe_dist/k8s_launcher.sh +++ b/tunix/experimental/examples/deepswe_dist/k8s_launcher.sh @@ -157,6 +157,11 @@ export WANDB_ENTITY=${WANDB_ENTITY:-} export LOG_DIR=${LOG_DIR:-} export TRAJECTORY_LOG_DIR=${TRAJECTORY_LOG_DIR:-} export RCP_LOGGING=${RCP_LOGGING:-false} +export VAL_START_AT=${VAL_START_AT:-} +export CHECKPOINT_STEP=${CHECKPOINT_STEP:-0} +export CHECKPOINT_TIMESTAMP_MS=${CHECKPOINT_TIMESTAMP_MS:-} +export SAMPLES_COUNT=${SAMPLES_COUNT:-0} +export IS_LAST_CHECKPOINT=${IS_LAST_CHECKPOINT:-false} export METRIC_LOGGER_DIR=${METRIC_LOGGER_DIR:-} export TARGET_ACCURACY=${TARGET_ACCURACY:-0.69} export TRAJECTORY_STORE_ROOT_DIR=${TRAJECTORY_STORE_ROOT_DIR:-${TRAJECTORY_STORE_ROOT:-}} @@ -407,6 +412,7 @@ start_orchestrator() { --tpu_topology="${TRAINER_TPU_SLICE}+${ROLLOUT_TPU_SLICE}" \ --target_accuracy=${TARGET_ACCURACY} \ ${METRIC_LOGGER_DIR:+--metric_logger_dir="${METRIC_LOGGER_DIR}"} \ + ${VAL_START_AT:+--val_start_at=${VAL_START_AT}} \ ${rcp_arg} \ ${debug_arg} \ " \ @@ -1003,6 +1009,7 @@ start_eval() { --tokenizer_path=${TOKENIZER_PATH} \ --model_absolute_path=${MAXTEXT_CKPT} \ --maxtext_model_name=${MAXTEXT_MODEL_NAME} \ + ${SCAN_LAYERS:+--scan_layers=${SCAN_LAYERS}} \ --mesh_fsdp=${ROLLOUT_MESH_FSDP:-2} \ --mesh_tp=${ROLLOUT_MESH_TP:-2} \ --vllm_utilization=${VLLM_GPU_MEMORY_UTILIZATION:-0.9} \ @@ -1034,6 +1041,13 @@ start_eval() { --use_agent_sandbox=${USE_AGENT_SANDBOX} \ --max_warmpool_size=${MAX_WARMPOOL_REPLICAS} \ --output_dir=${output_dir} \ + --rcp_logging=${RCP_LOGGING} \ + ${METRIC_LOGGER_DIR:+--metric_logger_dir=\"${METRIC_LOGGER_DIR}\"} \ + --target_accuracy=${TARGET_ACCURACY} \ + --checkpoint_step=${CHECKPOINT_STEP:-0} \ + ${CHECKPOINT_TIMESTAMP_MS:+--checkpoint_timestamp_ms=${CHECKPOINT_TIMESTAMP_MS}} \ + --samples_count=${SAMPLES_COUNT:-0} \ + --is_last_checkpoint=${IS_LAST_CHECKPOINT:-false} \ " \ | apply_manifest done diff --git a/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py b/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py index d9e3b2d074..56372307d4 100644 --- a/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py +++ b/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py @@ -376,6 +376,19 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: default=False, help="Enable MLPerf RCP (mllog) compliance logging.", ) + parser.add_argument( + "--val_start_at", + type=int, + default=( + int(os.getenv("VAL_START_AT")) + if os.getenv("VAL_START_AT", "").strip() + else None + ), + help=( + "First optimizer step to save/evaluate checkpoints from " + "(defaults to CEIL(2.5 + 3840 / global_batch_size))." + ), + ) parser.add_argument( "--metric_logger_dir", type=str, @@ -730,6 +743,58 @@ def main(argv: list[str], context: ProcessContext | None = None) -> None: max_staleness=args.max_staleness, ) + global_batch_size = int(args.batch_size) * int(args.num_generations) + val_start_step = ( + mllog_utils.compute_val_start_step(global_batch_size, args.val_start_at) + if args.rcp_logging or args.val_start_at is not None + else None + ) + manifest_file = ( + os.path.join( + args.metric_logger_dir.rstrip("/"), "eval_checkpoints.jsonl" + ) + if args.rcp_logging and args.metric_logger_dir + else "" + ) + + def _on_checkpoint_saved(ckpt_info: dict[str, Any]) -> None: + if not manifest_file: + return + step_num = int(ckpt_info["step"]) + samples_count = step_num * global_batch_size + ts_ms = int(ckpt_info["timestamp_ms"]) + ckpt_path = str(ckpt_info["checkpoint_path"]) + if not ckpt_path: + raise ValueError( + f"save_checkpoint for step {step_num} returned no checkpoint_path;" + " cannot write the eval manifest entry." + ) + record = { + "step": step_num, + "checkpoint_path": ckpt_path, + "timestamp_ms": ts_ms, + "samples_count": samples_count, + "global_batch_size": global_batch_size, + "batch_size": int(args.batch_size), + "num_generations": int(args.num_generations), + "val_start_at": int(val_start_step or 1), + "max_steps": int(args.max_steps), + "target_accuracy": float(args.target_accuracy), + "seed": int(args.seed), + "mllog_file": ( + mllog_utils.get_mllog_file_path( + metric_logger_dir=args.metric_logger_dir, seed=args.seed + ) + or "" + ), + } + mllog_utils.append_checkpoint_manifest(manifest_file, record) + logging.info( + "Appended checkpoint manifest entry for step=%d to %s", + step_num, + manifest_file, + ) + program = rl_program.StandardRLProgram( algo=algo, dataset=prompt_stream, @@ -778,6 +843,8 @@ def main(argv: list[str], context: ProcessContext | None = None) -> None: if args.rcp_logging else None, ), + val_start_step=val_start_step, + on_checkpoint_saved=_on_checkpoint_saved if manifest_file else None, ) if args.rcp_logging: @@ -802,7 +869,12 @@ def main(argv: list[str], context: ProcessContext | None = None) -> None: if program.last_step_result is not None else args.max_steps ) - mllog_utils.train_stop(args, step=completed_steps, status="success") + mllog_utils.train_stop( + args, + step=completed_steps, + status="success", + time_ms=program.last_step_timestamp_ms, + ) except BaseException as e: if args.rcp_logging: completed_steps = ( diff --git a/tunix/experimental/examples/recipes/README.md b/tunix/experimental/examples/recipes/README.md index 8ce81d295f..3ae534f39a 100644 --- a/tunix/experimental/examples/recipes/README.md +++ b/tunix/experimental/examples/recipes/README.md @@ -12,6 +12,7 @@ This directory contains executable recipe scripts for running distributed DeepSW | [`mlperf_397b_256_v7x.sh`](mlperf_397b_256_v7x.sh) | Qwen3.5-397B-A17B | `1x tpu7x:4x4x8` (128 chips)
`FSDP=32, TP=1, EP=2, CP=4` | `32x tpu7x:2x2x2` (256 chips)
`DP=1, TP=1, EP=16` | | [`mlperf_397b_512_v7x.sh`](mlperf_397b_512_v7x.sh) | Qwen3.5-397B-A17B | `1x tpu7x:4x4x8` (128 chips)
`FSDP=32, TP=1, EP=2, CP=4` | `32x tpu7x:2x2x4` (512 chips / 1024)
`DP=2, TP=1, EP=16` | | [`mlperf_397b_1024_v7x.sh`](mlperf_397b_1024_v7x.sh) | Qwen3.5-397B-A17B | `1x tpu7x:4x4x8` (128 chips)
`FSDP=32, TP=1, EP=2, CP=4` | `64x tpu7x:2x2x4` (1024 chips)
`DP=2, TP=1, EP=16` | +| [`mlperf_35b_eval.sh`](mlperf_35b_eval.sh) | Qwen3.5-35B-A3B (offline eval, pass@4) | None (no trainer) | `16x tpuv5:2x2x1` (64 chips)
`DP=2, FSDP=2, TP=2` | --- diff --git a/tunix/experimental/examples/recipes/mlperf_35b_128_v5p.sh b/tunix/experimental/examples/recipes/mlperf_35b_128_v5p.sh index c545415d07..705e43ead7 100755 --- a/tunix/experimental/examples/recipes/mlperf_35b_128_v5p.sh +++ b/tunix/experimental/examples/recipes/mlperf_35b_128_v5p.sh @@ -34,6 +34,14 @@ export ROLLOUT_TPU_SLICE="tpuv5:2x2x1" export ROLLOUT_MESH_EXPERT="${ROLLOUT_MESH_EXPERT:-4}" export ROLLOUT_REPLICAS="${ROLLOUT_REPLICAS:-16}" +# MLPerf RCP logging with deferred offline eval: from VAL_START_AT (default +# CEIL(2.5 + 3840 / global_batch_size) = 18) save a checkpoint every step and +# record it in ${METRIC_LOGGER_DIR}/eval_checkpoints.jsonl for +# mlperf_35b_eval.sh. Keep all of them (max_steps - VAL_START_AT + 1 <= 35). +export RCP_LOGGING="${RCP_LOGGING:-true}" +export CHECKPOINT_SAVE_INTERVAL_STEPS="${CHECKPOINT_SAVE_INTERVAL_STEPS:-1}" +export CHECKPOINT_MAX_TO_KEEP="${CHECKPOINT_MAX_TO_KEEP:-35}" + # vLLM Rollout Configuration (from paste.googleplex.com/5903655694368768) export VLLM_ADDITIONAL_CONFIG='{"sharding":{"sharding_strategy":{"expert_parallelism":4,"tensor_parallelism":1,"enable_dp_attention":true}},"custom_mamba_cache_multiplier":16,"maxtext_config":{"scan_layers":false,"attention":"vllm_rpa","allow_split_physical_axes":true,"use_multimodal":false,"prefuse_moe_weights":true}}' diff --git a/tunix/experimental/examples/recipes/mlperf_35b_eval.sh b/tunix/experimental/examples/recipes/mlperf_35b_eval.sh new file mode 100755 index 0000000000..9d018e0937 --- /dev/null +++ b/tunix/experimental/examples/recipes/mlperf_35b_eval.sh @@ -0,0 +1,78 @@ +#!/bin/bash +set -e + +DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +# Offline pass@4 evaluation of Qwen3.5-35B-A3B on the MLPerf validation split +# (251 instances x 4 rollouts) with 16x 4-chip v5p rollout slices (no trainer). +# +# MLPerf RCP: point CHECKPOINT_MANIFEST_FILE at a training run's +# ${METRIC_LOGGER_DIR}/eval_checkpoints.jsonl (written by mlperf_35b_128_v5p.sh). +# Checkpoints are evaluated in step order and eval_* events are appended to the +# training MLLOG until pass@4 >= TARGET_ACCURACY; run_stop is backdated to that +# checkpoint's weight-update timestamp (status=aborted if none converges). +# +# Without a manifest, MAXTEXT_CKPT is evaluated once. Set RCP_LOGGING=true to +# test RCP logging before post-training; CHECKPOINT_STEP, SAMPLES_COUNT and +# CHECKPOINT_TIMESTAMP_MS then default to mock values and events go to +# ${EVAL_OUTPUT_DIR}/mllog (not the training MLLOG) unless METRIC_LOGGER_DIR is set. + +# k8s has a 63 char limit on total label name, so keep job_prefix unique to your job and short +export JOB_PREFIX="${JOB_PREFIX:-${USER}}" +export MAXTEXT_OUTPUT_DIR="${MAXTEXT_OUTPUT_DIR:-gs://atwigg-trellis-europe-west4-dev/maxtext/${JOB_PREFIX}}" +export EVAL_OUTPUT_DIR="${EVAL_OUTPUT_DIR:-gs://atwigg-trellis-europe-west4-dev/eval_results/${JOB_PREFIX}}" +export TUNIX_IMAGE="${TUNIX_IMAGE:-gcr.io/cloud-tpu-multipod-dev/sanbao/tunix_stack:eval}" + +export REGION="europe-west4" +export CLUSTER="bodaborg-v5p-nap" +export K8S_NAMESPACE="trellis" + +export PATHWAYS_SERVER_IMAGE="${PATHWAYS_SERVER_IMAGE:-us-docker.pkg.dev/cloud-tpu-v2-images-dev/pathways/gke/datenglin/unsanitized_server:raiden_20260920_v2}" +export PATHWAYS_PROXY_IMAGE="${PATHWAYS_PROXY_IMAGE:-us-docker.pkg.dev/cloud-tpu-v2-images-dev/pathways/gke/datenglin/unsanitized_proxy_server:raiden_20260920_v2}" + +# Model configuration +export MODEL_NAME="Qwen3.5-35B-A3B" +export MODEL_ID="Qwen/Qwen3.5-35B-A3B" +export TOKENIZER_PATH="Qwen/Qwen3.5-35B-A3B" +export MAXTEXT_MODEL_NAME="qwen3.5-35b-a3b" +export MAXTEXT_CKPT="${MAXTEXT_CKPT:-gs://maxtext-model-checkpoints/qwen3.5-35b-a3b/scanned/0/items}" +export SCAN_LAYERS="${SCAN_LAYERS:-true}" +export CHECKPOINT_STORAGE_USE_OCDBT="${CHECKPOINT_STORAGE_USE_OCDBT:-false}" +export CHECKPOINT_STORAGE_USE_ZARR3="${CHECKPOINT_STORAGE_USE_ZARR3:-false}" + +# Rollout topology (Pathways 4-chip 2x2x1 slices, DP=2, FSDP=2, TP=2) +export WEIGHT_SYNC_MODE="none" +export ROLLOUT_JOBSET_YAML="jobset.pathways.yaml" +export ROLLOUT_TPU_SLICE="tpuv5:2x2x1" +export ROLLOUT_MESH_FSDP=2 +export ROLLOUT_MESH_TP=2 +export ROLLOUT_REPLICAS="${ROLLOUT_REPLICAS:-16}" +export VLLM_DATA_PARALLEL_SIZE=2 +export ENABLE_PREFIX_CACHING="${ENABLE_PREFIX_CACHING:-false}" + +# Validation: 4 rollouts per instance, temperature 0.1, top_p 0.95 +export DATASET_PATH="${DATASET_PATH:-gs://mlperf_dataset/benchmark-r2e-gym-easy}" +export DATASET_SPLIT="${DATASET_SPLIT:-validation}" +export NUM_GENERATIONS="${NUM_GENERATIONS:-4}" +export TEMPERATURE="0.1" +export TOP_P="0.95" +export STEP_TIMEOUT_SECS=60 +export REWARD_TIMEOUT_SECS=60 + +# Sandbox +export SANDBOX_NODE_SELECTOR_VAL="sandbox-cpu-pool" +export IMAGE_REWRITE_PREFIX="${IMAGE_REWRITE_PREFIX:-europe-west4-docker.pkg.dev/cloud-tpu-multipod-dev/tunix/}" + +# MLPerf RCP offline eval (see top of file) +export CHECKPOINT_MANIFEST_FILE="${CHECKPOINT_MANIFEST_FILE:-}" +if [[ -n "${CHECKPOINT_MANIFEST_FILE}" ]]; then + export RCP_LOGGING="${RCP_LOGGING:-true}" +else + export METRIC_LOGGER_DIR="${METRIC_LOGGER_DIR:-${EVAL_OUTPUT_DIR%/}/mllog}" +fi +export CHECKPOINT_STEP="${CHECKPOINT_STEP:-18}" +export SAMPLES_COUNT="${SAMPLES_COUNT:-4608}" +export CHECKPOINT_TIMESTAMP_MS="${CHECKPOINT_TIMESTAMP_MS:-$(date +%s)000}" +export IS_LAST_CHECKPOINT="${IS_LAST_CHECKPOINT:-true}" + +source "${DIR}/mlperf_base.sh" "${1:-eval}" "${@:2}" diff --git a/tunix/experimental/examples/recipes/mlperf_base.sh b/tunix/experimental/examples/recipes/mlperf_base.sh index 6fce6594bb..477b0729de 100755 --- a/tunix/experimental/examples/recipes/mlperf_base.sh +++ b/tunix/experimental/examples/recipes/mlperf_base.sh @@ -57,7 +57,7 @@ export TRAINER_PREFUSE_MOE_WEIGHTS="true" export ROLLOUT_PREFUSE_MOE_WEIGHTS="true" export VERIFY_WEIGHTS="true" export TRAINER_PADDED_MOE_MLP_DIM="" -export WEIGHT_SYNC_MODE="raiden" +export WEIGHT_SYNC_MODE="${WEIGHT_SYNC_MODE:-raiden}" export WEIGHT_SYNC_DISABLE_TIMEOUTS="${WEIGHT_SYNC_DISABLE_TIMEOUTS:-${DISABLE_WEIGHT_SYNC_TIMEOUTS:-0}}" export TPU_RAIDEN_DATA_NICS="${TPU_RAIDEN_DATA_NICS:-eth0}" @@ -81,8 +81,8 @@ export TRAINABLE_PARAMETERS_MASK='^(?!.*routed_experts/gate/kernel).*' # Qwen3.5 vocab. export EOS_TOKENS="${EOS_TOKENS:-248046,248044}" export TRAINER_BASE_NUM_KV_HEADS=2 -export ROLLOUT_MESH_FSDP=1 -export ROLLOUT_MESH_TP=1 +export ROLLOUT_MESH_FSDP="${ROLLOUT_MESH_FSDP:-1}" +export ROLLOUT_MESH_TP="${ROLLOUT_MESH_TP:-1}" # ============================================================================== # MLPerf RCP Logging @@ -153,10 +153,10 @@ export VLLM_ENABLE_V1_MULTIPROCESSING=0 export MAX_STEPS=${MAX_STEPS:-50} export BATCH_SIZE=${BATCH_SIZE:-16} export MINI_BATCH_SIZE=${MINI_BATCH_SIZE:-${BATCH_SIZE}} -export NUM_GENERATIONS=16 +export NUM_GENERATIONS="${NUM_GENERATIONS:-16}" export TRAIN_MICRO_BATCH_SIZE="${TRAIN_MICRO_BATCH_SIZE:-32}" export CHECKPOINT_SAVE_INTERVAL_STEPS=${CHECKPOINT_SAVE_INTERVAL_STEPS:-0} -export CHECKPOINT_MAX_TO_KEEP=10 +export CHECKPOINT_MAX_TO_KEEP="${CHECKPOINT_MAX_TO_KEEP:-10}" export CHECKPOINT_ASYNC=${CHECKPOINT_ASYNC:-true} export ENABLE_PATHWAYS_PERSISTENCE=${ENABLE_PATHWAYS_PERSISTENCE:-1} export MAX_STALENESS=${MAX_STALENESS:-1} @@ -167,8 +167,8 @@ export MAX_SEQ_TOKEN_PER_TPU=${MAX_SEQ_TOKEN_PER_TPU:-65536} export MAX_SEGMENTS_PER_PACKED_ROW=${MAX_SEGMENTS_PER_PACKED_ROW:-16} # Sampling Parameters (explicitly disable top-k, set top-p 1.0 and temperature 1.0) -export TEMPERATURE="1.0" -export TOP_P="1.0" +export TEMPERATURE="${TEMPERATURE:-1.0}" +export TOP_P="${TOP_P:-1.0}" export TOP_K="-1" # Algorithmic & Loss Hyperparameters @@ -231,8 +231,8 @@ export SANDBOX_NODE_SELECTOR_VAL="${SANDBOX_NODE_SELECTOR_VAL:-sandbox-np}" export MAX_WARMPOOL_REPLICAS=2 export ROLLOUT_MAX_CONCURRENCY="${ROLLOUT_MAX_CONCURRENCY:-256}" export MAX_CONCURRENCY="${MAX_CONCURRENCY:-256}" -export STEP_TIMEOUT_SECS=300 -export REWARD_TIMEOUT_SECS=180 +export STEP_TIMEOUT_SECS="${STEP_TIMEOUT_SECS:-300}" +export REWARD_TIMEOUT_SECS="${REWARD_TIMEOUT_SECS:-180}" export FLUSH_EVERY_N_STEPS=1 export MAX_TURNS=30 export MAX_PROMPT_LENGTH="${MAX_PROMPT_LENGTH:-4096}" @@ -259,5 +259,108 @@ fi if [[ "${MLPERF_NO_LAUNCH:-0}" != "1" ]]; then COMMAND="${1:-start}" shift || true + if [[ "${COMMAND}" == "eval" && -n "${CHECKPOINT_MANIFEST_FILE:-}" ]]; then + echo "Running sequential offline evaluation from manifest: ${CHECKPOINT_MANIFEST_FILE}" + # Stdlib-only (the launcher host has no JAX/tunix install). Prints one + # "step, samples_count, timestamp_ms, checkpoint_path, is_last, mllog_file" + # TSV row per checkpoint; a missing, empty or non-contiguous manifest aborts + # (set -e). + MANIFEST_ROWS_TSV="$(python3 -c ' +import json, subprocess, sys +path = sys.argv[1] +if path.startswith("gs://"): + text = subprocess.check_output(["gsutil", "cat", path], text=True) +else: + with open(path, encoding="utf-8") as f: + text = f.read() +records = sorted( + (json.loads(line) for line in text.splitlines() if line.strip()), + key=lambda r: int(r["step"]), +) +if not records: + sys.exit(f"Checkpoint manifest is empty: {path}") +steps = [int(r["step"]) for r in records] +first = int(records[0].get("val_start_at", steps[0])) +if steps != list(range(first, first + len(steps))): + sys.exit(f"Manifest steps must be contiguous from val_start_at={first}: {steps}") +for i, r in enumerate(records): + print("\t".join([ + str(int(r["step"])), + str(int(r["samples_count"])), + str(int(r["timestamp_ms"])), + str(r["checkpoint_path"]), + "true" if i == len(records) - 1 else "false", + str(r.get("mllog_file") or ""), + ])) +' "${CHECKPOINT_MANIFEST_FILE}")" + mapfile -t MANIFEST_ROWS <<< "${MANIFEST_ROWS_TSV}" + BASE_EVAL_OUTPUT_DIR="${EVAL_OUTPUT_DIR:-${MAXTEXT_OUTPUT_DIR}/eval_results}" + EVAL_JOBSET_NAME="${EVAL_JOBSET_NAME:-${JOB_PREFIX}-eval}" + for row in "${MANIFEST_ROWS[@]}"; do + IFS=$'\t' read -r STEP SAMPLES TS_MS CKPT_PATH IS_LAST MLLOG_FILE <<< "${row}" + + echo "=== Evaluating checkpoint step=${STEP} samples=${SAMPLES} is_last=${IS_LAST} path=${CKPT_PATH} ===" + export MAXTEXT_CKPT="${CKPT_PATH}" + export CHECKPOINT_STEP="${STEP}" + export SAMPLES_COUNT="${SAMPLES}" + export CHECKPOINT_TIMESTAMP_MS="${TS_MS}" + export IS_LAST_CHECKPOINT="${IS_LAST}" + # Append eval_* / run_stop to the training run's MLLOG. + export METRIC_LOGGER_DIR="${MLLOG_FILE:-${METRIC_LOGGER_DIR:-}}" + export EVAL_OUTPUT_DIR="${BASE_EVAL_OUTPUT_DIR%/}/step_${STEP}" + "${LAUNCHER}" --command eval --image "${TUNIX_IMAGE}" "$@" + + if [[ "${DRY_RUN:-false}" != "true" ]]; then + HEAD_JOBSET="${EVAL_JOBSET_NAME}" + if [[ "${ROLLOUT_REPLICAS:-1}" -gt 1 ]]; then + HEAD_JOBSET="${EVAL_JOBSET_NAME}-0" + fi + echo "Waiting for evaluation JobSet ${HEAD_JOBSET} (main container) in namespace ${K8S_NAMESPACE}..." + while true; do + if ! kubectl get jobset "${HEAD_JOBSET}" -n "${K8S_NAMESPACE}" &>/dev/null; then + echo "JobSet ${HEAD_JOBSET} no longer exists." + break + fi + MAIN_EXIT="$(kubectl get pods -n "${K8S_NAMESPACE}" -l "jobset.sigs.k8s.io/jobset-name=${HEAD_JOBSET},jobset.sigs.k8s.io/replicatedjob-name=proc" -o jsonpath='{.items[0].status.containerStatuses[?(@.name=="main")].state.terminated.exitCode}' 2>/dev/null || true)" + if [[ -n "${MAIN_EXIT}" ]]; then + echo "Main evaluation container finished with exit code ${MAIN_EXIT}." + break + fi + sleep 10 + done + "${LAUNCHER}" --command stop_eval --image "${TUNIX_IMAGE}" || true + + TARGET_REACHED="$( + python3 -c ' +import glob, json, os, subprocess, sys +out_dir = os.environ["EVAL_OUTPUT_DIR"].rstrip("/") +if out_dir.startswith("gs://"): + res = subprocess.run(["gsutil", "ls", f"{out_dir}/*/summary.json"], capture_output=True, text=True, check=False) + if res.returncode == 0 and res.stdout.strip(): + matches = sorted(line.strip() for line in res.stdout.splitlines() if line.strip()) + if matches: + res_cat = subprocess.run(["gsutil", "cat", matches[-1]], capture_output=True, text=True, check=False) + if res_cat.returncode == 0 and res_cat.stdout.strip(): + data = json.loads(res_cat.stdout) + print("true" if data.get("target_reached") else "false") + sys.exit(0) +else: + matches = sorted(glob.glob(f"{out_dir}/*/summary.json")) + if matches: + with open(matches[-1], "r", encoding="utf-8") as f: + data = json.load(f) + print("true" if data.get("target_reached") else "false") + sys.exit(0) +print("false") +' + )" + if [[ "${TARGET_REACHED}" == "true" ]]; then + echo "Target accuracy ${TARGET_ACCURACY} reached at step ${STEP}. Stopping offline evaluation loop." + break + fi + fi + done + exit 0 + fi exec "${LAUNCHER}" --command "${COMMAND}" --image "${TUNIX_IMAGE}" "$@" fi diff --git a/tunix/experimental/orchestrator/rl_program.py b/tunix/experimental/orchestrator/rl_program.py index b458777929..6bfcab0ef8 100644 --- a/tunix/experimental/orchestrator/rl_program.py +++ b/tunix/experimental/orchestrator/rl_program.py @@ -365,6 +365,8 @@ def __init__( ) = trajectory_queue_manager.GroupOrder.ARRIVAL, on_step_begin: Callable[[int], None] | None = None, on_step_end: Callable[[int, Any], None] | None = None, + val_start_step: int | None = None, + on_checkpoint_saved: Callable[[dict[str, Any]], None] | None = None, ): super().__init__() self.engine: rl_engine_interface.AbstractRLEngine | None = None @@ -499,6 +501,9 @@ def __init__( ) self.on_step_begin = on_step_begin self.on_step_end = on_step_end + self.val_start_step = val_start_step + self.on_checkpoint_saved = on_checkpoint_saved + self.last_step_timestamp_ms: int | None = None self._in_flight_rollouts = 0 self._window_release = asyncio.Event() self._dispatch_done = asyncio.Event() @@ -1449,12 +1454,20 @@ async def train_stage(self) -> None: async def _maybe_save_checkpoint() -> None: nonlocal checkpoint_saved + ckpt_ts_ms = time.time_ns() // 1_000_000 + self.last_step_timestamp_ms = ckpt_ts_ms optimizer_step = self.step + 1 if ( isinstance(step_result, dict) and step_result.get("train_step") is not None ): optimizer_step = int(step_result["train_step"]) + if ( + self.val_start_step is not None + and optimizer_step < self.val_start_step + ): + checkpoint_saved = True + return if isinstance( self.scored_q, trajectory_queue_manager.BatchOrderedQueueManager, @@ -1466,7 +1479,7 @@ async def _maybe_save_checkpoint() -> None: ) else: next_batch_idx = self.step + 1 - await self.engine.save_checkpoint( + save_resp = await self.engine.save_checkpoint( role=datatypes.Role.ACTOR, metadata={ "step": optimizer_step, @@ -1478,6 +1491,17 @@ async def _maybe_save_checkpoint() -> None: }, ) checkpoint_saved = True + if self.on_checkpoint_saved is not None: + # `TrainerWorker.save_checkpoint` always returns the saved path in + # `Response.metadata["checkpoint_path"]`. The callback runs in a + # worker thread (it may do blocking file/GCS I/O) and is awaited, so + # calls never overlap and its exceptions still propagate. + ckpt_info: dict[str, Any] = { + "step": optimizer_step, + "timestamp_ms": ckpt_ts_ms, + "checkpoint_path": save_resp.metadata["checkpoint_path"], + } + await asyncio.to_thread(self.on_checkpoint_saved, ckpt_info) while groups_consumed < self.full_batch_size: _t_gen = time.monotonic() diff --git a/tunix/experimental/train/abstract_trainer.py b/tunix/experimental/train/abstract_trainer.py index 3034e34303..dddd8a821c 100644 --- a/tunix/experimental/train/abstract_trainer.py +++ b/tunix/experimental/train/abstract_trainer.py @@ -164,6 +164,11 @@ def model_scope( f"{type(self).__name__} does not implement model_scope." ) + @property + def checkpoint_dir(self) -> str | None: + """Returns the root directory where checkpoints are saved, if configured.""" + return None + @abc.abstractmethod def save_checkpoint(self, metadata: Any, **kwargs) -> None: """Force the trainer to serialize its state (model + optimizer). diff --git a/tunix/experimental/train/peft_trainer_v2.py b/tunix/experimental/train/peft_trainer_v2.py index a6e4beafe8..21f31f1ded 100644 --- a/tunix/experimental/train/peft_trainer_v2.py +++ b/tunix/experimental/train/peft_trainer_v2.py @@ -1287,6 +1287,11 @@ def _shard(x: Any) -> Any: # answer to "which weights are live", and a subclass may rebind it. yield self.model, args, kwargs + @property + @override + def checkpoint_dir(self) -> str | None: + return self.config.checkpoint_root_directory + @override def save_checkpoint(self, metadata: Any = None, **kwargs) -> None: """Saves a checkpoint of the trainer state (model + optimizer). diff --git a/tunix/experimental/worker/trainer_worker.py b/tunix/experimental/worker/trainer_worker.py index 81ca6f391a..de24f8cf3e 100644 --- a/tunix/experimental/worker/trainer_worker.py +++ b/tunix/experimental/worker/trainer_worker.py @@ -14,6 +14,7 @@ """TrainerWorker implementation for role-based isolation.""" +from collections.abc import Mapping import contextlib from typing import Any, Callable, ContextManager, cast @@ -357,13 +358,22 @@ def per_token_logps( self.state = WorkerState.ERROR raise + def _resolve_checkpoint_path(self, metadata: Mapping[str, Any]) -> str: + """Resolves the saved Orbax model_params directory from the underlying trainer.""" + ckpt_dir = self._trainer.checkpoint_dir + step = metadata.get("step") + if ckpt_dir and step is not None: + return f"{str(ckpt_dir).rstrip('/')}/{int(step)}/model_params" + return "" + def save_checkpoint(self, metadata: Any, **kwargs) -> datatypes.Response: """Force the trainer to serialize its state (model + optimizer).""" self._ensure_ready() try: self._trainer.save_checkpoint(metadata, **kwargs) + ckpt_path = self._resolve_checkpoint_path(metadata) self._last_error = None - return self._response(checkpoint_saved=True) + return self._response(checkpoint_saved=True, checkpoint_path=ckpt_path) except Exception as exc: self._last_error = str(exc) self.state = WorkerState.ERROR diff --git a/tunix/utils/maxtext_utils.py b/tunix/utils/maxtext_utils.py index 46a551f79c..11fe3c11cb 100644 --- a/tunix/utils/maxtext_utils.py +++ b/tunix/utils/maxtext_utils.py @@ -806,6 +806,7 @@ def _init_state(self) -> None: wrap_with_tunix_adapter=wrap_with_tunix_adapter, tokenizer_pad_id=tokenizer_pad_id, ) + engine.checkpoint_dir = maxtext_config.checkpoint_dir # When `float32_gate_logits=True` and `weight_dtype=bfloat16`, both the # trainer and vLLM rollout models store gate/router/norm/GDN/logits_dense diff --git a/tunix/utils/mllog_utils.py b/tunix/utils/mllog_utils.py index 2562fbda61..7e2dbc4cb5 100644 --- a/tunix/utils/mllog_utils.py +++ b/tunix/utils/mllog_utils.py @@ -14,7 +14,9 @@ """MLPerf RCP Logging Utilities for Post-Training (GRPO).""" +import json import logging +import math import os from typing import Any, Callable, Iterable, Mapping, Optional import jax @@ -34,6 +36,34 @@ _local_log_path: Optional[str] = None +def _download_from_gcs_if_exists(gcs_path: str, local_path: str) -> None: + """Downloads an existing GCS mllog file so new events append to it.""" + try: + import fsspec # pylint: disable=g-import-not-at-top + + fs = fsspec.filesystem("gs") + if fs.exists(gcs_path): + fs.get(gcs_path, local_path) + except (ImportError, ModuleNotFoundError): + import tensorflow as tf # pylint: disable=g-import-not-at-top + + if tf.io.gfile.exists(gcs_path): + tf.io.gfile.copy(gcs_path, local_path, overwrite=True) + + +def get_mllog_file_path( + metric_logger_dir: Optional[str] = None, + seed: Optional[int] = None, +) -> Optional[str]: + """Returns the MLLOG file (GCS URI or local path) for metric_logger_dir.""" + if not metric_logger_dir: + return _gcs_target_path or _local_log_path + if metric_logger_dir.endswith(".out") or metric_logger_dir.endswith(".log"): + return metric_logger_dir + seed_val = seed if seed is not None else 1 + return os.path.join(metric_logger_dir.rstrip("/"), f"seed_{seed_val}.out") + + def _parse_topology_devices( topology: Optional[str], rollout_replicas: int = 1 ) -> Optional[int]: @@ -71,7 +101,7 @@ def _flush_to_gcs_if_needed() -> None: fs = fsspec.filesystem("gs") fs.put(_local_log_path, _gcs_target_path) - except Exception: # pylint: disable=broad-exception-caught + except (ImportError, ModuleNotFoundError): import tensorflow as tf # pylint: disable=g-import-not-at-top tf.io.gfile.makedirs(os.path.dirname(_gcs_target_path)) @@ -120,6 +150,8 @@ def configure_logger( abs_filename = os.path.abspath(filename) _local_log_path = abs_filename os.makedirs(os.path.dirname(abs_filename), exist_ok=True) + if _gcs_target_path and not os.path.exists(abs_filename): + _download_from_gcs_if_exists(_gcs_target_path, abs_filename) existing_files = [ os.path.abspath(getattr(h, "baseFilename", "")) for h in getattr(mllogger.logger, "handlers", []) @@ -227,16 +259,22 @@ def train_start(args=None, step: int = 0, samples_count: Optional[int] = None): _flush_to_gcs_if_needed() -def block_stop(step: int = 0, samples_count: Optional[int] = None): +def block_stop( + step: int = 0, + samples_count: Optional[int] = None, + time_ms: Optional[int] = None, +): """Marks the end of a training block.""" if _is_master_process() and mllogger is not None: metadata = {"step": int(step)} if samples_count is not None: metadata[getattr(constants, "SAMPLES_COUNT", "samples_count")] = int(samples_count) + extra_kwargs = {} if time_ms is None else {"time_ms": int(time_ms)} mllogger.end( key=getattr(constants, "BLOCK_STOP", "block_stop"), metadata=metadata, + **extra_kwargs, ) @@ -245,8 +283,22 @@ def train_stop( step: Optional[int] = None, samples_count: Optional[int] = None, status: str = "success", + time_ms: Optional[int] = None, ): - """Marks the end of a training block and the training run.""" + """Marks the end of the last training block. + + run_stop is not emitted here: the offline evaluator emits it (backdated to + the passing checkpoint's weight-update timestamp) via log_offline_eval_step. + + Args: + args: Optional namespace providing max_steps, batch_size, num_generations. + step: Last completed optimizer step (defaults to args.max_steps). + samples_count: Cumulative training samples (defaults to step * gbs). + status: Unused; kept for backward compatibility. + time_ms: Optional timestamp (ms) to backdate block_stop to, e.g. the last + weight update. + """ + del status if args is not None: if step is None: step = getattr(args, "max_steps", 0) @@ -257,20 +309,26 @@ def train_stop( samples_count = int(step) * global_batch_size step_val = 0 if step is None else int(step) - block_stop(step=step_val, samples_count=samples_count) - run_stop(status=status, samples_count=samples_count) + block_stop(step=step_val, samples_count=samples_count, time_ms=time_ms) + _flush_to_gcs_if_needed() -def start_eval(step: int = 0, samples_count: Optional[int] = None): +def start_eval( + step: int = 0, + samples_count: Optional[int] = None, + time_ms: Optional[int] = None, +): """Marks the start of an evaluation interval.""" if _is_master_process() and mllogger is not None: metadata = {"step": int(step)} if samples_count is not None: metadata[getattr(constants, "SAMPLES_COUNT", "samples_count")] = int(samples_count) + extra_kwargs = {} if time_ms is None else {"time_ms": int(time_ms)} mllogger.start( key=getattr(constants, "EVAL_START", "eval_start"), metadata=metadata, + **extra_kwargs, ) @@ -279,6 +337,7 @@ def end_eval( accuracy: float = 0.0, samples_count: Optional[int] = None, validation_time: Optional[float] = None, + time_ms: Optional[int] = None, ): """Marks the end of an evaluation interval and records eval accuracy.""" if _is_master_process() and mllogger is not None: @@ -286,11 +345,13 @@ def end_eval( if samples_count is not None: metadata[getattr(constants, "SAMPLES_COUNT", "samples_count")] = int(samples_count) + extra_kwargs = {} if time_ms is None else {"time_ms": int(time_ms)} if validation_time is not None: mllogger.event( key="tracked_stats", value={"validation_time": float(validation_time)}, metadata={"step": int(step)}, + **extra_kwargs, ) eval_accuracy_metadata = {} @@ -301,11 +362,147 @@ def end_eval( key=getattr(constants, "EVAL_ACCURACY", "eval_accuracy"), value=float(accuracy), metadata=eval_accuracy_metadata, + **extra_kwargs, ) mllogger.end( key=getattr(constants, "EVAL_STOP", "eval_stop"), metadata=metadata, + **extra_kwargs, + ) + + +def log_offline_eval_step( + step: int, + samples_count: int, + eval_accuracy: float, + target_accuracy: float = 0.69, + checkpoint_timestamp_ms: Optional[int] = None, + is_last_checkpoint: bool = False, + validation_time: Optional[float] = None, + emit_start_eval: bool = True, +) -> bool: + """Logs offline eval events for one checkpoint and a backdated run_stop. + + run_stop(status="success") is emitted when eval_accuracy reaches + target_accuracy; run_stop(status="aborted") is emitted when the final + checkpoint misses it. Both use checkpoint_timestamp_ms as time_ms so + checkpoint serialization and offline eval are excluded from time-to-train. + + Args: + step: Optimizer step of the evaluated checkpoint. + samples_count: Cumulative training samples at step. + eval_accuracy: pass@4 accuracy of the checkpoint. + target_accuracy: Convergence threshold. + checkpoint_timestamp_ms: Weight-update timestamp attached to the + checkpoint. + is_last_checkpoint: Whether this is the final checkpoint to evaluate. + validation_time: Optional eval wall time in seconds. + emit_start_eval: Whether to emit eval_start (set False if the caller + already emitted it when eval began). + + Returns: + True if target_accuracy was reached. + """ + passed = float(eval_accuracy) >= float(target_accuracy) + if not (_is_master_process() and mllogger is not None): + return passed + + if emit_start_eval: + start_eval(step=int(step), samples_count=int(samples_count)) + end_eval( + step=int(step), + accuracy=float(eval_accuracy), + samples_count=int(samples_count), + validation_time=validation_time, + ) + if passed or is_last_checkpoint: + # run_stop flushes to GCS. + run_stop( + status="success" if passed else "aborted", + samples_count=int(samples_count), + time_ms=checkpoint_timestamp_ms, ) + else: + _flush_to_gcs_if_needed() + return passed + + +def compute_val_start_step( + global_batch_size: int, + val_start_at_override: Optional[int] = None, +) -> int: + """Returns the first step to validate: CEIL(2.5 + 3840 / global_batch_size).""" + if val_start_at_override is not None and int(val_start_at_override) > 0: + return int(val_start_at_override) + if int(global_batch_size) <= 0: + raise ValueError( + f"global_batch_size must be positive, got {global_batch_size}" + ) + return int(math.ceil(2.5 + 3840.0 / float(global_batch_size))) + + +def _read_manifest_text(manifest_path: str) -> str: + """Returns the manifest content (local or gs://), or "" if it is missing.""" + if not manifest_path.startswith("gs://"): + if not os.path.exists(manifest_path): + return "" + with open(manifest_path, "r", encoding="utf-8") as f: + return f.read() + try: + import fsspec # pylint: disable=g-import-not-at-top + + fs = fsspec.filesystem("gs") + if not fs.exists(manifest_path): + return "" + with fs.open(manifest_path, "r", encoding="utf-8") as f: + return f.read() + except (ImportError, ModuleNotFoundError): + import tensorflow as tf # pylint: disable=g-import-not-at-top + + if not tf.io.gfile.exists(manifest_path): + return "" + with tf.io.gfile.GFile(manifest_path, "r") as f: + return f.read() + + +def _write_manifest_text(manifest_path: str, content: str) -> None: + """Writes the manifest content to a local path or gs:// URI.""" + if not manifest_path.startswith("gs://"): + os.makedirs(os.path.dirname(os.path.abspath(manifest_path)), exist_ok=True) + with open(manifest_path, "w", encoding="utf-8") as f: + f.write(content) + return + try: + import fsspec # pylint: disable=g-import-not-at-top + + fs = fsspec.filesystem("gs") + with fs.open(manifest_path, "w", encoding="utf-8") as f: + f.write(content) + except (ImportError, ModuleNotFoundError): + import tensorflow as tf # pylint: disable=g-import-not-at-top + + with tf.io.gfile.GFile(manifest_path, "w") as f: + f.write(content) + + +def append_checkpoint_manifest( + manifest_path: str, + record: Mapping[str, Any], +) -> None: + """Upserts a checkpoint record (keyed by step) into a JSONL manifest.""" + if not manifest_path: + return + records_by_step: dict[int, dict[str, Any]] = {} + for line in _read_manifest_text(manifest_path).splitlines(): + if line.strip(): + parsed = json.loads(line) + records_by_step[int(parsed["step"])] = parsed + records_by_step[int(record["step"])] = dict(record) + content = "".join( + json.dumps(records_by_step[s], ensure_ascii=False) + "\n" + for s in sorted(records_by_step) + ) + _write_manifest_text(manifest_path, content) def check_eval( @@ -822,16 +1019,22 @@ def rcp_logger(metrics_buffer: Any) -> None: rcp_metrics_logger = create_rcp_metrics_logger -def run_stop(status: str = "success", samples_count: Optional[int] = None): +def run_stop( + status: str = "success", + samples_count: Optional[int] = None, + time_ms: Optional[int] = None, +): """Marks the end of the training run.""" if _is_master_process() and mllogger is not None: metadata = {"status": status} if samples_count is not None: metadata[getattr(constants, "SAMPLES_COUNT", "samples_count")] = int(samples_count) + extra_kwargs = {} if time_ms is None else {"time_ms": int(time_ms)} mllogger.end( key=getattr(constants, "RUN_STOP", "run_stop"), metadata=metadata, + **extra_kwargs, ) _flush_to_gcs_if_needed() @@ -882,7 +1085,8 @@ def init_print( except (TypeError, AttributeError): pass if eval_samples is None: - eval_samples = 256 + # v6.1 qwen35_397b_grpo validation split size (compliance: == 251). + eval_samples = 251 # Parallelism dimensions from meshes train_tp = 1 @@ -961,6 +1165,11 @@ def init_print( getattr(constants, "PIPELINE_PARALLELISM", "pipeline_parallelism"): 1, getattr(constants, "CONTEXT_PARALLELISM", "context_parallelism"): train_sp, getattr(constants, "EXPERT_PARALLELISM", "expert_parallelism"): getattr(args, "train_mesh_expert", 1), + # Mandatory v6.1 precision and run-config disclosures. + "lowest_numerical_precision_in_linear": "bfloat16", + "lowest_numerical_precision_in_attn": "bfloat16", + "lowest_numerical_precision_in_comm": "bfloat16", + "config_filename": args.model_id or "qwen35_397b_grpo", "generation_backend": getattr(args, "rollout_engine", "vllm"), "generation_tensor_parallelism": rollout_tp, "generation_pipeline_parallelism": 1,