Skip to content
Open
21 changes: 10 additions & 11 deletions autotest/utils/check_metric.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,8 @@


def _normalize_sft_metric_cfg(metric: str, value) -> tuple[float, int | None, str]:
"""Accept ``metric: threshold`` or ``metric: {threshold, aggregate, method}``.
"""Accept ``metric: threshold`` or ``metric: {threshold, aggregate,
method}``.

Percentile aggregation is opt-in only (explicit ``aggregate`` in config).
Comparison method defaults to ``relative``; token counts default to ``absolute``.
Expand Down Expand Up @@ -127,7 +128,8 @@ def _align_first_phase_steps(phase, base_steps, cur_steps, cur_metrics, kind="SF


def _resolve_first_phase_baseline_steps(resume_base_path: str) -> int | None:
"""Return step count of companion ``tracker.jsonl`` for a resume baseline."""
"""Return step count of companion ``tracker.jsonl`` for a resume
baseline."""
if not resume_base_path.endswith("tracker-resume.jsonl"):
return None
first_path = resume_base_path[: -len("tracker-resume.jsonl")] + "tracker.jsonl"
Expand Down Expand Up @@ -284,7 +286,8 @@ def _should_run_memory_gradient_check(
drift_threshold: float,
max_rel_drift: float,
) -> bool:
"""Run leak heuristic only when baseline drift is loose or current swing grew."""
"""Run leak heuristic only when baseline drift is loose or current swing
grew."""
tight_vs_baseline = max_rel_drift < drift_threshold * MEMORY_GRADIENT_BASELINE_DRIFT_SKIP_RATIO
base_range = float(max(base_vals) - min(base_vals))
cur_range = float(max(cur_vals) - min(cur_vals))
Expand Down Expand Up @@ -370,9 +373,8 @@ def _find_nonfinite(values: list[float]) -> int | None:
def _annotate_mtp_only_failure(fail_metric: dict[str, str], metric_list: list[str]) -> None:
"""Clarify MTP-only failures when llm/local anchors still pass.

MTP CE is expected to be much noisier than next-token llm/local loss. A lone
MTP threshold miss with healthy llm/local is usually pack/noise, not a
train-infer bug — unless values are non-finite (already hard-failed above).
MTP CE is expected to be much noisier than next-token llm/local loss. A lone MTP threshold miss with healthy
llm/local is usually pack/noise, not a train-infer bug — unless values are non-finite (already hard-failed above).
"""
mtp_key = "loss/reduced_mtp_loss"
if mtp_key not in fail_metric:
Expand Down Expand Up @@ -428,9 +430,7 @@ def check_result(case_name, base_path, cur_path, check_metric, phase=None):
if metric in SFT_FINITE_LOSS_METRICS:
bad_idx = _find_nonfinite(cur_metrics[metric])
if bad_idx is not None:
fail_metric[metric] = (
f"{metric} is non-finite at step {bad_idx}: {cur_metrics[metric][bad_idx]!r}"
)
fail_metric[metric] = f"{metric} is non-finite at step {bad_idx}: {cur_metrics[metric][bad_idx]!r}"
continue
bad_idx = _find_nonfinite(base_metrics[metric])
if bad_idx is not None:
Expand Down Expand Up @@ -665,8 +665,7 @@ def check_rl_result(case_name, base_path, cur_path, assert_info, phase=None):
if not passed:
if method == "value":
fail_metric[metric] = (
f"{metric} value {cur_val:.6f} does not satisfy {operator} {threshold} "
f"at step {report_idx}"
f"{metric} value {cur_val:.6f} does not satisfy {operator} {threshold} at step {report_idx}"
)
else:
fail_metric[metric] = (
Expand Down
41 changes: 38 additions & 3 deletions recipe/trace/viewer/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
import argparse
import http.server
import json
import logging
import os
import threading
import time
Expand All @@ -28,6 +29,9 @@
from recipe.trace.viewer.render import render_rollout_trace_html, write_rollout_trace_html


logger = logging.getLogger(__name__)


_JAEGER_PROXY_PREFIX = "/jaeger"
JAEGER_DEFAULT_QUERY_URL = "http://127.0.0.1:16686"
_PROXY_TIMEOUT_S = 10.0
Expand Down Expand Up @@ -122,6 +126,28 @@ def _train_step_cache_key(train_step: str | int | None) -> str:
return text or "latest"


def _empty_payload(
*,
trace_jsonl_path: Path | str,
service_name: str | None,
run_id: str | None,
) -> dict[str, Any]:
return {
"title": "XTuner Rollout Trace Viewer",
"generated_at_s": None,
"source": "trace_jsonl",
"jaeger_query_url": None,
"jaeger_link_url": None,
"service_name": service_name,
"run_id": run_id,
"trace_jsonl_path": os.fspath(Path(trace_jsonl_path).expanduser()),
"available_train_steps": [],
"samples": [],
"sample_count": 0,
"unavailable": "trace JSONL not readable yet",
}


def fetch_rollout_view_payload_from_trace_jsonl(
trace_jsonl_path: Path | str,
*,
Expand Down Expand Up @@ -181,7 +207,6 @@ def current_source_signature() -> Any:
load_base_payload,
source_signature=current_source_signature,
)
payload_cache.get(train_step)

Comment thread
matrix72c marked this conversation as resolved.
class Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self) -> None:
Expand Down Expand Up @@ -211,7 +236,15 @@ def _query_train_step(self, query: str) -> str | int | None:
return values[-1]

def _payload(self, selected_train_step: str | int | None) -> dict[str, Any]:
return payload_cache.get(selected_train_step)
try:
return payload_cache.get(selected_train_step)
except (OSError, ValueError) as exc:
logger.warning("Trace payload unavailable for %s: %s", self.path, exc)
return _empty_payload(
trace_jsonl_path=trace_jsonl_path,
service_name=service_name,
run_id=run_id,
)

def _send_json(self, payload: dict[str, Any]) -> None:
self._send_bytes(json.dumps(payload, ensure_ascii=False).encode("utf-8"), "application/json")
Expand Down Expand Up @@ -302,7 +335,9 @@ def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=0)
parser.add_argument("--output", type=Path, default=None)
parser.add_argument("--train-step", default="latest", help="Initial train step to render: latest, all, or a step value.")
parser.add_argument(
"--train-step", default="latest", help="Initial train step to render: latest, all, or a step value."
)
return parser.parse_args(argv)


