Hint users when wrongly execute it with partial ranks in dumper (#19014)
This commit is contained in:
@@ -1,5 +1,8 @@
|
||||
import io
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
@@ -10,6 +13,7 @@ import torch.distributed as dist
|
||||
from sglang.srt.debug_utils.dumper import (
|
||||
_collect_megatron_parallel_info,
|
||||
_collect_sglang_parallel_info,
|
||||
_collective_with_timeout,
|
||||
_Dumper,
|
||||
_materialize_value,
|
||||
_obj_to_dict,
|
||||
@@ -25,6 +29,17 @@ register_cuda_ci(est_time=30, suite="nightly-2-gpu", nightly=True)
|
||||
register_amd_ci(est_time=60, suite="nightly-amd", nightly=True)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _capture_stdout():
|
||||
captured = io.StringIO()
|
||||
old_stdout = sys.stdout
|
||||
sys.stdout = captured
|
||||
try:
|
||||
yield captured
|
||||
finally:
|
||||
sys.stdout = old_stdout
|
||||
|
||||
|
||||
class TestDumperPureFunctions:
|
||||
def test_get_truncated_value(self):
|
||||
assert get_truncated_value(None) is None
|
||||
@@ -86,6 +101,34 @@ class TestTorchSave:
|
||||
assert "skip the tensor" in captured.out
|
||||
|
||||
|
||||
class TestCollectiveTimeout:
|
||||
def test_watchdog_fires_on_timeout(self):
|
||||
block_event = threading.Event()
|
||||
output = ""
|
||||
|
||||
def run_with_timeout():
|
||||
nonlocal output
|
||||
with _capture_stdout() as captured:
|
||||
_collective_with_timeout(
|
||||
lambda: block_event.wait(),
|
||||
operation_name="test_blocked_op",
|
||||
timeout_seconds=2,
|
||||
)
|
||||
output = captured.getvalue()
|
||||
|
||||
worker = threading.Thread(target=run_with_timeout)
|
||||
worker.start()
|
||||
|
||||
time.sleep(4)
|
||||
block_event.set()
|
||||
worker.join(timeout=5)
|
||||
|
||||
print(f"Captured output: {output!r}")
|
||||
assert "WARNING" in output
|
||||
assert "test_blocked_op" in output
|
||||
assert "2s" in output
|
||||
|
||||
|
||||
class TestDumperDistributed:
|
||||
def test_basic(self, tmp_path):
|
||||
with temp_set_env(
|
||||
@@ -125,6 +168,32 @@ class TestDumperDistributed:
|
||||
not_exist=["tensor_skip"],
|
||||
)
|
||||
|
||||
def test_collective_timeout(self):
|
||||
with temp_set_env(allow_sglang=True, SGLANG_DUMPER_ENABLE="1"):
|
||||
run_distributed_test(self._test_collective_timeout_func)
|
||||
|
||||
@staticmethod
|
||||
def _test_collective_timeout_func(rank):
|
||||
dumper = _Dumper(
|
||||
enable=True,
|
||||
base_dir=Path("/tmp"),
|
||||
partial_name=None,
|
||||
enable_http_server=False,
|
||||
collective_timeout=3,
|
||||
)
|
||||
|
||||
with _capture_stdout() as captured:
|
||||
if rank != 0:
|
||||
time.sleep(6)
|
||||
dumper.on_forward_pass_start()
|
||||
|
||||
output = captured.getvalue()
|
||||
print(f"Rank {rank} captured output: {output!r}")
|
||||
|
||||
if rank == 0:
|
||||
assert "WARNING" in output, f"Expected WARNING in rank 0 output: {output}"
|
||||
assert "has not completed after 3s" in output
|
||||
|
||||
def test_http_enable(self):
|
||||
with temp_set_env(allow_sglang=True, SGLANG_DUMPER_ENABLE="0"):
|
||||
run_distributed_test(self._test_http_func)
|
||||
|
||||
Reference in New Issue
Block a user