Hint users when wrongly execute it with partial ranks in dumper (#19014)
This commit is contained in:
@@ -57,6 +57,7 @@ class _Dumper:
|
|||||||
partial_name: Optional[str] = None,
|
partial_name: Optional[str] = None,
|
||||||
enable_http_server: bool = True,
|
enable_http_server: bool = True,
|
||||||
cleanup_previous: bool = False,
|
cleanup_previous: bool = False,
|
||||||
|
collective_timeout: int = 60,
|
||||||
):
|
):
|
||||||
# Config
|
# Config
|
||||||
self._enable = enable
|
self._enable = enable
|
||||||
@@ -68,6 +69,7 @@ class _Dumper:
|
|||||||
self._enable_grad = enable_grad
|
self._enable_grad = enable_grad
|
||||||
self._enable_model_value = enable_model_value
|
self._enable_model_value = enable_model_value
|
||||||
self._enable_model_grad = enable_model_grad
|
self._enable_model_grad = enable_model_grad
|
||||||
|
self._collective_timeout = collective_timeout
|
||||||
|
|
||||||
# States
|
# States
|
||||||
self._partial_name = partial_name
|
self._partial_name = partial_name
|
||||||
@@ -96,6 +98,7 @@ class _Dumper:
|
|||||||
"SGLANG_ENABLE_DUMPER_HTTP_SERVER", "1"
|
"SGLANG_ENABLE_DUMPER_HTTP_SERVER", "1"
|
||||||
),
|
),
|
||||||
cleanup_previous=get_bool_env_var("SGLANG_DUMPER_CLEANUP_PREVIOUS", "0"),
|
cleanup_previous=get_bool_env_var("SGLANG_DUMPER_CLEANUP_PREVIOUS", "0"),
|
||||||
|
collective_timeout=60,
|
||||||
)
|
)
|
||||||
|
|
||||||
def on_forward_pass_start(self):
|
def on_forward_pass_start(self):
|
||||||
@@ -119,11 +122,13 @@ class _Dumper:
|
|||||||
if self._http_server_handled:
|
if self._http_server_handled:
|
||||||
return
|
return
|
||||||
self._http_server_handled = True
|
self._http_server_handled = True
|
||||||
_start_maybe_http_server(self)
|
_start_maybe_http_server(self, timeout_seconds=self._collective_timeout)
|
||||||
|
|
||||||
def _ensure_partial_name(self):
|
def _ensure_partial_name(self):
|
||||||
if self._partial_name is None:
|
if self._partial_name is None:
|
||||||
self._partial_name = _get_partial_name()
|
self._partial_name = _get_partial_name(
|
||||||
|
timeout_seconds=self._collective_timeout
|
||||||
|
)
|
||||||
print(f"[Dumper] Choose partial_name={self._partial_name}")
|
print(f"[Dumper] Choose partial_name={self._partial_name}")
|
||||||
|
|
||||||
def set_ctx(self, **kwargs):
|
def set_ctx(self, **kwargs):
|
||||||
@@ -331,11 +336,37 @@ def _torch_save(value, path: str):
|
|||||||
print(f"[Dumper] Observe error={e} when saving data, skip the tensor")
|
print(f"[Dumper] Observe error={e} when saving data, skip the tensor")
|
||||||
|
|
||||||
|
|
||||||
def _get_partial_name():
|
def _collective_with_timeout(fn, operation_name: str, timeout_seconds: int = 60):
|
||||||
|
completed = threading.Event()
|
||||||
|
|
||||||
|
def watchdog():
|
||||||
|
if not completed.wait(timeout=timeout_seconds):
|
||||||
|
print(
|
||||||
|
f"\n[Dumper] WARNING: '{operation_name}' has not completed after "
|
||||||
|
f"{timeout_seconds}s. This usually means not all ranks are "
|
||||||
|
f"participating in this collective operation.\n",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
thread = threading.Thread(target=watchdog, daemon=True)
|
||||||
|
thread.start()
|
||||||
|
try:
|
||||||
|
return fn()
|
||||||
|
finally:
|
||||||
|
completed.set()
|
||||||
|
|
||||||
|
|
||||||
|
def _get_partial_name(timeout_seconds: int = 60):
|
||||||
rank = _get_rank()
|
rank = _get_rank()
|
||||||
object_list = [str(time.time()) if rank == 0 else None]
|
object_list = [str(time.time()) if rank == 0 else None]
|
||||||
|
|
||||||
if dist.is_initialized():
|
if dist.is_initialized():
|
||||||
dist.broadcast_object_list(object_list, device="cuda")
|
_collective_with_timeout(
|
||||||
|
lambda: dist.broadcast_object_list(object_list, device="cuda"),
|
||||||
|
operation_name="broadcast_object_list in _get_partial_name",
|
||||||
|
timeout_seconds=timeout_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
return object_list[0]
|
return object_list[0]
|
||||||
|
|
||||||
|
|
||||||
@@ -494,14 +525,16 @@ def _collect_megatron_parallel_info():
|
|||||||
# -------------------------------------- http control server ------------------------------------------
|
# -------------------------------------- http control server ------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _start_maybe_http_server(dumper):
|
def _start_maybe_http_server(dumper, timeout_seconds: int = 60):
|
||||||
http_port = get_int_env_var("SGLANG_DUMPER_SERVER_PORT", 40000)
|
http_port = get_int_env_var("SGLANG_DUMPER_SERVER_PORT", 40000)
|
||||||
zmq_base_port = get_int_env_var("SGLANG_DUMPER_ZMQ_BASE_PORT", 16800)
|
zmq_base_port = get_int_env_var("SGLANG_DUMPER_ZMQ_BASE_PORT", 16800)
|
||||||
if http_port <= 0:
|
if http_port <= 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
local_handler = _DumperRpcHandler(dumper)
|
local_handler = _DumperRpcHandler(dumper)
|
||||||
rpc_handles = _create_zmq_rpc_handles(local_handler, base_port=zmq_base_port)
|
rpc_handles = _create_zmq_rpc_handles(
|
||||||
|
local_handler, base_port=zmq_base_port, timeout_seconds=timeout_seconds
|
||||||
|
)
|
||||||
|
|
||||||
if _get_rank() == 0:
|
if _get_rank() == 0:
|
||||||
handler_class = _make_dumper_http_handler(rpc_handles=rpc_handles)
|
handler_class = _make_dumper_http_handler(rpc_handles=rpc_handles)
|
||||||
@@ -549,7 +582,9 @@ class _DumperRpcHandler:
|
|||||||
# -------------------------------------- zmq rpc ------------------------------------------
|
# -------------------------------------- zmq rpc ------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _create_zmq_rpc_handles(handler, base_port: int) -> Optional[List["_ZmqRpcHandle"]]:
|
def _create_zmq_rpc_handles(
|
||||||
|
handler, base_port: int, timeout_seconds: int = 60
|
||||||
|
) -> Optional[List["_ZmqRpcHandle"]]:
|
||||||
import zmq
|
import zmq
|
||||||
|
|
||||||
rank = _get_rank()
|
rank = _get_rank()
|
||||||
@@ -578,7 +613,11 @@ def _create_zmq_rpc_handles(handler, base_port: int) -> Optional[List["_ZmqRpcHa
|
|||||||
|
|
||||||
if dist.is_initialized():
|
if dist.is_initialized():
|
||||||
all_addresses = [None] * world_size
|
all_addresses = [None] * world_size
|
||||||
dist.all_gather_object(all_addresses, local_addr)
|
_collective_with_timeout(
|
||||||
|
lambda: dist.all_gather_object(all_addresses, local_addr),
|
||||||
|
operation_name="all_gather_object in _create_zmq_rpc_handles",
|
||||||
|
timeout_seconds=timeout_seconds,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
all_addresses = [local_addr]
|
all_addresses = [local_addr]
|
||||||
print(f"[Dumper.ZmqRpc] rank={rank} all_addresses={all_addresses}")
|
print(f"[Dumper.ZmqRpc] rank={rank} all_addresses={all_addresses}")
|
||||||
|
|||||||
@@ -1,5 +1,8 @@
|
|||||||
|
import io
|
||||||
import sys
|
import sys
|
||||||
|
import threading
|
||||||
import time
|
import time
|
||||||
|
from contextlib import contextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -10,6 +13,7 @@ import torch.distributed as dist
|
|||||||
from sglang.srt.debug_utils.dumper import (
|
from sglang.srt.debug_utils.dumper import (
|
||||||
_collect_megatron_parallel_info,
|
_collect_megatron_parallel_info,
|
||||||
_collect_sglang_parallel_info,
|
_collect_sglang_parallel_info,
|
||||||
|
_collective_with_timeout,
|
||||||
_Dumper,
|
_Dumper,
|
||||||
_materialize_value,
|
_materialize_value,
|
||||||
_obj_to_dict,
|
_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)
|
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:
|
class TestDumperPureFunctions:
|
||||||
def test_get_truncated_value(self):
|
def test_get_truncated_value(self):
|
||||||
assert get_truncated_value(None) is None
|
assert get_truncated_value(None) is None
|
||||||
@@ -86,6 +101,34 @@ class TestTorchSave:
|
|||||||
assert "skip the tensor" in captured.out
|
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:
|
class TestDumperDistributed:
|
||||||
def test_basic(self, tmp_path):
|
def test_basic(self, tmp_path):
|
||||||
with temp_set_env(
|
with temp_set_env(
|
||||||
@@ -125,6 +168,32 @@ class TestDumperDistributed:
|
|||||||
not_exist=["tensor_skip"],
|
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):
|
def test_http_enable(self):
|
||||||
with temp_set_env(allow_sglang=True, SGLANG_DUMPER_ENABLE="0"):
|
with temp_set_env(allow_sglang=True, SGLANG_DUMPER_ENABLE="0"):
|
||||||
run_distributed_test(self._test_http_func)
|
run_distributed_test(self._test_http_func)
|
||||||
|
|||||||
Reference in New Issue
Block a user