Expand Down
227 changes: 227 additions & 0 deletions tests/rl/test_trace.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,9 @@
import sys
import unittest
from pathlib import Path
from tempfile import TemporaryDirectory
from unittest import mock
from urllib.request import urlopen


def _run_trace_utils(repo_root: Path, command: str) -> dict:
Expand All @@ -22,6 +25,230 @@ def _run_trace_utils(repo_root: Path, command: str) -> dict:


class TestTrace(unittest.TestCase):
def test_external_collector_skips_local_collector_and_propagates_endpoint(self):
from xtuner.v1.rl.trace import runtime as trace_runtime

class Provider:
def shutdown(self):
return None

with TemporaryDirectory() as temp_dir:
config = trace_runtime.TraceConfig(
enabled=True,
output_dir=temp_dir,
external_otlp_endpoint="http://otel-collector.namespace.svc:4317",
)
handle = trace_runtime._build_trace_runtime_handle(config)

self.assertFalse(handle.start_local_collector)
self.assertIsNone(handle.collector_port)
self.assertIsNone(handle.runtime.trace_jsonl_path)
self.assertEqual(
handle.env_vars["OTEL_EXPORTER_OTLP_ENDPOINT"],
"http://otel-collector.namespace.svc:4317",
)
self.assertNotIn("XTUNER_OTEL_JSONL_PATH", handle.env_vars)

with (
mock.patch.object(trace_runtime._OTelCollector, "start") as start_collector,
mock.patch.object(trace_runtime, "_configure_tracer_provider", return_value=Provider()),
):
handle.start()
handle.close()
start_collector.assert_not_called()
trace_runtime.clear_trace_env()

def test_local_collector_remains_the_default(self):
from xtuner.v1.rl.trace import runtime as trace_runtime

with TemporaryDirectory() as temp_dir:
with (
mock.patch.object(trace_runtime, "find_free_ports", return_value=[4317]),
mock.patch.object(trace_runtime, "_local_advertised_host", return_value="10.0.0.1"),
):
handle = trace_runtime._build_trace_runtime_handle(
trace_runtime.TraceConfig(enabled=True, output_dir=temp_dir)
)

self.assertTrue(handle.start_local_collector)
self.assertIsNotNone(handle.collector_port)
self.assertIsNotNone(handle.runtime.trace_jsonl_path)
self.assertTrue(handle.runtime.trace_jsonl_path.is_file())
self.assertEqual(handle.endpoint, "http://10.0.0.1:4317")

def test_external_trace_jsonl_is_owned_by_collector_and_not_propagated(self):
from xtuner.v1.rl.trace import runtime as trace_runtime

with TemporaryDirectory() as temp_dir:
trace_path = Path(temp_dir) / "shared" / "traces.jsonl"
handle = trace_runtime._build_trace_runtime_handle(
trace_runtime.TraceConfig(
enabled=True,
output_dir=Path(temp_dir) / "runs",
external_otlp_endpoint="http://otel-collector.namespace.svc:4317",
external_trace_jsonl_path=trace_path,
xtuner_viewer_enabled=True,
)
)

self.assertEqual(handle.runtime.trace_jsonl_path, trace_path)
self.assertNotIn("XTUNER_OTEL_JSONL_PATH", handle.env_vars)
self.assertFalse(trace_path.parent.exists())
self.assertFalse(trace_path.exists())

def test_external_viewer_lazily_loads_trace_jsonl(self):
from recipe.trace.viewer.server import start_rollout_trace_viewer

with TemporaryDirectory() as temp_dir:
trace_path = Path(temp_dir) / "shared" / "traces.jsonl"
viewer = start_rollout_trace_viewer(
None,
service_name="xtuner-rollout",
run_id="run-1",
trace_jsonl_path=trace_path,
host="127.0.0.1",
port=0,
train_step="all",
)
try:
self.assertTrue(viewer.thread.is_alive())
self.assertFalse(trace_path.exists())
Comment thread
matrix72c marked this conversation as resolved.

