From 2b767d215141ef42329b9ee805316845645abec9 Mon Sep 17 00:00:00 2001 From: maxtext authors Date: Sun, 30 Aug 2026 11:43:44 -0700 Subject: [PATCH] Add ML Diagnostics scalar backend and options wiring for MaxText RL. Integrates Google Cloud ML Diagnostics (google-cloud-mldiagnostics) with MaxText/Tunix RL training loops by providing a custom MLDiagScalarBackend implementing the Tunix LoggingBackend protocol. PiperOrigin-RevId: 973504402 --- .../trainers/post_train/rl/mldiag_backend.py | 178 ++++++++++ .../trainers/post_train/rl/train_rl.py | 39 ++- tests/unit/mldiag_backend_test.py | 319 ++++++++++++++++++ 3 files changed, 530 insertions(+), 6 deletions(-) create mode 100644 src/maxtext/trainers/post_train/rl/mldiag_backend.py create mode 100644 tests/unit/mldiag_backend_test.py diff --git a/src/maxtext/trainers/post_train/rl/mldiag_backend.py b/src/maxtext/trainers/post_train/rl/mldiag_backend.py new file mode 100644 index 0000000000..b5d88d90b2 --- /dev/null +++ b/src/maxtext/trainers/post_train/rl/mldiag_backend.py @@ -0,0 +1,178 @@ +# Copyright 2023–2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""ML Diagnostics Scalar Backend for Tunix RL. + +Implements the metrax / Tunix LoggingBackend protocol (log_scalar, close) +to stream training and evaluation scalar metrics directly into Google Cloud ML +Diagnostics. +""" + +from __future__ import annotations + +import logging +import math +from typing import Any + +import jax +import numpy as np + +try: + import google_cloud_mldiagnostics as mldiag + + mldiag_metrics = getattr(mldiag, "metrics", None) + metric_types = getattr(mldiag, "metric_types", None) + _HAS_MLDIAG = mldiag is not None and mldiag_metrics is not None +except ImportError: + try: + # pylint: disable=g-import-not-at-top + from maxtext.common.gcloud_stub import mldiagnostics_modules + + mldiag, _ = mldiagnostics_modules() + mldiag_metrics = getattr(mldiag, "metrics", None) if mldiag else None + metric_types = getattr(mldiag, "metric_types", None) if mldiag else None + _HAS_MLDIAG = mldiag is not None and mldiag_metrics is not None + except Exception: # pylint: disable=broad-exception-caught + mldiag = None + mldiag_metrics = None + metric_types = None + _HAS_MLDIAG = False + +from maxtext.common import managed_mldiagnostics + +_EXACT_METRIC_MAP = { + "loss": "LOSS", + "learning_rate": "LEARNING_RATE", + "grad_norm": "GRADIENT_NORM", + "total_weights": "TOTAL_WEIGHTS", + "step_time": "STEP_TIME", + "global_step_time": "STEP_TIME", + "throughput": "THROUGHPUT", + "latency": "LATENCY", +} + + +def _normalize_metric_event(event: str) -> str: + """Normalizes enum string representations in metric event names. + + In Python 3.11+, enum interpolation in f-strings can format enums as + 'Mode.TRAIN' or 'Mode.EVAL' instead of their string values ('train', + 'eval'). This normalizes them back to lowercase standard modes. + + Args: + event: The raw metric event name. + + Returns: + The normalized metric event name. + """ + return ( + event.replace("Mode.TRAIN", "train") + .replace("Mode.Train", "train") + .replace("Mode.EVAL", "eval") + .replace("Mode.Eval", "eval") + ) + + +def _extract_scalar(value: Any) -> int | float | None: + """Extracts a primitive Python float or int from an array, tensor, or scalar. + + Args: + value: A scalar, array, or tensor value. + + Returns: + An int or float scalar value, or None if extraction fails or value is bool. + """ + if isinstance(value, (str, bytes)): + return None + try: + if hasattr(value, "item"): + val = value.item() + elif isinstance(value, (int, float, np.number)): + val = value + else: + val = float(value) + + if isinstance(val, (str, bytes, bool, np.bool_)): + return None + if isinstance(val, (int, np.integer)): + return int(val) + if isinstance(val, (float, np.floating)): + return float(val) + return float(val) + except (TypeError, ValueError, AttributeError, RuntimeError): + return None + + +class MLDiagScalarBackend: + """Routes all scalar metrics to Google Cloud ML Diagnostics.""" + + def __init__(self, config: Any | None = None) -> None: + """Initializes the ML Diagnostics scalar backend.""" + if not _HAS_MLDIAG: + logging.warning( + "google_cloud_mldiagnostics is not installed; MLDiagScalarBackend is" + " disabled." + ) + if config is not None: + managed_mldiagnostics.ManagedMLDiagnostics(config) + + def log_scalar( + self, + event: str, + value: Any, + step: int | None = None, + **kwargs: Any, + ) -> None: + """Logs a single scalar metric to Google Cloud ML Diagnostics. + + Args: + event: Hierarchical metric name (e.g. 'actor/train/loss', + 'rewards/train/score/mean'). + value: Metric scalar value (can be float, int, or jnp/np array). + step: Training or evaluation step index. + **kwargs: Additional metadata keywords. + """ + if not _HAS_MLDIAG or mldiag_metrics is None or jax.process_index() != 0: + return + + val = _extract_scalar(value) + if val is None or math.isnan(val) or math.isinf(val): + return + + try: + event = _normalize_metric_event(event) + metric_key = event + if event in _EXACT_METRIC_MAP: + enum_attr = _EXACT_METRIC_MAP[event] + if metric_types is not None and hasattr( + metric_types.MetricType, enum_attr + ): + metric_key = getattr(metric_types.MetricType, enum_attr) + + int_step = int(step) if step is not None else None + timestamp = kwargs.get("timestamp") + if timestamp is not None: + mldiag_metrics.record( + metric_key, val, step=int_step, timestamp=timestamp + ) + else: + mldiag_metrics.record(metric_key, val, step=int_step) + except Exception as e: # pylint: disable=broad-exception-caught + logging.warning( + "Failed to record metric '%s' to ML Diagnostics: %s", event, e + ) + + def close(self) -> None: + """Closes the ML Diagnostics backend.""" + pass diff --git a/src/maxtext/trainers/post_train/rl/train_rl.py b/src/maxtext/trainers/post_train/rl/train_rl.py index 3296971007..36e7af9de4 100644 --- a/src/maxtext/trainers/post_train/rl/train_rl.py +++ b/src/maxtext/trainers/post_train/rl/train_rl.py @@ -450,7 +450,8 @@ def create_rl_components( # pylint: disable=too-many-positional-arguments # Setup checkpointing if trainer_config.enable_checkpointing: checkpointing_options = ocp.CheckpointManagerOptions( - save_interval_steps=trainer_config.checkpoint_period, max_to_keep=trainer_config.max_num_checkpoints_to_keep + save_interval_steps=trainer_config.checkpoint_period, + max_to_keep=trainer_config.max_num_checkpoints_to_keep, ) checkpoint_dir = trainer_config.checkpoint_dir else: @@ -458,15 +459,41 @@ def create_rl_components( # pylint: disable=too-many-positional-arguments checkpoint_dir = None # Set up micro batching - train_micro_batch_size = None if trainer_config.train_micro_batch_size == -1 else trainer_config.train_micro_batch_size + train_micro_batch_size = ( + None + if trainer_config.train_micro_batch_size == -1 + else trainer_config.train_micro_batch_size + ) rollout_micro_batch_size = ( - None if trainer_config.rollout_micro_batch_size == -1 else trainer_config.rollout_micro_batch_size + None + if trainer_config.rollout_micro_batch_size == -1 + else trainer_config.rollout_micro_batch_size ) # Setup metrics logging - metrics_logging_options = metrics_logger.MetricsLoggerOptions( - log_dir=trainer_config.tensorboard_dir, flush_every_n_steps=trainer_config.log_period - ) + if getattr(trainer_config, "managed_mldiagnostics", False): + # pylint: disable=g-import-not-at-top,import-outside-toplevel + from maxtext.trainers.post_train.rl import mldiag_backend + # pylint: enable=g-import-not-at-top,import-outside-toplevel + + backend_list = [lambda: mldiag_backend.MLDiagScalarBackend(trainer_config)] + if getattr(trainer_config, "enable_tensorboard", False): + backend_list.append( + lambda: metrics_logger.TensorboardBackend( # pylint: disable=unnecessary-lambda + log_dir=trainer_config.tensorboard_dir, + flush_every_n_steps=trainer_config.log_period, + ) + ) + metrics_logging_options = metrics_logger.MetricsLoggerOptions( + log_dir=trainer_config.tensorboard_dir or "", + flush_every_n_steps=trainer_config.log_period, + backend_kwargs={"custom_backend": backend_list}, + ) + else: + metrics_logging_options = metrics_logger.MetricsLoggerOptions( + log_dir=trainer_config.tensorboard_dir or "", + flush_every_n_steps=trainer_config.log_period, + ) profiler_options = None if trainer_config.profiler == "xplane": diff --git a/tests/unit/mldiag_backend_test.py b/tests/unit/mldiag_backend_test.py new file mode 100644 index 0000000000..96bee6dc33 --- /dev/null +++ b/tests/unit/mldiag_backend_test.py @@ -0,0 +1,319 @@ +# Copyright 2023–2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for MLDiagScalarBackend.""" + +from unittest import mock + +from absl.testing import absltest +import jax.numpy as jnp +import numpy as np + +from maxtext.trainers.post_train.rl import mldiag_backend + +_RECORD_TARGET = ( + "maxtext.src.maxtext.trainers.post_train.rl.mldiag_backend." + "mldiag_metrics.record" +) +_METRICS_TARGET = ( + "maxtext.src.maxtext.trainers.post_train.rl.mldiag_backend." + "mldiag_metrics" +) +_MLDIAG_TARGET = ( + "maxtext.src.maxtext.trainers.post_train.rl.mldiag_backend.mldiag" +) + + +class MLDiagScalarBackendTest(absltest.TestCase): + + def test_extract_scalar_valid_types(self): + self.assertEqual(mldiag_backend._extract_scalar(42), 42) + self.assertEqual(mldiag_backend._extract_scalar(3.14), 3.14) + self.assertEqual(mldiag_backend._extract_scalar(np.float32(1.5)), 1.5) + self.assertEqual(mldiag_backend._extract_scalar(np.int64(10)), 10) + val = mldiag_backend._extract_scalar(jnp.array(2.718)) + self.assertIsNotNone(val) + self.assertAlmostEqual(val, 2.718, places=3) + self.assertEqual(mldiag_backend._extract_scalar(np.array([5.0])), 5.0) + + def test_extract_scalar_type_preservation(self): + val_int = mldiag_backend._extract_scalar(42) + self.assertIsInstance(val_int, int) + self.assertEqual(val_int, 42) + + val_np_int = mldiag_backend._extract_scalar(np.int64(10)) + self.assertIsInstance(val_np_int, int) + self.assertEqual(val_np_int, 10) + + val_np_int32 = mldiag_backend._extract_scalar(np.int32(7)) + self.assertIsInstance(val_np_int32, int) + self.assertEqual(val_np_int32, 7) + + val_float = mldiag_backend._extract_scalar(3.14) + self.assertIsInstance(val_float, float) + self.assertEqual(val_float, 3.14) + + val_np_float = mldiag_backend._extract_scalar(np.float32(1.5)) + self.assertIsInstance(val_np_float, float) + self.assertEqual(val_np_float, 1.5) + + def test_extract_scalar_invalid_types(self): + self.assertIsNone(mldiag_backend._extract_scalar(None)) + self.assertIsNone(mldiag_backend._extract_scalar(True)) + self.assertIsNone(mldiag_backend._extract_scalar(False)) + self.assertIsNone(mldiag_backend._extract_scalar(np.bool_(True))) + self.assertIsNone(mldiag_backend._extract_scalar(np.bool_(False))) + self.assertIsNone(mldiag_backend._extract_scalar(np.array(True))) + self.assertIsNone(mldiag_backend._extract_scalar(np.array([False]))) + self.assertIsNone(mldiag_backend._extract_scalar([1, 2, 3])) + self.assertIsNone(mldiag_backend._extract_scalar(np.array([1.0, 2.0]))) + self.assertIsNone(mldiag_backend._extract_scalar("not_a_number")) + self.assertIsNone(mldiag_backend._extract_scalar("123")) + self.assertIsNone(mldiag_backend._extract_scalar("3.14")) + self.assertIsNone(mldiag_backend._extract_scalar(b"123")) + self.assertIsNone(mldiag_backend._extract_scalar(b"3.14")) + self.assertIsNone(mldiag_backend._extract_scalar(b"not_a_number")) + self.assertIsNone(mldiag_backend._extract_scalar(np.array("123"))) + self.assertIsNone(mldiag_backend._extract_scalar(np.array(["123"]))) + self.assertIsNone(mldiag_backend._extract_scalar(np.array(b"123"))) + self.assertIsNone(mldiag_backend._extract_scalar(np.str_("123"))) + self.assertIsNone(mldiag_backend._extract_scalar(np.bytes_(b"123"))) + + def test_log_scalar_non_zero_process_index_noop(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=1): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("actor/train/loss", 0.5, step=1) + mock_record.assert_not_called() + + def test_log_scalar_nan_inf_none_ignored(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("actor/train/loss", float("nan"), step=1) + backend.log_scalar("actor/train/loss", float("inf"), step=1) + backend.log_scalar("actor/train/loss", -float("inf"), step=1) + backend.log_scalar("actor/train/loss", None, step=1) + mock_record.assert_not_called() + + def test_log_scalar_string_and_bytes_ignored(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("actor/train/loss", "0.5", step=1) + backend.log_scalar("actor/train/loss", "123", step=1) + backend.log_scalar("actor/train/loss", b"0.5", step=1) + backend.log_scalar("actor/train/loss", np.array("0.5"), step=1) + backend.log_scalar("actor/train/loss", True, step=1) + mock_record.assert_not_called() + + def test_log_scalar_exact_metric_mapping(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("loss", 0.125, step=10) + mock_record.assert_called_once() + args, kwargs = mock_record.call_args + expected_key = ( + mldiag_backend.metric_types.MetricType.LOSS + if ( + mldiag_backend.metric_types is not None + and hasattr(mldiag_backend.metric_types.MetricType, "LOSS") + ) + else "loss" + ) + self.assertEqual(args[0], expected_key) + self.assertAlmostEqual(args[1], 0.125) + self.assertEqual(kwargs.get("step"), 10) + + def test_log_scalar_all_exact_metric_mappings(self): + backend = mldiag_backend.MLDiagScalarBackend() + for event_name, enum_name in mldiag_backend._EXACT_METRIC_MAP.items(): + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar(event_name, 1.23, step=7) + mock_record.assert_called_once() + args, kwargs = mock_record.call_args + expected_key = ( + getattr(mldiag_backend.metric_types.MetricType, enum_name) + if ( + mldiag_backend.metric_types is not None + and hasattr(mldiag_backend.metric_types.MetricType, enum_name) + ) + else event_name + ) + self.assertEqual(args[0], expected_key) + self.assertAlmostEqual(args[1], 1.23) + self.assertEqual(kwargs.get("step"), 7) + + def test_log_scalar_hierarchical_event_preserves_full_name(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("actor/train/loss", 0.125, step=10) + mock_record.assert_called_once_with("actor/train/loss", 0.125, step=10) + + def test_log_scalar_multi_model_metrics_no_collision(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("actor/train/loss", 0.25, step=1) + backend.log_scalar("critic/train/loss", 0.5, step=1) + backend.log_scalar("rewards/score", 1.0, step=1) + self.assertEqual(mock_record.call_count, 3) + self.assertEqual( + mock_record.call_args_list[0], + mock.call("actor/train/loss", 0.25, step=1), + ) + self.assertEqual( + mock_record.call_args_list[1], + mock.call("critic/train/loss", 0.5, step=1), + ) + self.assertEqual( + mock_record.call_args_list[2], + mock.call("rewards/score", 1.0, step=1), + ) + + def test_log_scalar_unmapped_metric_passes_full_event_name(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("custom/reward/score", 1.0, step=5) + mock_record.assert_called_once_with("custom/reward/score", 1.0, step=5) + + def test_log_scalar_step_zero_preserved(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("actor/train/loss", 0.5, step=0) + mock_record.assert_called_once_with("actor/train/loss", 0.5, step=0) + + def test_log_scalar_exception_safe(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch( + _RECORD_TARGET, + side_effect=RuntimeError("RPC failure"), + ): + # Should not raise exception + backend.log_scalar("actor/train/loss", 0.5, step=1) + + def test_log_scalar_passes_timestamp_kwarg(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar( + "custom/reward/score", + 1.0, + step=5, + timestamp="2026-08-30T12:00:00Z", + ) + mock_record.assert_called_once_with( + "custom/reward/score", + 1.0, + step=5, + timestamp="2026-08-30T12:00:00Z", + ) + + def test_close_callable(self): + backend = mldiag_backend.MLDiagScalarBackend() + # Should execute cleanly without error + backend.close() + + def test_init_with_config_calls_managed_mldiagnostics(self): + mock_config = mock.MagicMock() + with mock.patch.object( + mldiag_backend.managed_mldiagnostics, + "ManagedMLDiagnostics", + ) as mock_managed: + mldiag_backend.MLDiagScalarBackend(mock_config) + mock_managed.assert_called_once_with(mock_config) + + mock_managed.reset_mock() + with mock.patch.object( + mldiag_backend.managed_mldiagnostics, + "ManagedMLDiagnostics", + ) as mock_managed: + mldiag_backend.MLDiagScalarBackend(config=mock_config) + mock_managed.assert_called_once_with(mock_config) + + def test_init_without_config_does_not_call_managed_mldiagnostics(self): + with mock.patch.object( + mldiag_backend.managed_mldiagnostics, + "ManagedMLDiagnostics", + ) as mock_managed: + mldiag_backend.MLDiagScalarBackend() + mldiag_backend.MLDiagScalarBackend(None) + mock_managed.assert_not_called() + + def test_normalize_metric_event(self): + self.assertEqual( + mldiag_backend._normalize_metric_event("actor/Mode.TRAIN/loss"), + "actor/train/loss", + ) + self.assertEqual( + mldiag_backend._normalize_metric_event("rewards/Mode.EVAL/score/mean"), + "rewards/eval/score/mean", + ) + self.assertEqual( + mldiag_backend._normalize_metric_event("global/Mode.Train/throughput"), + "global/train/throughput", + ) + self.assertEqual( + mldiag_backend._normalize_metric_event("Mode.Eval/latency"), + "eval/latency", + ) + self.assertEqual( + mldiag_backend._normalize_metric_event("actor/train/loss"), + "actor/train/loss", + ) + self.assertEqual( + mldiag_backend._normalize_metric_event("loss"), + "loss", + ) + self.assertEqual( + mldiag_backend._normalize_metric_event(""), + "", + ) + + def test_log_scalar_normalizes_mode_enum_in_event_name(self): + backend = mldiag_backend.MLDiagScalarBackend() + with mock.patch("jax.process_index", return_value=0): + with mock.patch(_RECORD_TARGET) as mock_record: + backend.log_scalar("actor/Mode.TRAIN/loss", 0.5, step=10) + backend.log_scalar("rewards/Mode.EVAL/score/mean", 1.5, step=10) + backend.log_scalar("global/Mode.Train/throughput", 100.0, step=10) + backend.log_scalar("Mode.Eval/latency", 20.0, step=10) + + self.assertEqual(mock_record.call_count, 4) + self.assertEqual( + mock_record.call_args_list[0], + mock.call("actor/train/loss", 0.5, step=10), + ) + self.assertEqual( + mock_record.call_args_list[1], + mock.call("rewards/eval/score/mean", 1.5, step=10), + ) + self.assertEqual( + mock_record.call_args_list[2], + mock.call("global/train/throughput", 100.0, step=10), + ) + self.assertEqual( + mock_record.call_args_list[3], + mock.call("eval/latency", 20.0, step=10), + ) + + +if __name__ == "__main__": + absltest.main()