Skip to content
Merged
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
34 changes: 34 additions & 0 deletions areal/api/cli_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -1902,6 +1902,36 @@ class TensorBoardConfig:
path: str | None = None


@dataclass
class TrackioConfig:
"""Configuration for Trackio experiment tracking (Hugging Face).

Trackio is a lightweight, local-first experiment tracking library
with a wandb-compatible API. Dashboards can be viewed locally or
deployed to Hugging Face Spaces.

See: https://github.com/gradio-app/trackio
"""

mode: str = "disabled"
"""Tracking mode. One of "disabled", "online", or "local"."""
project: str | None = None
"""Project name. Defaults to experiment_name if not set."""
name: str | None = None
"""Run name. Defaults to trial_name if not set."""
space_id: str | None = None
"""HF Space ID for remote dashboard deployment (e.g. "user/my-space").
When set, metrics are also pushed to the specified Hugging Face Space."""

Comment on lines +1906 to +1925

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

To improve robustness and prevent unexpected behavior with invalid modes, it's a good practice to validate the mode field. Many other configuration dataclasses in this file (e.g., NormConfig, RecoverConfig) already use a __post_init__ method for validation. I suggest adding one here to ensure mode is one of the expected values ('disabled', 'online', or 'local').

@dataclass
class TrackioConfig:
    """Configuration for Trackio experiment tracking (Hugging Face).

    Trackio is a lightweight, local-first experiment tracking library
    with a wandb-compatible API. Dashboards can be viewed locally or
    deployed to Hugging Face Spaces.

    See: https://github.com/gradio-app/trackio
    """

    mode: str = "disabled"
    """Tracking mode. One of "disabled", "online", or "local"."""
    project: str | None = None
    """Project name. Defaults to experiment_name if not set."""
    name: str | None = None
    """Run name. Defaults to trial_name if not set."""
    space_id: str | None = None
    """HF Space ID for remote dashboard deployment (e.g. "user/my-space").
    When set, metrics are also pushed to the specified Hugging Face Space."""

    def __post_init__(self):
        """Validate Trackio configuration."""
        valid_modes = {"disabled", "online", "local"}
        if self.mode not in valid_modes:
            raise ValueError(
                f"Invalid trackio mode: '{self.mode}'. Must be one of {valid_modes}."
            )

def __post_init__(self):
"""Validate Trackio configuration."""
valid_modes = {"disabled", "online", "local"}
if self.mode not in valid_modes:
raise ValueError(
f"Invalid trackio mode: '{self.mode}'. Must be one of {valid_modes}."
)


@dataclass
class StatsLoggerConfig:
"""Configuration for experiment statistics logging and tracking services."""
Expand All @@ -1921,6 +1951,10 @@ class StatsLoggerConfig:
default_factory=TensorBoardConfig,
metadata={"help": "TensorBoard configuration. Only 'path' field required."},
)
trackio: TrackioConfig = field(
default_factory=TrackioConfig,
metadata={"help": "Trackio configuration (Hugging Face experiment tracking)."},
)


@dataclass
Expand Down
10 changes: 9 additions & 1 deletion areal/utils/logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -414,7 +414,7 @@ def setup_file_logging(


def log_swanlab_wandb_tensorboard(data, step=None, summary_writer=None):
# Logs data to SwanLab wandb TensorBoard.
# Logs data to SwanLab, wandb, TensorBoard, and Trackio.

global _LATEST_LOG_STEP
if step is None:
Expand All @@ -435,6 +435,14 @@ def log_swanlab_wandb_tensorboard(data, step=None, summary_writer=None):

wandb.log(data, step=step)

# trackio
try:
import trackio

trackio.log(data, step=step)
except (ModuleNotFoundError, ImportError):
pass

# tensorboard
if summary_writer is not None:
for key, val in data.items():
Expand Down
18 changes: 18 additions & 0 deletions areal/utils/stats_logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

import swanlab
import torch.distributed as dist
import trackio
import wandb
from tensorboardX import SummaryWriter

Expand Down Expand Up @@ -90,6 +91,19 @@ def init(self):
logdir=self.get_log_path(self.config),
mode=swanlab_config.mode,
)