with urlopen(f"{viewer.url}/api/trace?train_step=all", timeout=2) as response:
self.assertEqual(response.status, 200)
payload = json.load(response)
self.assertEqual(payload["sample_count"], 0)
self.assertEqual(payload["samples"], [])
self.assertIn("unavailable", payload)

trace_path.parent.mkdir(parents=True)
trace_path.write_text(
json.dumps(
{
"traceID": "trace-1",
"processes": {
"p1": {
"serviceName": "xtuner-rollout",
"tags": [{"key": "run.id", "value": "run-1"}],
}
},
"spans": [
{
"traceID": "trace-1",
"spanID": "span-1",
"operationName": "rollout.generate",
"processID": "p1",
"startTime": 1_000,
"duration": 2_000,
"tags": [{"key": "xtuner.rollout_id", "value": "rollout-1"}],
}
],
}
)
+ "\n",
encoding="utf-8",
)

with urlopen(f"{viewer.url}/api/trace?train_step=all", timeout=2) as response:
payload = json.load(response)

self.assertEqual(payload["sample_count"], 1)
self.assertEqual(payload["samples"][0]["rollout_id"], "rollout-1")
finally:
viewer.close()

def test_external_viewer_process_starts_before_trace_jsonl_exists(self):
from xtuner.v1.rl.trace import runtime as trace_runtime

class Provider:
def shutdown(self):
return None

with TemporaryDirectory() as temp_dir:
trace_path = Path(temp_dir) / "shared" / "traces.jsonl"
handle = trace_runtime._build_trace_runtime_handle(
trace_runtime.TraceConfig(
enabled=True,
output_dir=Path(temp_dir) / "runs",
external_otlp_endpoint="http://otel-collector.namespace.svc:4317",
external_trace_jsonl_path=trace_path,
xtuner_viewer_enabled=True,
xtuner_viewer_port=0,
)
)
try:
with mock.patch.object(trace_runtime, "_configure_tracer_provider", return_value=Provider()):
handle.start()
self.assertIsNotNone(handle.xtuner_viewer_process)
self.assertIsNone(handle.xtuner_viewer_process.poll())
self.assertFalse(trace_path.exists())
finally:
handle.close()
trace_runtime.clear_trace_env()

def test_ray_child_inherits_external_endpoint_without_trace_jsonl(self):
from xtuner.v1.rl.trace import runtime as trace_runtime

class Provider:
def shutdown(self):
return None

with TemporaryDirectory() as temp_dir:
trace_path = Path(temp_dir) / "shared" / "traces.jsonl"
driver_handle = trace_runtime._build_trace_runtime_handle(
trace_runtime.TraceConfig(
enabled=True,
output_dir=Path(temp_dir) / "runs",
external_otlp_endpoint="http://otel-collector.namespace.svc:4317",
external_trace_jsonl_path=trace_path,
)
)

with (
mock.patch.object(trace_runtime, "_RUNTIME", None),
mock.patch.object(trace_runtime, "register_atexit_once"),
mock.patch.object(trace_runtime._OTelCollector, "start") as start_collector,
mock.patch.object(trace_runtime, "_configure_tracer_provider", return_value=Provider()) as configure,
mock.patch.dict(os.environ, driver_handle.env_vars, clear=True),
):
self.assertTrue(trace_runtime.ensure_trace_runtime_from_env())
runtime = trace_runtime.current_trace_runtime()
self.assertIsNotNone(runtime)
self.assertEqual(runtime.mode, "inherited")
self.assertIsNone(runtime.trace_jsonl_path)
self.assertNotIn("XTUNER_OTEL_JSONL_PATH", trace_runtime.get_trace_env_vars())
configure.assert_called_once_with(
service_name="xtuner-rollout",
run_id=driver_handle.runtime.run_id,
endpoint="http://otel-collector.namespace.svc:4317",
protocol="grpc",
)
start_collector.assert_not_called()
trace_runtime.close_trace()

def test_external_viewer_requires_shared_trace_jsonl(self):
from pydantic import ValidationError

from xtuner.v1.rl.trace import runtime as trace_runtime

with self.assertRaisesRegex(ValidationError, "external_trace_jsonl_path"):
trace_runtime.TraceConfig(
enabled=True,
external_otlp_endpoint="http://otel-collector.namespace.svc:4317",
xtuner_viewer_enabled=True,
)

def test_external_trace_jsonl_requires_external_endpoint(self):
from pydantic import ValidationError

from xtuner.v1.rl.trace import runtime as trace_runtime

with self.assertRaisesRegex(ValidationError, "external_otlp_endpoint"):
trace_runtime.TraceConfig(
enabled=True,
external_trace_jsonl_path="/shared/traces.jsonl",
)

def test_trace_span_records_attributes_events_and_errors(self):
repo_root = Path(__file__).resolve().parents[2]
output = _run_trace_utils(repo_root, "record-span")
Expand Down
Loading
Loading