This guide describes how to capture, monitor, and visualize training, system, and performance metrics in MaxDiffusion using the Google Cloud ML Diagnostics SDK (google-cloud-mldiagnostics).
MaxDiffusion integrates with Google Cloud ML Diagnostics to provide real-time telemetry during training runs on Cloud TPUs:
- Workload Metrics: In multi-host JAX jobs, step-level metrics (loss, step time, learning rate, gradient norm, parameter weights, custom activations) are buffered and dispatched from the master node (process index 0) to prevent duplicate logs.
- System & Accelerator Metrics: The SDK automatically runs background daemon threads on all worker hosts to capture hardware utilization (
tpu_duty_cycle,hbm_utilization,host_cpu_utilization,host_memory_utilization). - Cloud Logging Sink: Metrics are written to Google Cloud Logging.
- Control Plane UI: The Diagnostics Console automatically discovers and renders standard and custom metric plots.
MaxDiffusion automatically translates internal scalar keys to canonical MetricType enums expected by the Control Plane UI:
- Loss (
loss): Training loss value per step (mapped fromlearning/loss). - Learning Rate (
learning_rate): Current optimizer learning rate (mapped fromlearning/current_learning_rate). - Gradient Norm (
gradient_norm): Global L2 norm of model gradients (mapped fromlearning/grad_norm). - Total Weights (
total_weights): Total trainable model parameter count (mapped fromlearning/total_weights). - Step Time (
step_time): Duration of each training step in seconds (mapped fromperf/step_time_seconds). - TFLOPS (
tflops): Hardware compute throughput per accelerator in TFLOP/s (mapped fromperf/per_device_tflops_per_sec).
Any key in metrics["scalar"] that is not part of _METRICS_TO_MANAGED is treated as a Custom Metric:
- Retains its raw string name (e.g.,
"custom/latents_mean","snr_loss_weight","cross_attn_entropy"). - Are dynamically discovered by the Control Plane UI and rendered in dedicated chart cards (
Over TimeandOver Steps).
When enable_ml_diagnostics=True is enabled, the SDK automatically captures:
tpu_duty_cycle: Core accelerator compute utilization percentage.hbm_utilization: High Bandwidth Memory consumed percentage.host_cpu_utilization: Host CPU usage percentage.host_memory_utilization: Host system RAM usage percentage.
Metric mapping and dispatch are centralized in train_utils.py and max_utils.py. Authors of training scripts can integrate metrics using two steps:
Initialize the run at the start of training:
from maxdiffusion import max_utils
max_utils.ensure_machinelearning_job_runs(config)Inside the trainer's training_loop():
from maxdiffusion import train_utils
# Record standard step metrics (and any custom metrics in train_metric["scalar"]):
train_utils.record_scalar_metrics(
train_metric,
step_time_delta,
self.per_device_tflops,
learning_rate_scheduler(step),
)
if self.config.write_metrics:
train_utils.write_metrics(writer, local_metrics_file, running_gcs_metrics, train_metric, step, self.config)Enable ML Diagnostics via YAML configuration files (using enable_ml_diagnostics: True) or command-line flags:
# src/maxdiffusion/configs/base_2_base.yml
run_name: "my-training-run"
enable_ml_diagnostics: True
write_metrics: True
log_period: 10Note
To enable automated profiling and on-demand XProf traces alongside metrics, see the ML Diagnostics Profiling Guide.
Run command:
python -m src.maxdiffusion.train src/maxdiffusion/configs/base_2_base.yml \
run_name=my-training-run \
output_dir=gs://my-bucket/output \
enable_ml_diagnostics=True \
write_metrics=TrueInspect metric logs directly using gcloud:
# Query loss metrics
gcloud logging read 'logName="projects/<PROJECT_ID>/logs/ml_diagnostics_metric" AND resource.labels.namespace="loss"' \
--limit=5 \
--format="json"
# Query custom metrics
gcloud logging read 'logName="projects/<PROJECT_ID>/logs/ml_diagnostics_metric" AND resource.labels.namespace="custom/latents_mean"' \
--limit=5 \
--format="json"
# Query hardware metrics
gcloud logging read 'logName="projects/<PROJECT_ID>/logs/ml_diagnostics_metric" AND resource.labels.namespace="hbm_utilization"' \
--limit=5 \
--format="json"- Open Google Cloud Console and navigate to Hypercompute Clusters → Diagnostics.
- Select your cluster and active
MachineLearningRun. - Inspect:
- Model Metrics: View predefined plots for
loss,learning_rate,gradient_norm, andtotal_weights. - Custom Metrics: View dynamically generated charts for all
custom/*metrics over time and steps. - Performance: View
step_time,tflops,tpu_duty_cycle, andhbm_utilization.
- Model Metrics: View predefined plots for