# trackio init
self._trackio_enabled = False
trackio_config = self.config.trackio
if trackio_config.mode != "disabled":
trackio.init(
project=trackio_config.project or self.config.experiment_name,
name=trackio_config.name or self.config.trial_name,
config=exp_config_dict,
space_id=trackio_config.space_id,
)
self._trackio_enabled = True

# tensorboard logging
self.summary_writer = None
if self.config.tensorboard.path is not None:
Expand All @@ -111,6 +125,8 @@ def close(self):
)
wandb.finish()
swanlab.finish()
if getattr(self, "_trackio_enabled", False):
trackio.finish()
if self.summary_writer is not None:
self.summary_writer.close()

Expand All @@ -133,6 +149,8 @@ def commit(self, epoch: int, step: int, global_step: int, data: dict | list[dict
self.print_stats(item)
wandb.log(item, step=log_step + i)
swanlab.log(item, step=log_step + i)
if getattr(self, "_trackio_enabled", False):
trackio.log(item, step=log_step + i)
if self.summary_writer is not None:
for key, val in item.items():
self.summary_writer.add_scalar(f"{key}", val, log_step + i)
Expand Down
1 change: 1 addition & 0 deletions docs/generate_cli_docs.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,7 @@ def categorize_dataclasses(
"WandBConfig",
"SwanlabConfig",
"TensorBoardConfig",
"TrackioConfig",
"SaverConfig",
"EvaluatorConfig",
"RecoverConfig",
Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@ dependencies = [
# Monitoring and logging
"wandb",
"tensorboardx",
"trackio",
"colorama",
"colorlog",
"swanboard==0.1.9b1",
Expand Down
195 changes: 195 additions & 0 deletions tests/test_trackio_backend.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,195 @@
"""Tests for Trackio experiment tracking backend integration."""

from dataclasses import fields
from unittest.mock import MagicMock, patch

from areal.api.cli_args import (
StatsLoggerConfig,
TrackioConfig,
)


class TestTrackioConfig:
"""Tests for TrackioConfig dataclass."""

def test_default_mode_is_disabled(self):
"""TrackioConfig should default to disabled mode."""
config = TrackioConfig()
assert config.mode == "disabled"

def test_default_optional_fields_are_none(self):
"""Optional fields should default to None."""
config = TrackioConfig()
assert config.project is None
assert config.name is None
assert config.space_id is None

def test_custom_values(self):
"""TrackioConfig should accept custom values."""
config = TrackioConfig(
mode="online",
project="my-project",
name="my-run",
space_id="user/my-space",
)
assert config.mode == "online"
assert config.project == "my-project"
assert config.name == "my-run"
assert config.space_id == "user/my-space"

def test_invalid_mode_raises_error(self):
"""TrackioConfig should reject invalid mode values."""
import pytest

with pytest.raises(ValueError, match="Invalid trackio mode"):
TrackioConfig(mode="invalid")

def test_all_valid_modes_accepted(self):
"""TrackioConfig should accept all valid mode values."""
for mode in ("disabled", "online", "local"):
config = TrackioConfig(mode=mode)
assert config.mode == mode


class TestStatsLoggerConfigTrackio:
"""Tests for Trackio field in StatsLoggerConfig."""

def test_trackio_field_exists(self):
"""StatsLoggerConfig should have a trackio field."""
field_names = [f.name for f in fields(StatsLoggerConfig)]
assert "trackio" in field_names

def test_trackio_field_default_is_disabled(self):
"""StatsLoggerConfig.trackio should default to disabled TrackioConfig."""
config = StatsLoggerConfig(
experiment_name="test_exp",
trial_name="trial_0",
fileroot="/tmp/test",
)
assert isinstance(config.trackio, TrackioConfig)
assert config.trackio.mode == "disabled"


def _make_test_config(trackio_config=None):
"""Create a minimal BaseExperimentConfig for testing StatsLogger."""
from areal.api.cli_args import BaseExperimentConfig

config = BaseExperimentConfig(
experiment_name="test_exp",
trial_name="trial_0",
total_train_epochs=1,
)
config.stats_logger.experiment_name = "test_exp"
config.stats_logger.trial_name = "trial_0"
config.stats_logger.fileroot = "/tmp/test"
if trackio_config is not None:
config.stats_logger.trackio = trackio_config
return config


def _make_ft_spec():
"""Create a mock FinetuneSpec for testing."""
from areal.api import FinetuneSpec

ft_spec = MagicMock(spec=FinetuneSpec)
ft_spec.total_train_epochs = 1
ft_spec.steps_per_epoch = 10
ft_spec.total_train_steps = 10
return ft_spec


class TestStatsLoggerTrackioIntegration:
"""Tests for Trackio integration in StatsLogger (mocked)."""

@patch("areal.utils.stats_logger.trackio")
@patch("areal.utils.stats_logger.wandb")
@patch("areal.utils.stats_logger.swanlab")
@patch("areal.utils.stats_logger.dist")
def test_trackio_init_called_when_enabled(
self, mock_dist, mock_swanlab, mock_wandb, mock_trackio
):
"""trackio.init() should be called when mode is not disabled."""
mock_dist.is_initialized.return_value = False

from areal.utils.stats_logger import StatsLogger

config = _make_test_config(TrackioConfig(mode="online"))
logger = StatsLogger(config, _make_ft_spec())
mock_trackio.init.assert_called_once()
assert logger._trackio_enabled is True

@patch("areal.utils.stats_logger.trackio")
@patch("areal.utils.stats_logger.wandb")
@patch("areal.utils.stats_logger.swanlab")
@patch("areal.utils.stats_logger.dist")
def test_trackio_not_init_when_disabled(
self, mock_dist, mock_swanlab, mock_wandb, mock_trackio
):
"""trackio.init() should NOT be called when mode is disabled."""
mock_dist.is_initialized.return_value = False

from areal.utils.stats_logger import StatsLogger

config = _make_test_config() # trackio defaults to disabled
logger = StatsLogger(config, _make_ft_spec())
mock_trackio.init.assert_not_called()
assert logger._trackio_enabled is False

@patch("areal.utils.stats_logger.trackio")
@patch("areal.utils.stats_logger.wandb")
@patch("areal.utils.stats_logger.swanlab")
@patch("areal.utils.stats_logger.dist")
def test_trackio_log_called_on_commit(
self, mock_dist, mock_swanlab, mock_wandb, mock_trackio
):
"""trackio.log() should be called during commit when enabled."""
mock_dist.is_initialized.return_value = False

from areal.utils.stats_logger import StatsLogger

config = _make_test_config(TrackioConfig(mode="online"))
logger = StatsLogger(config, _make_ft_spec())
mock_trackio.log.reset_mock()

data = {"loss/avg": 0.5, "reward/avg": 1.0}
logger.commit(epoch=0, step=0, global_step=0, data=data)
mock_trackio.log.assert_called_once_with(data, step=0)

@patch("areal.utils.stats_logger.trackio")
@patch("areal.utils.stats_logger.wandb")
@patch("areal.utils.stats_logger.swanlab")
@patch("areal.utils.stats_logger.dist")
def test_trackio_finish_called_on_close(
self, mock_dist, mock_swanlab, mock_wandb, mock_trackio
):
"""trackio.finish() should be called during close when enabled."""
mock_dist.is_initialized.return_value = False

from areal.utils.stats_logger import StatsLogger

config = _make_test_config(TrackioConfig(mode="online"))
logger = StatsLogger(config, _make_ft_spec())
mock_trackio.finish.reset_mock()

logger.close()
mock_trackio.finish.assert_called_once()

@patch("areal.utils.stats_logger.trackio")
@patch("areal.utils.stats_logger.wandb")
@patch("areal.utils.stats_logger.swanlab")
@patch("areal.utils.stats_logger.dist")
def test_trackio_not_logged_when_disabled(
self, mock_dist, mock_swanlab, mock_wandb, mock_trackio
):
"""trackio.log() should NOT be called during commit when disabled."""
mock_dist.is_initialized.return_value = False

from areal.utils.stats_logger import StatsLogger

config = _make_test_config() # trackio defaults to disabled
logger = StatsLogger(config, _make_ft_spec())
mock_trackio.log.reset_mock()

data = {"loss/avg": 0.5}
logger.commit(epoch=0, step=0, global_step=0, data=data)
mock_trackio.log.assert_not_called()
Loading