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
178 changes: 178 additions & 0 deletions src/maxtext/trainers/post_train/rl/mldiag_backend.py
Original file line number Diff line number Diff line change
@@ -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
39 changes: 33 additions & 6 deletions src/maxtext/trainers/post_train/rl/train_rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -450,23 +450,50 @@ 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:
checkpointing_options = None
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":
Expand Down
Loading
Loading