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
134 changes: 132 additions & 2 deletions tests/unit_tests/observability/test_structured_logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
set_step,
)
from torchtitan.observability.structured_logger.structured_logging import (
_get_structured_logger_init_args,
_structured_logger,
event_extra,
ExtraFields,
Expand Down Expand Up @@ -73,16 +74,43 @@ def structured_logger_fixture():
import torchtitan.observability.structured_logger.structured_logging as sl_mod

tl = _structured_logger
orig = (tl.handlers[:], tl.level, tl.propagate)
root_logger = logging.getLogger()
orig = (
tl.handlers[:],
tl.level,
tl.propagate,
sl_mod._disabled,
sl_mod._structured_logger_init_args,
root_logger.handlers[:],
)
# Reset the module-level init sentinel so init_structured_logger re-runs
# for each test (otherwise the second call short-circuits as "already
# initialized").
sl_mod._is_initialized = False
sl_mod._disabled = False
sl_mod._structured_logger_init_args = None
yield tl
tl.handlers, tl.level, tl.propagate = orig
(
tl.handlers,
tl.level,
tl.propagate,
sl_mod._disabled,
init_args,
root_logger.handlers,
) = orig
sl_mod._structured_logger_init_args = init_args
sl_mod._is_initialized = False


@pytest.fixture
def external_logger():
source_logger = logging.getLogger("external_library")
orig = (source_logger.handlers[:], source_logger.level, source_logger.propagate)
source_logger.handlers = []
yield source_logger
source_logger.handlers, source_logger.level, source_logger.propagate = orig


# ---------------------------------------------------------------------------
# Step context tests (hybrid ContextVar)
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -506,6 +534,45 @@ def test_has_time_us(self):
assert "time_us" in parsed
assert isinstance(parsed["time_us"], int)

def test_checkpoint_context_fields(self):
fmt = TraceJsonlFormatter(rank=0, source="test")
record = logging.LogRecord(
name="test",
level=logging.INFO,
pathname="test.py",
lineno=1,
msg="test",
args=None,
exc_info=None,
)
for key, value in event_extra("log_metric").items():
setattr(record, key, value)
record.context = []
record.measured_from_start_time_ms = 123456

parsed = json.loads(fmt.format(record))

assert parsed["context"] == []
assert parsed["measured_from_start_time_ms"] == 123456

def test_omits_missing_context(self):
fmt = TraceJsonlFormatter(rank=0, source="test")
record = logging.LogRecord(
name="test",
level=logging.INFO,
pathname="test.py",
lineno=1,
msg="test",
args=None,
exc_info=None,
)
for key, value in event_extra("log_metric").items():
setattr(record, key, value)

parsed = json.loads(fmt.format(record))

assert "context" not in parsed


# ---------------------------------------------------------------------------
# TraceEventsOnlyFilter
Expand Down Expand Up @@ -607,6 +674,15 @@ def test_second_call_is_noop(self, tmp_path, structured_logger_fixture):

assert len(structured_logger_fixture.handlers) == handler_count

def test_records_resolved_init_args(self, tmp_path, structured_logger_fixture):
init_structured_logger(rank=17, source="rl_trainer", output_dir=str(tmp_path))

assert _get_structured_logger_init_args() == (
"rl_trainer",
str(tmp_path),
17,
)


class TestFactoryMechanism:
def test_default_creates_jsonl(self, tmp_path, structured_logger_fixture):
Expand Down Expand Up @@ -646,6 +722,60 @@ def fake_factory(*, structured_logger, rank, source, output_dir, **kw):
)


# ---------------------------------------------------------------------------
# External structured logging
# ---------------------------------------------------------------------------


class TestExternalStructuredLogging:
def test_init_forwards_only_structured_records_from_other_loggers(
self, tmp_path, structured_logger_fixture, external_logger
):
init_structured_logger(rank=0, source="trainer", output_dir=str(tmp_path))
external_logger.setLevel(logging.INFO)

logging.getLogger("external_library.worker").info(
"external metric",
extra={
"log_type": "event",
"log_type_name": "log_metric",
"event_name": "train.step.e2e.latency_ms",
"step": 7,
"value": 12.5,
"context": ["source:test"],
},
)
logging.getLogger("external_library.worker").info("plain text")

for handler in structured_logger_fixture.handlers:
handler.flush()
trace_dir = os.path.join(str(tmp_path), "structured_logs")
jsonl_files = [f for f in os.listdir(trace_dir) if f.endswith(".jsonl")]
with open(os.path.join(trace_dir, jsonl_files[0])) as f:
lines = [json.loads(line) for line in f if line.strip()]

assert len(lines) == 1
assert lines[0]["logger_name"] == "external_library.worker"
assert lines[0]["log_type_name"] == "log_metric"
assert lines[0]["event_name"] == "train.step.e2e.latency_ms"
assert lines[0]["step"] == 7
assert lines[0]["value"] == 12.5
assert lines[0]["context"] == ["source:test"]
assert structured_logger_fixture.propagate is False

