From 3ecc9d4ed860613948c02dc689e6b47be17611e3 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Wed, 9 Sep 2026 09:28:38 -0400 Subject: [PATCH 1/3] fix(train): defer W&B authentication to SDK --- skyrl/train/utils/utils.py | 4 ---- skyrl/utils/log.py | 3 --- tests/utils/test_log.py | 14 ++++++++++++++ 3 files changed, 14 insertions(+), 7 deletions(-) create mode 100644 tests/utils/test_log.py diff --git a/skyrl/train/utils/utils.py b/skyrl/train/utils/utils.py index 0e4fdc1843..ef3c712ac5 100644 --- a/skyrl/train/utils/utils.py +++ b/skyrl/train/utils/utils.py @@ -536,10 +536,6 @@ def validate_generator_cfg(cfg: SkyRLTrainConfig): "for multi-turn generation" ) - # TODO(tgriggs): use a more modular config validation - if cfg.trainer.logger == "wandb": - assert os.environ.get("WANDB_API_KEY"), "`WANDB_API_KEY` is required for `wandb` logger" - if cfg.generator.sampling_params.logprobs is not None: assert isinstance(cfg.generator.sampling_params.logprobs, int) if cfg.generator.sampling_params.logprobs > 1: diff --git a/skyrl/utils/log.py b/skyrl/utils/log.py index dad5cde16e..d2e36c6461 100644 --- a/skyrl/utils/log.py +++ b/skyrl/utils/log.py @@ -1,5 +1,4 @@ import logging -import os from enum import Enum from pathlib import Path from typing import Any @@ -113,8 +112,6 @@ def __init__(self, config: dict[str, Any], **kwargs): super().__init__(config, **kwargs) if wandb is None: raise RuntimeError("wandb not installed") - if not os.environ.get("WANDB_API_KEY"): - raise ValueError("WANDB_API_KEY environment variable not set") self.run = wandb.init(config=config, **kwargs) # type: ignore[union-attr] def log(self, metrics: dict[str, Any], step: int | None = None) -> None: diff --git a/tests/utils/test_log.py b/tests/utils/test_log.py new file mode 100644 index 0000000000..9a43d1d85d --- /dev/null +++ b/tests/utils/test_log.py @@ -0,0 +1,14 @@ +from unittest.mock import MagicMock + +from skyrl.utils import log + + +def test_wandb_tracker_delegates_authentication_to_wandb(monkeypatch): + monkeypatch.delenv("WANDB_API_KEY", raising=False) + wandb = MagicMock() + monkeypatch.setattr(log, "wandb", wandb) + + tracker = log.WandbTracker(config={"model": "test"}, project="project") + + wandb.init.assert_called_once_with(config={"model": "test"}, project="project") + assert tracker.run is wandb.init.return_value From c2fa48bc33811b86e017c0e2a2095387a46a6da9 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Wed, 9 Sep 2026 09:29:21 -0400 Subject: [PATCH 2/3] [test][train] Cover SDK-owned W&B auth --- tests/train/test_config.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/tests/train/test_config.py b/tests/train/test_config.py index 4f11b34edf..00328aef4c 100644 --- a/tests/train/test_config.py +++ b/tests/train/test_config.py @@ -25,6 +25,7 @@ from skyrl.train.utils.utils import ( prepare_runtime_environment, validate_cfg, + validate_generator_cfg, validate_inference_engine_cfg, ) from tests.train.util import example_dummy_config @@ -46,6 +47,14 @@ def _make_validated_test_config(): return cfg +def test_validate_generator_cfg_defers_wandb_authentication(monkeypatch): + cfg = _make_validated_test_config() + cfg.trainer.logger = "wandb" + monkeypatch.delenv("WANDB_API_KEY", raising=False) + + validate_generator_cfg(cfg) + + # Helper dataclasses for testing @dataclass class _SimpleConfig(BaseConfig): From 984720f2ef7bc32d9f8dc3f82372290e94bfebe4 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Thu, 10 Sep 2026 00:29:37 -0400 Subject: [PATCH 3/3] [fix][train] Forward W&B runtime settings --- skyrl/train/utils/utils.py | 7 ++++--- tests/train/test_config.py | 26 ++++++++++++++++++++++++++ 2 files changed, 30 insertions(+), 3 deletions(-) diff --git a/skyrl/train/utils/utils.py b/skyrl/train/utils/utils.py index ef3c712ac5..9e6c44a982 100644 --- a/skyrl/train/utils/utils.py +++ b/skyrl/train/utils/utils.py @@ -882,9 +882,10 @@ def prepare_runtime_environment(cfg: SkyRLTrainConfig) -> dict[str, str]: # TODO: this can be removed if we standardize on env files. # But it's helpful for a quickstart - if os.environ.get("WANDB_API_KEY"): - logger.info("Exporting wandb api key to ray runtime env") - env_vars["WANDB_API_KEY"] = os.environ["WANDB_API_KEY"] + for var_name in ("WANDB_API_KEY", "WANDB_MODE", "WANDB_BASE_URL"): + if value := os.environ.get(var_name): + logger.info(f"Exporting `{var_name}` to ray runtime env") + env_vars[var_name] = value if os.environ.get("MLFLOW_TRACKING_URI"): logger.info("Exporting mlflow tracking uri to ray runtime env") diff --git a/tests/train/test_config.py b/tests/train/test_config.py index 00328aef4c..c97d940c81 100644 --- a/tests/train/test_config.py +++ b/tests/train/test_config.py @@ -183,6 +183,32 @@ def test_runtime_env_forwards_te_block_scale_mode(monkeypatch): assert env_vars["NVTE_FP8_BLOCK_SCALING_FP32_SCALES"] == "1" +def test_runtime_env_forwards_wandb_environment(monkeypatch): + wandb_environment = { + "WANDB_API_KEY": "test-api-key", + "WANDB_MODE": "offline", + "WANDB_BASE_URL": "https://wandb.example.com", + } + for var_name, value in wandb_environment.items(): + monkeypatch.setenv(var_name, value) + monkeypatch.setattr(train_utils, "peer_access_supported", lambda **_kwargs: True) + + env_vars = prepare_runtime_environment(example_dummy_config()) + + assert {var_name: env_vars[var_name] for var_name in wandb_environment} == wandb_environment + + +def test_runtime_env_omits_unset_wandb_environment(monkeypatch): + wandb_environment = {"WANDB_API_KEY", "WANDB_MODE", "WANDB_BASE_URL"} + for var_name in wandb_environment: + monkeypatch.delenv(var_name, raising=False) + monkeypatch.setattr(train_utils, "peer_access_supported", lambda **_kwargs: True) + + env_vars = prepare_runtime_environment(example_dummy_config()) + + assert wandb_environment.isdisjoint(env_vars) + + def test_runtime_env_supports_fsdp_without_megatron_configs(monkeypatch): monkeypatch.delenv("NVTE_FP8_BLOCK_SCALING_FP32_SCALES", raising=False) monkeypatch.delenv("NVTE_FP8_BLOCK_AMAX_EPSILON", raising=False)