Skip to content

Commit 88b5e0e

Browse files
feat: MLFlow exporter (#2606)
<!-- Thanks for contributing to NeMo Gym! Please fill out the sections below. --> ## What does this PR do? This PR extends the experiment tracking possibilities by adding MLFlow backend. Till now there was a W&B exporter wired in the code, that tracks configuration, metrics and rollouts. This PR adds an exporter abstraction, moves existing functionality into W&B backend and adds MLFlow backend. Both can be configured at the same time - artifacts are sent to both tracking servers. **Deprecated field:** `upload_rollouts_to_wandb` has been renamed to `upload_rollouts`. If the user provides an old key, it is mapped to the new one and a deprecation warning is emitted. ## Followup 1. Export file artifacts. The original W&B code logs rollouts only, while NEL was logging more. This PR just adds a new backend and doesn't provide full feature parity with NEL. Tracked with #2619 2. Placement of export initialization seems sub-optimal. In my opinion it belongs to CLI, not config parsing. As a side effect we don't have the user-friendly error handling on config error in exporter. Tracked with #2620 3. Clean `RewardProfiler` up. It looks like it uses `wandb.Histrogram` for no good reason and carries some other obsolete code. Tracked with #2621 ## Usage example ``` gym eval run \ --benchmark gpqa \ --model-type openai_model \ --output gpqa-results/rollouts.jsonl \ --split benchmark \ +mlflow_tracking_uri=https://mlflow.my-server.com/ \ +mlflow_experiment_name=my-eval \ +mlflow_run_name=gpqa ``` This will run the evaluation as usual and create an MLFlow experiment for the run: ``` ... INFO: Shutting down INFO: Waiting for application shutdown. INFO: Application shutdown complete. INFO: Finished server process [2090575] NeMo Gym finished! Shutting down Ray cluster spun up by NeMo Gym... 🏃 View run gpqa at: https://mlflow.my-server.com/#/experiments/1740/runs/f088e196471048c9bf20e9beb9b771f5 🧪 View experiment at: https://mlflow.my-server.com/#/experiments/1740 ``` ## Checklist - [x] I have read the [contributing guidelines](https://docs.nvidia.com/nemo/gym/latest/contribute/development-setup). - [x] The change is focused; unrelated "drive-by" edits are tracked as separate issues/PRs. - [x] Tests added or updated and pass locally, or N/A for docs-only / non-code changes (so CI unit/server checks pass when applicable). - [x] Pre-commit checks pass locally (`pre-commit run --all-files`) (so CI lint/format/copyright pass). - [x] All commits have DCO sign-off (`git commit -s`) (so the DCO check passes). --------- Signed-off-by: Marta Stepniewska-Dziubinska <martas@nvidia.com> Co-authored-by: bxyu-nvidia <bxyu@nvidia.com>
1 parent fbd379f commit 88b5e0e

18 files changed

Lines changed: 1272 additions & 111 deletions

‎benchmarks/nemotron_3.5_super/sbatch_external_vllm.sh‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,7 @@ experiment_name=$EXPERIMENT_NAME/slurm_job_id_\$SLURM_JOB_ID/date_\$(date +%Y%m%
4444
# +uv_venv_dir=/opt/uv_venvs is from the container.
4545
# +skip_venv_if_present=true will reuse the venvs baked into the container if possible.
4646
# ++use_absolute_ip=true: Necessary for communication between harness in sandbox and Gym model servers
47-
# ++upload_rollouts_to_wandb=false: Rollouts file is massive. We leave on the cluster.
47+
# ++upload_rollouts=false: Rollouts file is massive. We leave on the cluster.
4848
# global_aiohttp_connector_limit_per_host: 16k concurrent requests should be enough. We can raise further if our inference is efficient enough to support.
4949
# port_range_low, port_range_high: Move into ephemeral ports
5050
gym eval run \
@@ -62,7 +62,7 @@ gym eval run \
6262
++policy_base_url=http://\$(getent hosts "\$ROUTER_NODE" | awk 'NR == 1 {print \$1}'):$ROUTER_SERVER_PORT/v1 \
6363
++policy_api_key=dummy_api_key \
6464
++policy_model_name=$MODEL \
65-
++upload_rollouts_to_wandb=false \
65+
++upload_rollouts=false \
6666
++global_aiohttp_connector_limit_per_host=16384 \
6767
++port_range_low=63000 \
6868
++port_range_high=64000

‎benchmarks/osworld/prepare.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -376,7 +376,7 @@ def write_env(
376376
f"output_jsonl_fpath: {_yaml_string(output_jsonl)}",
377377
"num_repeats: 1",
378378
f"num_samples_in_parallel: {num_samples_in_parallel}",
379-
"upload_rollouts_to_wandb: false",
379+
"upload_rollouts: false",
380380
"responses_create_params:",
381381
f" max_output_tokens: {max_output_tokens}",
382382
f" temperature: {temperature}",

‎fern/versions/latest/pages/reference/configuration.mdx‎

Lines changed: 59 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -199,6 +199,11 @@ use_absolute_ip: true # Bind servers to the host's IP instead of 127.0
199199

200200
# Optional: validation behavior
201201
error_on_almost_servers: true # Exit on invalid configs (default: true)
202+
203+
# Optional: experiment tracking (see Experiment Tracking below)
204+
wandb_project: gym-dev
205+
wandb_name: my-run
206+
wandb_api_key: your-wandb-key
202207
```
203208
204209
### Multi-Node Configuration
@@ -211,6 +216,60 @@ error_on_almost_servers: true # Exit on invalid configs (default: true)
211216

212217
---
213218

219+
## Experiment Tracking
220+
221+
Gym exports the resolved config, aggregate metrics, and rollouts to any tracking backend you
222+
configure. Weights & Biases and MLflow are both supported, and both can be on at once. A backend
223+
is used only when every key it needs is set; otherwise it is silently skipped.
224+
225+
### Weights & Biases
226+
227+
| Key | Required | Description |
228+
|-----|----------|-------------|
229+
| `wandb_project` | yes | W&B project to log to. |
230+
| `wandb_name` | yes | Run name. |
231+
| `wandb_api_key` | yes | W&B API key. |
232+
233+
### MLflow
234+
235+
| Key | Required | Description |
236+
|-----|----------|-------------|
237+
| `mlflow_tracking_uri` | yes | Tracking server URI. |
238+
| `mlflow_experiment_name` | yes | Experiment to log to. Created if it does not exist. |
239+
| `mlflow_run_name` | yes | Run name. |
240+
| `mlflow_tracking_token` | no | Bearer token. Omit for unauthenticated servers. |
241+
242+
```bash
243+
gym eval run \
244+
--benchmark gpqa \
245+
--model-type openai_model \
246+
+mlflow_tracking_uri=https://mlflow.example.com/ \
247+
+mlflow_experiment_name=my-experiment \
248+
+mlflow_run_name=gpqa
249+
```
250+
251+
<Note>
252+
`mlflow_tracking_uri` and `mlflow_tracking_token` are shared with the GitLab model registry used by
253+
`gym dataset upload|download`. Setting only those two does not enable the exporter — it also needs
254+
an experiment and run name.
255+
</Note>
256+
257+
### Rollout upload
258+
259+
Rollouts are uploaded to every configured backend by default. Turn this off when they are large:
260+
261+
```bash
262+
gym eval run ... +upload_rollouts=false
263+
```
264+
265+
Metrics and config are still exported. The old name `upload_rollouts_to_wandb` is deprecated and will be removed in a future release.
266+
267+
<Note>
268+
Exporting is best-effort. Errors that happen during export (e.g., unavailable server) are reported as warnings and the run continues.
269+
</Note>
270+
271+
---
272+
214273
## Command Line Usage
215274

216275
To run servers, use `gym env start`. NeMo Gym uses [Hydra](https://hydra.cc/) for configuration management.

‎nemo_gym/config_types.py‎

Lines changed: 72 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -739,11 +739,62 @@ def is_almost_server(server_type_config_dict: Any) -> bool:
739739

740740

741741
########################################
742-
# Weights and Biases
742+
# Exporter backends
743743
########################################
744744

745745

746-
class WANDBConfig(BaseModel):
746+
class ExporterConfig(BaseModel):
747+
"""Credentials and run identity for one exporter backend.
748+
749+
The exporter registry validates these against the global config to decide which backends to
750+
open, which is why they live here rather than next to the backend: checking availability must
751+
not require importing a tracking SDK.
752+
"""
753+
754+
@property
755+
def is_available(self) -> bool:
756+
"""Whether every field the backend needs to connect is set."""
757+
raise NotImplementedError
758+
759+
760+
DEPRECATED_UPLOAD_ROLLOUTS_KEY = "upload_rollouts_to_wandb"
761+
762+
763+
class UploadRolloutsConfigMixin(BaseModel):
764+
"""`upload_rollouts` plus back-compat for the W&B-specific name it replaced.
765+
766+
The flag gates rollout upload for every configured exporter, not just W&B, so the old name is
767+
accepted for one deprecation cycle and mapped onto the new field.
768+
"""
769+
770+
upload_rollouts: bool = Field(
771+
default=True,
772+
description=(
773+
"Upload the rollouts to every configured exporter. Sometimes this should be off "
774+
"because the rollouts are massive. Default: True"
775+
),
776+
)
777+
778+
@model_validator(mode="before")
779+
@classmethod
780+
def map_deprecated_upload_rollouts(cls, data):
781+
if not isinstance(data, dict) or DEPRECATED_UPLOAD_ROLLOUTS_KEY not in data:
782+
return data
783+
784+
data = dict(data)
785+
legacy = data.pop(DEPRECATED_UPLOAD_ROLLOUTS_KEY)
786+
warnings.warn(
787+
f"`{DEPRECATED_UPLOAD_ROLLOUTS_KEY}` is deprecated; use `upload_rollouts`, which "
788+
"gates rollout upload for every configured exporter.",
789+
DeprecationWarning,
790+
stacklevel=2,
791+
)
792+
# An explicit `upload_rollouts` wins, so callers can migrate without removing the old key.
793+
data.setdefault("upload_rollouts", legacy)
794+
return data
795+
796+
797+
class WANDBConfig(ExporterConfig):
747798
wandb_project: Optional[str] = None
748799
wandb_name: Optional[str] = None
749800
wandb_api_key: Optional[str] = None
@@ -754,6 +805,25 @@ def is_available(self) -> bool:
754805
return self.wandb_project and self.wandb_name and self.wandb_api_key and self.wandb_api_key != "****"
755806

756807

808+
class MLFlowConfig(ExporterConfig):
809+
"""Also used for the GitLab model registry, which needs only the URI and token."""
810+
811+
mlflow_tracking_uri: Optional[str] = None
812+
mlflow_tracking_token: Optional[str] = None
813+
mlflow_experiment_name: Optional[str] = None
814+
mlflow_run_name: Optional[str] = None
815+
816+
@property
817+
def is_available(self) -> bool:
818+
# The token is optional: unauthenticated tracking servers are possible.
819+
return (
820+
self.mlflow_tracking_uri
821+
and self.mlflow_experiment_name
822+
and self.mlflow_run_name
823+
and self.mlflow_tracking_token != "****"
824+
)
825+
826+
757827
########################################
758828
# Aggregate Metrics
759829
########################################

‎nemo_gym/exporters/__init__.py‎

Lines changed: 103 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,103 @@
1+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2+
# SPDX-License-Identifier: Apache-2.0
3+
#
4+
# Licensed under the Apache License, Version 2.0 (the "License");
5+
# you may not use this file except in compliance with the License.
6+
# You may obtain a copy of the License at
7+
#
8+
# http://www.apache.org/licenses/LICENSE-2.0
9+
#
10+
# Unless required by applicable law or agreed to in writing, software
11+
# distributed under the License is distributed on an "AS IS" BASIS,
12+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
# See the License for the specific language governing permissions and
14+
# limitations under the License.
15+
"""Run-metadata exporters.
16+
17+
`setup_exporters` opens one exporter per configured backend during global config resolution; call
18+
sites then fan out through the module-level `export_*` helpers.
19+
20+
A backend's module — and with it a multi-hundred-millisecond tracking SDK import — is loaded only
21+
once its config says it is wired up. That is why the registry pairs each backend with a config
22+
model from `config_types` rather than asking the exporter class whether it is available.
23+
"""
24+
25+
import atexit
26+
import logging
27+
from importlib import import_module
28+
from typing import Any, Optional
29+
30+
from omegaconf import DictConfig
31+
32+
from nemo_gym.config_types import ExporterConfig, MLFlowConfig, WANDBConfig
33+
from nemo_gym.exporters.base import BaseExporter
34+
35+
36+
logger = logging.getLogger(__name__)
37+
38+
# Backend name -> (config model, "module:class").
39+
EXPORTER_REGISTRY: dict[str, tuple[type[ExporterConfig], str]] = {
40+
"wandb": (WANDBConfig, "nemo_gym.exporters.wandb:WandbExporter"),
41+
"mlflow": (MLFlowConfig, "nemo_gym.exporters.mlflow:MLflowExporter"),
42+
}
43+
44+
_EXPORTERS: list[BaseExporter] = []
45+
46+
47+
def _load_exporter_class(class_path: str) -> type[BaseExporter]:
48+
module_name, class_name = class_path.split(":")
49+
return getattr(import_module(module_name), class_name)
50+
51+
52+
def get_exporters() -> list[BaseExporter]:
53+
"""The exporters opened for this process. Empty when no backend is configured."""
54+
return list(_EXPORTERS)
55+
56+
57+
def setup_exporters(global_config_dict: DictConfig) -> list[BaseExporter]:
58+
"""Open every backend that is fully configured, and log the run config to each.
59+
60+
Replaces any previously opened exporters. A backend that fails to start is skipped with a
61+
warning: telemetry must not take the run down with it.
62+
"""
63+
teardown_exporters()
64+
65+
for name, (config_model, class_path) in EXPORTER_REGISTRY.items():
66+
if not config_model.model_validate(global_config_dict).is_available:
67+
continue
68+
try:
69+
exporter = _load_exporter_class(class_path)(global_config_dict)
70+
exporter.setup()
71+
except Exception as e:
72+
logger.warning(f"Exporter {name} failed to start; continuing without it: {e}", exc_info=True)
73+
continue
74+
75+
exporter.export_config()
76+
_EXPORTERS.append(exporter)
77+
78+
# Registered here rather than at import so this lands after any hook a backend installed while
79+
# opening. atexit runs LIFO, so ours goes first: the wandb SDK closes its service in its own
80+
# hook, which would leave our teardown talking to a dead socket. Re-registering is harmless
81+
# because `teardown_exporters` is idempotent.
82+
atexit.register(teardown_exporters)
83+
return get_exporters()
84+
85+
86+
def teardown_exporters() -> None:
87+
"""Close all open exporters. Safe to call repeatedly and when none are open."""
88+
while _EXPORTERS:
89+
exporter = _EXPORTERS.pop()
90+
try:
91+
exporter.teardown()
92+
except Exception as e:
93+
logger.warning(f"Exporter {exporter.name} failed to shut down cleanly: {e}", exc_info=True)
94+
95+
96+
def export_metrics(metrics: dict[str, Any], step: Optional[int] = None) -> None:
97+
for exporter in _EXPORTERS:
98+
exporter.export_metrics(metrics, step)
99+
100+
101+
def export_rollouts(rollouts: list[dict[str, Any]]) -> None:
102+
for exporter in _EXPORTERS:
103+
exporter.export_rollouts(rollouts)

‎nemo_gym/exporters/base.py‎

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,90 @@
1+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2+
# SPDX-License-Identifier: Apache-2.0
3+
#
4+
# Licensed under the Apache License, Version 2.0 (the "License");
5+
# you may not use this file except in compliance with the License.
6+
# You may obtain a copy of the License at
7+
#
8+
# http://www.apache.org/licenses/LICENSE-2.0
9+
#
10+
# Unless required by applicable law or agreed to in writing, software
11+
# distributed under the License is distributed on an "AS IS" BASIS,
12+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
# See the License for the specific language governing permissions and
14+
# limitations under the License.
15+
import logging
16+
from abc import ABC, abstractmethod
17+
from copy import deepcopy
18+
from typing import Any, ClassVar, Optional
19+
20+
from omegaconf import DictConfig
21+
22+
from nemo_gym.secret_utils import recursively_hide_secrets
23+
24+
25+
logger = logging.getLogger(__name__)
26+
27+
28+
class BaseExporter(ABC):
29+
"""Sink for the run metadata a NeMo Gym eval produces: config, metrics, rollouts.
30+
31+
A backend (W&B, MLflow, ...) subclasses this and wraps its own run handle. Lifecycle is
32+
`setup` -> any number of `log_*` calls -> `teardown`. Whether a backend runs at all is decided
33+
by its `ExporterConfig` in the registry, not here.
34+
35+
Exporters are best-effort telemetry: a failing tracking server must not fail the eval. Call
36+
sites go through the `export_*` wrappers, which swallow and log backend exceptions. Subclasses
37+
implement the unwrapped `_*` hooks and may raise freely.
38+
"""
39+
40+
# Config key prefix and identifier used in logs, e.g. "wandb".
41+
name: ClassVar[str]
42+
43+
def __init__(self, global_config_dict: DictConfig) -> None:
44+
self.global_config_dict = global_config_dict
45+
46+
@abstractmethod
47+
def setup(self) -> None:
48+
"""Open the backing run. Called once, after config resolution."""
49+
50+
@abstractmethod
51+
def teardown(self) -> None:
52+
"""Close the backing run. Must be safe to call when `setup` failed or never ran."""
53+
54+
@abstractmethod
55+
def _log_config(self, config_dict: DictConfig) -> None:
56+
"""Record the resolved run config. Secrets are already masked by `export_config`."""
57+
58+
@abstractmethod
59+
def _log_metrics(self, metrics: dict[str, Any], step: Optional[int] = None) -> None:
60+
"""Record scalar metrics. `step` is the training step, or None for a one-shot eval."""
61+
62+
@abstractmethod
63+
def _log_rollouts(self, rollouts: list[dict[str, Any]]) -> None:
64+
"""Record rollout results as a table. Rollouts are sanitized dicts, one per rollout."""
65+
66+
def export_config(self) -> None:
67+
# `global_config_dict` holds live credentials (the backend needs its own API key to connect),
68+
# so mask a copy rather than shipping it to a tracking server as-is.
69+
config_dict_to_log = deepcopy(self.global_config_dict)
70+
recursively_hide_secrets(config_dict_to_log)
71+
self._guard("config", self._log_config, config_dict_to_log)
72+
73+
def export_metrics(self, metrics: dict[str, Any], step: Optional[int] = None) -> None:
74+
self._guard("metrics", self._log_metrics, metrics, step)
75+
76+
def export_rollouts(self, rollouts: list[dict[str, Any]]) -> None:
77+
self._guard("rollouts", self._log_rollouts, rollouts)
78+
79+
def _guard(self, what: str, fn, *args) -> None:
80+
try:
81+
fn(*args)
82+
except Exception as e:
83+
logger.warning(f"Exporter {self.name} failed to log {what}; continuing: {e}", exc_info=True)
84+
85+
def __enter__(self) -> "BaseExporter":
86+
self.setup()
87+
return self
88+
89+
def __exit__(self, exc_type, exc_value, traceback) -> None:
90+
self.teardown()

0 commit comments

Comments
 (0)