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
85 changes: 85 additions & 0 deletions docs/metrics.md
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@ For Reinforcement Learning jobs, Tunix collects additional metrics related to
the RL algorithm (e.g., PPO), reward modeling, and generation.

#### Rewards & Scores

* **`rewards/sum`**: The sum of rewards for a trajectory.
* **`rewards/mean`**, **`rewards/max`**, **`rewards/min`**: Statistics of the
rewards across the batch.
Expand All @@ -67,6 +68,7 @@ the RL algorithm (e.g., PPO), reward modeling, and generation.
individual reward components are logged by name.

#### Policy & Value (PPO)

* **`advantages/mean`**, **`advantages/max`**, **`advantages/min`**:
Statistics of the advantages.
* **`returns/mean`**, **`returns/max`**, **`returns/min`**: Statistics of
Expand All @@ -81,6 +83,7 @@ the RL algorithm (e.g., PPO), reward modeling, and generation.
enabled).

#### Generation & Data

* **`prompts`**: The input prompts used for generation.
* **`completions`**: The text completions generated by the model.
* **`completions/mean_length`**, **`completions/max_length`**,
Expand Down Expand Up @@ -210,6 +213,88 @@ depends on the execution environment.
* `log_dir`: Directory where event files are written.
* `flush_every_n_steps`: Frequency of flushing logs to disk (default: 100).

### Experimental: OpenTelemetry double-write