def test_second_init_restores_a_removed_root_forwarder(
self, tmp_path, structured_logger_fixture
):
root_logger = logging.getLogger()
init_structured_logger(rank=0, source="trainer", output_dir=str(tmp_path))
forwarder = root_logger.handlers[-1]
root_logger.handlers = []

init_structured_logger(rank=0, source="trainer", output_dir=str(tmp_path))

assert root_logger.handlers == [forwarder]


# ---------------------------------------------------------------------------
# No-op flag
# ---------------------------------------------------------------------------
Expand Down
104 changes: 104 additions & 0 deletions tests/unit_tests/test_torch_checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

import dataclasses
import json
import logging
import queue
import unittest
from concurrent.futures import Future
Expand All @@ -27,6 +28,7 @@
SyncCheckpointSaverConfig,
)
from torch_checkpointing.default_resharder import DefaultResharder
from torch_checkpointing.logging_utils import checkpoint_logging_context
from torchtitan.components.checkpointer import (
BaseCheckpointManager,
CheckpointManager,
Expand Down Expand Up @@ -548,3 +550,105 @@ def test_last_step_uses_synchronous_manager_and_model_only_payload(self) -> None
sync_config = build.call_args.args[0]
self.assertIsNone(sync_config.pre_finalize_callback)
manager.close()

def test_save_stamps_the_step_on_backend_events(self) -> None:
# The backend reads this context when it builds its own events and
# exports it to the async save subprocess. Without it every forwarded
# backend metric carries step=None, which makes them hard to line up
# against the training step they belong to.
config = TorchCheckpointingManager.Config(
enable=True,
interval=1,
keep_latest_k=0,
initial_load_model_only=False,
)
manager, backend_manager = self._build_manager(config)
self.addCleanup(checkpoint_logging_context.import_context, {})

self.assertTrue(manager.save(curr_step=7))

self.assertEqual(7, checkpoint_logging_context.get("step"))
backend_manager.save_result.set_result(None)
manager.close()

def test_subprocess_logging_initializes_and_delegates(self) -> None:
calls = []
init_fn = mock.Mock(side_effect=lambda *_args: calls.append("existing"))
with mock.patch.object(
manager_module.sl,
"init_structured_logger",
side_effect=lambda **_kwargs: calls.append("structured"),
) as init_structured_logger:
manager_module._init_subprocess_logging(
("rl_trainer", "/tmp/output", 17),
init_fn,
("argument",),
)

init_structured_logger.assert_called_once_with(
source="rl_trainer",
output_dir="/tmp/output",
rank=17,
)
init_fn.assert_called_once_with("argument")
self.assertEqual(["existing", "structured"], calls)

def _init_subprocess_logging(self) -> None:
with mock.patch.object(manager_module.sl, "init_structured_logger"):
manager_module._init_subprocess_logging(
("training", "/tmp/output", 0), None, ()
)

def test_subprocess_logging_only_overrides_suppressed_inherited_level(self) -> None:
root_logger = logging.getLogger()
backend_logger = logging.getLogger(manager_module._BACKEND_LOGGER_NAME)
self.addCleanup(root_logger.setLevel, root_logger.level)
self.addCleanup(backend_logger.setLevel, backend_logger.level)
for root_level, backend_level, expected_level in (
(logging.WARNING, logging.NOTSET, logging.INFO),
(logging.WARNING, logging.DEBUG, logging.DEBUG),
(logging.WARNING, logging.WARNING, logging.WARNING),
(logging.DEBUG, logging.NOTSET, logging.NOTSET),
):
with self.subTest(root_level=root_level, backend_level=backend_level):
root_logger.setLevel(root_level)
backend_logger.setLevel(backend_level)

self._init_subprocess_logging()

self.assertEqual(expected_level, backend_logger.level)

def test_async_manager_composes_subprocess_logging_initializer(self) -> None:
original_init_fn = mock.Mock()
config = TorchCheckpointingManager.Config(
enable=True,
keep_latest_k=0,
initial_load_model_only=False,
)
backend_config = _default_backend_config()
backend_config.subprocess_init_fn = original_init_fn
backend_config.subprocess_init_args = ("argument",)

with mock.patch.object(
manager_module,
"_get_structured_logger_init_args",
return_value=("rl_trainer", "/tmp/output", 17),
):
manager, _ = self._build_manager(
config,
backend_config=backend_config,
)

self.assertIs(
manager._manager_config.subprocess_init_fn,
manager_module._init_subprocess_logging,
)
self.assertEqual(
(
("rl_trainer", "/tmp/output", 17),
original_init_fn,
("argument",),
),
manager._manager_config.subprocess_init_args,
)
manager.close()
Loading
Loading