Tunix is evaluating [OpenTelemetry](https://opentelemetry.io/) as a
vendor-neutral instrumentation layer. As an opt-in first step, `MetricsLogger`
supports a **double-write** mode: every scalar passed to `log()` is emitted both
through the existing Metrax/`jax.monitoring` backends (unchanged, still the
default) and as an OpenTelemetry gauge measurement. This lets you compare the
two pipelines side by side before any switch is made.

Enable it with the `enable_opentelemetry` flag (requires `pip install
'google-tunix[otel]'`):

```python
# OTLP exporter: pip install opentelemetry-exporter-otlp
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
OTLPMetricExporter,
)
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader
from tunix.sft import metrics_logger

reader = PeriodicExportingMetricReader(OTLPMetricExporter())
provider = MeterProvider(metric_readers=[reader])

options = metrics_logger.MetricsLoggerOptions(
log_dir="/tmp/logs", # Existing backends keep working unchanged.
enable_opentelemetry=True, # Additionally emit OpenTelemetry gauges.
)
logger = metrics_logger.MetricsLogger(options, otel_meter_provider=provider)
```

Key properties of the double-write path:

* **Off by default.** With the flag unset, behavior is identical to previous
releases; the OpenTelemetry API does not even need to be installed.
* **Application-owned providers.** Tunix never configures exporters or shuts
down providers. Configure the global `MeterProvider` (or inject one via the
keyword-only `otel_meter_provider` argument) and own its flush/shutdown.
`MetricsLogger.close()` only closes the Metrax backends.
* **Process policy.** Like the Metrax backends, only JAX process 0 emits.
* **Naming.** Known metrics map to stable instrument names (`loss` →
`tunix.training.loss`, `perplexity` → `tunix.training.perplexity`,
`learning_rate` → `tunix.training.learning_rate`, `grad_norm` →
`tunix.training.gradient.norm`); other metric keys are normalized into the
`tunix.*` namespace (e.g. `rewards/score mean` →
`tunix.rewards.score.mean`). The legacy prefix and mode become the
low-cardinality attributes `tunix.metrics.prefix` and `tunix.training.mode`,
and the logical step is emitted as a separate `tunix.training.step` gauge.

#### OpenTelemetry → Weights & Biases

W&B ingests OpenTelemetry traces natively (via the
[Weave OTLP endpoint](https://docs.wandb.ai/weave/guides/tracking/otel)), but
has no OTLP endpoint for run metrics. Tunix therefore ships
`tunix.sft.otel_wandb.WandbMetricsExporter`, an OpenTelemetry SDK metric
exporter that forwards the gauges above to `wandb.log`, using the familiar
`{prefix}/{mode}/{name}` chart keys (e.g. `actor/train/tunix.training.loss`) and
the `tunix.training.step` gauge as the W&B step:

```python
import wandb
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader
from tunix.sft import metrics_logger
from tunix.sft import otel_wandb

run = wandb.init(project="my-project")
reader = PeriodicExportingMetricReader(
otel_wandb.WandbMetricsExporter(run), export_interval_millis=10_000
)
provider = MeterProvider(metric_readers=[reader])

options = metrics_logger.MetricsLoggerOptions(
log_dir="/tmp/logs",
enable_opentelemetry=True,
)
logger = metrics_logger.MetricsLogger(options, otel_meter_provider=provider)
```

Because the Metrax `WandbBackend` also remains active by default, you can point
both pipelines at the same W&B project and compare the charts directly.

### Custom metric logger

You can integrate any logging service by creating a custom backend that conforms
Expand Down
7 changes: 7 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -69,10 +69,17 @@ test = [
"pytest-xdist",
"gcsfs",
"chex",
"opentelemetry-sdk>=1.25.0",
"wandb", # Offline end-to-end test for the OpenTelemetry W&B exporter
]
cli = [
"tensorflow",
]
otel = [
# Experimental OpenTelemetry double-write path for MetricsLogger.
"opentelemetry-api>=1.25.0",
"opentelemetry-sdk>=1.25.0",
]
frozenlake = [
"gymnasium",
]
Expand Down
41 changes: 41 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
"""Shared test configuration.

The test environment installs ``wandb`` for the OpenTelemetry W&B exporter
end-to-end test. Installing it has a side effect on the legacy path: the
default Metrax ``WandbBackend`` becomes constructible, registers a
process-global ``jax.monitoring`` listener that receives JAX's own internal
compile metrics, and shares the module-global ``wandb.run`` across logger
instances — one logger's ``close()`` (``wandb.finish()``) then breaks any
other still-registered backend in the process.

To keep the legacy default-backend behavior identical to an environment
without ``wandb`` (which is how CI has always run), the default
``WandbBackend`` is replaced with a stub that raises ``ImportError``, which
``MetricsLoggerOptions.create_backends`` already handles by skipping the
backend. Tests that exercise wandb do so explicitly: the OpenTelemetry
exporter tests pass an offline ``wandb.init`` run object directly and never
touch the Metrax backend.
"""

import os

import pytest

os.environ.setdefault("WANDB_MODE", "disabled")


class _WandbBackendUnavailable:

def __init__(self, *args, **kwargs):
raise ImportError(
"The default Metrax WandbBackend is disabled in tests; construct a"
" wandb run explicitly instead."
)


@pytest.fixture(autouse=True)
def _disable_default_wandb_backend(monkeypatch):
monkeypatch.setattr(
"tunix.sft.metrics_logger.WandbBackend", _WandbBackendUnavailable
)
yield
128 changes: 128 additions & 0 deletions tests/sft/metrics_logger_test.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Metrics logger unittest."""

import collections
import copy
import os
import tempfile
Expand All @@ -8,6 +9,9 @@
from absl.testing import absltest
import jax
import metrax.logging as metrax_logging
import numpy as np
from opentelemetry.sdk import metrics as otel_sdk_metrics
from opentelemetry.sdk.metrics import export as otel_sdk_export
from tunix.sft import metrics_logger
from tunix.utils import env_utils

Expand Down Expand Up @@ -229,5 +233,129 @@ def test_backend_kwargs_are_passed_to_backends(self):
)


def _gauge_points(metrics_data):
"""Maps instrument name to a list of (value, attributes) data points."""
points = collections.defaultdict(list)
if metrics_data is None:
return points
for resource_metrics in metrics_data.resource_metrics:
for scope_metrics in resource_metrics.scope_metrics:
for metric in scope_metrics.metrics:
for point in metric.data.data_points:
points[metric.name].append(
(getattr(point, "value", None), dict(point.attributes))
)
return points


class OpenTelemetryDoubleWriteTest(absltest.TestCase):
"""Tests the opt-in OpenTelemetry double-write path."""

def setUp(self):
super().setUp()
self._temp_dir_obj = tempfile.TemporaryDirectory()
self.log_dir = self._temp_dir_obj.name
self.backend = mock.Mock(spec=metrax_logging.LoggingBackend)
self.enter_context(
mock.patch.object(jax.monitoring, "register_scalar_listener")
)
self.mock_record_scalar = self.enter_context(
mock.patch.object(jax.monitoring, "record_scalar")
)
self.reader = otel_sdk_export.InMemoryMetricReader()
self.meter_provider = otel_sdk_metrics.MeterProvider(
metric_readers=[self.reader]
)

def _make_logger(self, enable_opentelemetry=True):
options = metrics_logger.MetricsLoggerOptions(
log_dir=self.log_dir,
backend_kwargs={"custom_backend": [lambda: self.backend]},
enable_opentelemetry=enable_opentelemetry,
)
return metrics_logger.MetricsLogger(
metrics_logger_options=options,
otel_meter_provider=self.meter_provider,
)

def test_disabled_by_default_and_legacy_path_unchanged(self):
logger = self._make_logger(enable_opentelemetry=False)
logger.log("actor", "loss", 0.5, metrics_logger.Mode.TRAIN, 1)

self.mock_record_scalar.assert_called_once_with(
"actor/train/loss", 0.5, step=1
)
self.assertEqual(_gauge_points(self.reader.get_metrics_data()), {})
self.assertAlmostEqual(logger.get_metric("actor", "loss", "train"), 0.5)

def test_double_write_emits_otel_and_legacy(self):
logger = self._make_logger()
logger.log("actor", "loss", 0.5, metrics_logger.Mode.TRAIN, 7)

# Legacy jax.monitoring write is unchanged.
self.mock_record_scalar.assert_called_once_with(
"actor/train/loss", 0.5, step=7
)
# OpenTelemetry gauge is additionally emitted.
points = _gauge_points(self.reader.get_metrics_data())
expected_attributes = {
"tunix.training.mode": "train",
"tunix.metrics.prefix": "actor",
}
self.assertEqual(
points["tunix.training.loss"], [(0.5, expected_attributes)]
)
self.assertEqual(points["tunix.training.step"], [(7, expected_attributes)])
# Local history is unchanged.
self.assertAlmostEqual(logger.get_metric("actor", "loss", "train"), 0.5)

def test_dynamic_metric_names_are_normalized(self):
logger = self._make_logger()
logger.log("actor", "rewards/score mean", 1.5, "train", 1)

points = _gauge_points(self.reader.get_metrics_data())
self.assertIn("tunix.rewards.score.mean", points)

def test_non_scalar_values_stay_in_history_only(self):
logger = self._make_logger()
logger.log("actor", "loss", np.array([0.5, 0.6]), "train", 1)

points = _gauge_points(self.reader.get_metrics_data())
self.assertNotIn("tunix.training.loss", points)
self.assertLen(logger.get_metric_history("actor", "loss", "train"), 1)

def test_repeated_step_emits_step_gauge_once(self):
logger = self._make_logger()
logger.log("actor", "loss", 0.5, "train", 3)
logger.log("actor", "perplexity", 1.6, "train", 3)

points = _gauge_points(self.reader.get_metrics_data())
self.assertLen(points["tunix.training.step"], 1)

@mock.patch.object(jax, "process_index", return_value=1)
def test_no_otel_emission_on_secondary_process(self, mock_process_index):
del mock_process_index
logger = self._make_logger()
logger.log("actor", "loss", 0.5, "train", 1)

self.assertEqual(_gauge_points(self.reader.get_metrics_data()), {})

def test_missing_otel_api_raises_when_enabled(self):
with mock.patch.object(metrics_logger, "otel_metrics", new=None):
with self.assertRaisesRegex(ImportError, "google-tunix\\[otel\\]"):
self._make_logger()

def test_close_leaves_meter_provider_running(self):
logger = self._make_logger()
logger.log("actor", "loss", 0.5, "train", 1)
logger.close()

self.backend.close.assert_called_once()
logger_after_close = self._make_logger()
logger_after_close.log("actor", "loss", 0.25, "train", 2)
points = _gauge_points(self.reader.get_metrics_data())
self.assertIn(0.25, [value for value, _ in points["tunix.training.loss"]])


if __name__ == "__main__":
absltest.main()
Loading
Loading