[Elastic EP] Fix recovery lifecycle and add manual coverage (#31744)
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -11,9 +11,10 @@ from sglang.srt.distributed import get_world_group, parallel_state
|
|||||||
from sglang.srt.distributed.utils import get_global_tcp_store
|
from sglang.srt.distributed.utils import get_global_tcp_store
|
||||||
from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata
|
from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata
|
||||||
from sglang.srt.managers.schedule_batch import ServerArgs
|
from sglang.srt.managers.schedule_batch import ServerArgs
|
||||||
from sglang.srt.utils import broadcast_pyobj, is_cpu, is_cuda
|
from sglang.srt.utils import is_cpu, is_cuda
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.eplb.eplb_manager import EPLBManager
|
from sglang.srt.eplb.eplb_manager import EPLBManager
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -462,7 +463,8 @@ def maybe_recover_ep_ranks(
|
|||||||
*,
|
*,
|
||||||
tp_group: parallel_state.GroupCoordinator,
|
tp_group: parallel_state.GroupCoordinator,
|
||||||
eplb_manager: EPLBManager,
|
eplb_manager: EPLBManager,
|
||||||
random_seed: int,
|
model_config: ModelConfig,
|
||||||
|
moe_ep_rank: int,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
# TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device
|
# TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device
|
||||||
# synchronization, and this function is on the forward-path.
|
# synchronization, and this function is on the forward-path.
|
||||||
@@ -489,17 +491,13 @@ def maybe_recover_ep_ranks(
|
|||||||
if ranks_to_recover and try_recover_ranks(ranks_to_recover):
|
if ranks_to_recover and try_recover_ranks(ranks_to_recover):
|
||||||
eplb_manager.reset_generator()
|
eplb_manager.reset_generator()
|
||||||
broadcast_global_expert_location_metadata(
|
broadcast_global_expert_location_metadata(
|
||||||
|
model_config=model_config,
|
||||||
|
moe_ep_rank=moe_ep_rank,
|
||||||
src_rank=get_healthy_expert_location_src_rank(
|
src_rank=get_healthy_expert_location_src_rank(
|
||||||
invoked_in_elastic_ep_rejoin_path=False
|
invoked_in_elastic_ep_rejoin_path=False
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
ElasticEPStateManager.instance().reset()
|
ElasticEPStateManager.instance().reset()
|
||||||
broadcast_pyobj(
|
|
||||||
[random_seed],
|
|
||||||
parallel_state.get_world_group().rank,
|
|
||||||
parallel_state.get_world_group().cpu_group,
|
|
||||||
src=parallel_state.get_world_group().ranks[0],
|
|
||||||
)
|
|
||||||
logger.info(f"recover ranks {ranks_to_recover} done")
|
logger.info(f"recover ranks {ranks_to_recover} done")
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|||||||
@@ -885,6 +885,12 @@ class Scheduler(
|
|||||||
if model_runner.token_to_kv_pool.post_capture_active:
|
if model_runner.token_to_kv_pool.post_capture_active:
|
||||||
model_runner.post_capture_resize_kv_pool()
|
model_runner.post_capture_resize_kv_pool()
|
||||||
|
|
||||||
|
if (
|
||||||
|
self.server_args.elastic_ep_backend is not None
|
||||||
|
and self.server_args.ep_join_mode == "recover"
|
||||||
|
):
|
||||||
|
model_runner.post_capture_elastic_ep_recover()
|
||||||
|
|
||||||
# Dispatch the model worker
|
# Dispatch the model worker
|
||||||
if self.spec_algorithm.is_none():
|
if self.spec_algorithm.is_none():
|
||||||
self.model_worker = self.tp_worker
|
self.model_worker = self.tp_worker
|
||||||
|
|||||||
@@ -342,8 +342,8 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
self.world_group = get_world_group()
|
self.world_group = get_world_group()
|
||||||
|
|
||||||
# Sync random seed across TP workers.
|
# Sync random seed across TP workers.
|
||||||
# Scale joiners cannot enter the launch-time WORLD broadcast.
|
# Elastic joiners cannot enter the launch-time WORLD broadcast.
|
||||||
if server_args.is_ep_scale_joiner:
|
if server_args.is_ep_joiner:
|
||||||
self.random_seed = server_args.random_seed
|
self.random_seed = server_args.random_seed
|
||||||
else:
|
else:
|
||||||
self.random_seed = broadcast_pyobj(
|
self.random_seed = broadcast_pyobj(
|
||||||
|
|||||||
@@ -399,15 +399,11 @@ class ModelRunner:
|
|||||||
def _initialize_elastic_ep_joiner(self) -> None:
|
def _initialize_elastic_ep_joiner(self) -> None:
|
||||||
if not (
|
if not (
|
||||||
self.server_args.elastic_ep_backend is not None
|
self.server_args.elastic_ep_backend is not None
|
||||||
and self.server_args.is_ep_joiner
|
and self.server_args.is_ep_scale_joiner
|
||||||
):
|
):
|
||||||
return
|
return
|
||||||
|
|
||||||
is_scale_join = self.server_args.ep_join_mode == "scale"
|
join_effective_ep_size = self.server_args.ep_join_rank_offset + self.ps.tp_size
|
||||||
if is_scale_join:
|
|
||||||
join_effective_ep_size = (
|
|
||||||
self.server_args.ep_join_rank_offset + self.ps.tp_size
|
|
||||||
)
|
|
||||||
dist.barrier(group=self.tp_group.cpu_group)
|
dist.barrier(group=self.tp_group.cpu_group)
|
||||||
if self.ps.tp_rank == 0:
|
if self.ps.tp_rank == 0:
|
||||||
register_scale_cohort(
|
register_scale_cohort(
|
||||||
@@ -418,20 +414,12 @@ class ModelRunner:
|
|||||||
self.server_args.override(
|
self.server_args.override(
|
||||||
"elastic_ep.scale_join", ep_size=join_effective_ep_size
|
"elastic_ep.scale_join", ep_size=join_effective_ep_size
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
join_process_groups()
|
|
||||||
|
|
||||||
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
||||||
broadcast_global_expert_location_metadata(
|
broadcast_global_expert_location_metadata(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
moe_ep_rank=global_ep_rank,
|
moe_ep_rank=global_ep_rank,
|
||||||
src_rank=(
|
src_rank=0,
|
||||||
0
|
|
||||||
if is_scale_join
|
|
||||||
else get_healthy_expert_location_src_rank(
|
|
||||||
invoked_in_elastic_ep_rejoin_path=True
|
|
||||||
)
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
set_global_expert_distribution_recorder(
|
set_global_expert_distribution_recorder(
|
||||||
ExpertDistributionRecorder.init_new(
|
ExpertDistributionRecorder.init_new(
|
||||||
@@ -441,10 +429,6 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
if not is_scale_join:
|
|
||||||
ElasticEPStateManager.instance().reset()
|
|
||||||
return
|
|
||||||
|
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
enable_joiner_all_gather,
|
enable_joiner_all_gather,
|
||||||
update_dp_attention_post_scale,
|
update_dp_attention_post_scale,
|
||||||
@@ -839,6 +823,27 @@ class ModelRunner:
|
|||||||
resize.capped_max_running_requests
|
resize.capped_max_running_requests
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def post_capture_elastic_ep_recover(self):
|
||||||
|
join_process_groups()
|
||||||
|
|
||||||
|
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
||||||
|
broadcast_global_expert_location_metadata(
|
||||||
|
model_config=self.model_config,
|
||||||
|
moe_ep_rank=global_ep_rank,
|
||||||
|
src_rank=get_healthy_expert_location_src_rank(
|
||||||
|
invoked_in_elastic_ep_rejoin_path=True
|
||||||
|
),
|
||||||
|
)
|
||||||
|
set_global_expert_distribution_recorder(
|
||||||
|
ExpertDistributionRecorder.init_new(
|
||||||
|
self.server_args,
|
||||||
|
get_global_expert_location_metadata(),
|
||||||
|
rank=global_ep_rank,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
ElasticEPStateManager.instance().reset()
|
||||||
|
|
||||||
def init_attention_backends(self):
|
def init_attention_backends(self):
|
||||||
"""Initialize attention backends only (no cuda graph capture)."""
|
"""Initialize attention backends only (no cuda graph capture)."""
|
||||||
# Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture
|
# Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture
|
||||||
@@ -1089,7 +1094,7 @@ class ModelRunner:
|
|||||||
dist_barrier_after_load(
|
dist_barrier_after_load(
|
||||||
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
is_ep_scale_joiner=self.server_args.is_ep_scale_joiner,
|
is_ep_joiner=self.server_args.is_ep_joiner,
|
||||||
)
|
)
|
||||||
|
|
||||||
def maybe_init_dwdp(self):
|
def maybe_init_dwdp(self):
|
||||||
@@ -1822,7 +1827,8 @@ class ModelRunner:
|
|||||||
recovered = maybe_recover_ep_ranks(
|
recovered = maybe_recover_ep_ranks(
|
||||||
tp_group=self.tp_group,
|
tp_group=self.tp_group,
|
||||||
eplb_manager=self.eplb_manager,
|
eplb_manager=self.eplb_manager,
|
||||||
random_seed=self.server_args.random_seed,
|
model_config=self.model_config,
|
||||||
|
moe_ep_rank=self._elastic_global_rank(),
|
||||||
)
|
)
|
||||||
if recovered:
|
if recovered:
|
||||||
self.forward_pass_id = 0
|
self.forward_pass_id = 0
|
||||||
|
|||||||
@@ -301,11 +301,11 @@ def dist_barrier_after_load(
|
|||||||
*,
|
*,
|
||||||
elastic_ep_backend: Optional[str],
|
elastic_ep_backend: Optional[str],
|
||||||
tp_rank: int,
|
tp_rank: int,
|
||||||
is_ep_scale_joiner: bool = False,
|
is_ep_joiner: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
if elastic_ep_backend == "mooncake":
|
if elastic_ep_backend == "mooncake":
|
||||||
# Mooncake does not support `monitored_barrier`
|
# Mooncake does not support `monitored_barrier`
|
||||||
if not is_ep_scale_joiner:
|
if not is_ep_joiner:
|
||||||
dist.barrier(group=get_tp_group().cpu_group)
|
dist.barrier(group=get_tp_group().cpu_group)
|
||||||
else:
|
else:
|
||||||
# Handle the case where some ranks do not finish loading.
|
# Handle the case where some ranks do not finish loading.
|
||||||
|
|||||||
@@ -8984,7 +8984,7 @@ class PortArgs:
|
|||||||
# (no availability-based search). If incrementing would
|
# (no availability-based search). If incrementing would
|
||||||
# overflow the valid TCP range, decrement instead.
|
# overflow the valid TCP range, decrement instead.
|
||||||
NUM_DERIVED_PORTS = 5
|
NUM_DERIVED_PORTS = 5
|
||||||
if server_args.is_ep_scale_joiner:
|
if server_args.is_ep_joiner:
|
||||||
port_base = server_args.port + ZMQ_TCP_PORT_DELTA
|
port_base = server_args.port + ZMQ_TCP_PORT_DELTA
|
||||||
if port_base + NUM_DERIVED_PORTS > 65535:
|
if port_base + NUM_DERIVED_PORTS > 65535:
|
||||||
port_base = server_args.port - ZMQ_TCP_PORT_DELTA
|
port_base = server_args.port - ZMQ_TCP_PORT_DELTA
|
||||||
@@ -9004,7 +9004,7 @@ class PortArgs:
|
|||||||
assert worker_ports is not None
|
assert worker_ports is not None
|
||||||
scheduler_input_port = worker_ports[dp_rank]
|
scheduler_input_port = worker_ports[dp_rank]
|
||||||
|
|
||||||
is_joiner = server_args.is_ep_scale_joiner
|
is_joiner = server_args.is_ep_joiner
|
||||||
# Under SGLANG_DISTRIBUTED_INIT_METHOD_OVERRIDE, SGLang never binds
|
# Under SGLANG_DISTRIBUTED_INIT_METHOD_OVERRIDE, SGLang never binds
|
||||||
# dist_init_port / nccl_port (rendezvous uses the externally-managed
|
# dist_init_port / nccl_port (rendezvous uses the externally-managed
|
||||||
# store; see distributed/bootstrap.py:_resolve_dist_init_method), so
|
# store; see distributed/bootstrap.py:_resolve_dist_init_method), so
|
||||||
|
|||||||
@@ -0,0 +1,255 @@
|
|||||||
|
"""Manual single-host Elastic EP recovery test.
|
||||||
|
|
||||||
|
Run:
|
||||||
|
|
||||||
|
CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7 python -m pytest \
|
||||||
|
test/manual/ep/test_elastic_recover.py -v -s
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import shlex
|
||||||
|
import subprocess
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.server_fixtures.disaggregation_fixture import get_rdma_devices_args
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_MODEL_NAME_FOR_TEST_MLA,
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
CustomTestCase,
|
||||||
|
try_cached_model,
|
||||||
|
)
|
||||||
|
from sglang.utils import wait_for_http_ready
|
||||||
|
|
||||||
|
TEST_MODEL = os.environ.get(
|
||||||
|
"SGLANG_ELASTIC_RECOVER_TEST_MODEL",
|
||||||
|
try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST_MLA),
|
||||||
|
)
|
||||||
|
EP_SIZE = 8
|
||||||
|
LOCAL_EP_SIZE = 4
|
||||||
|
DIST_INIT_ADDR = os.environ.get("SGLANG_ELASTIC_RECOVER_DIST_INIT", "127.0.0.1:25555")
|
||||||
|
PRIMARY_PORT = int(os.environ.get("SGLANG_ELASTIC_RECOVER_PRIMARY_PORT", "21000"))
|
||||||
|
JOINER_PORT = int(os.environ.get("SGLANG_ELASTIC_RECOVER_JOINER_PORT", "22000"))
|
||||||
|
RECOVER_WAIT_SECONDS = float(os.environ.get("SGLANG_ELASTIC_RECOVER_WAIT_SECONDS", "5"))
|
||||||
|
RECOVER_TIMEOUT_SECONDS = float(
|
||||||
|
os.environ.get("SGLANG_ELASTIC_RECOVER_TIMEOUT_SECONDS", "300")
|
||||||
|
)
|
||||||
|
RANDOM_SEED = int(os.environ.get("SGLANG_ELASTIC_RECOVER_RANDOM_SEED", "42"))
|
||||||
|
ib_devices = get_rdma_devices_args()
|
||||||
|
|
||||||
|
|
||||||
|
def _visible_device_ids() -> list[str]:
|
||||||
|
visible = os.environ.get("CUDA_VISIBLE_DEVICES")
|
||||||
|
if visible:
|
||||||
|
return [device.strip() for device in visible.split(",") if device.strip()]
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
return [str(index) for index in range(torch.cuda.device_count())]
|
||||||
|
except Exception:
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def _server_args(node_rank: int, port: int, recover: bool = False) -> list[str]:
|
||||||
|
args = [
|
||||||
|
"sglang",
|
||||||
|
"serve",
|
||||||
|
"--model-path",
|
||||||
|
TEST_MODEL,
|
||||||
|
"--host",
|
||||||
|
"127.0.0.1",
|
||||||
|
"--port",
|
||||||
|
str(port),
|
||||||
|
"--device",
|
||||||
|
"cuda",
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--tp",
|
||||||
|
str(EP_SIZE),
|
||||||
|
"--dp",
|
||||||
|
str(EP_SIZE),
|
||||||
|
"--nnodes",
|
||||||
|
"2",
|
||||||
|
"--node-rank",
|
||||||
|
str(node_rank),
|
||||||
|
"--dist-init-addr",
|
||||||
|
DIST_INIT_ADDR,
|
||||||
|
"--random-seed",
|
||||||
|
str(RANDOM_SEED),
|
||||||
|
"--enable-dp-attention",
|
||||||
|
"--enable-dp-lm-head",
|
||||||
|
"--elastic-ep-backend",
|
||||||
|
"mooncake",
|
||||||
|
"--mooncake-ib-device",
|
||||||
|
ib_devices,
|
||||||
|
"--moe-a2a-backend",
|
||||||
|
"mooncake",
|
||||||
|
"--deepep-mode",
|
||||||
|
"low_latency",
|
||||||
|
"--moe-dense-tp-size",
|
||||||
|
"1",
|
||||||
|
"--disable-custom-all-reduce",
|
||||||
|
"--enable-eplb",
|
||||||
|
"--ep-num-redundant-experts",
|
||||||
|
"72",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"512",
|
||||||
|
"--cuda-graph-max-bs-decode",
|
||||||
|
"16",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.5",
|
||||||
|
]
|
||||||
|
if recover:
|
||||||
|
args.extend(["--elastic-ep-join-mode", "recover"])
|
||||||
|
extra_args = os.environ.get("SGLANG_ELASTIC_RECOVER_EXTRA_SERVER_ARGS", "")
|
||||||
|
return args + shlex.split(extra_args)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(
|
||||||
|
len(_visible_device_ids()) >= EP_SIZE,
|
||||||
|
"Elastic EP recovery E2E needs 8 visible GPUs.",
|
||||||
|
)
|
||||||
|
class TestElasticRecover4To4(CustomTestCase):
|
||||||
|
"""Kill one four-rank node and recover it with a fresh process group."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.base_url = f"http://127.0.0.1:{PRIMARY_PORT}"
|
||||||
|
cls.processes: list[subprocess.Popen] = []
|
||||||
|
cls.log_files = []
|
||||||
|
cls.log_paths: dict[str, Path] = {}
|
||||||
|
visible_devices = _visible_device_ids()
|
||||||
|
|
||||||
|
cls.primary = cls._launch(
|
||||||
|
node_rank=0,
|
||||||
|
port=PRIMARY_PORT,
|
||||||
|
visible_devices=visible_devices[:LOCAL_EP_SIZE],
|
||||||
|
name="primary",
|
||||||
|
)
|
||||||
|
cls.initial_joiner = cls._launch(
|
||||||
|
node_rank=1,
|
||||||
|
port=JOINER_PORT,
|
||||||
|
visible_devices=visible_devices[LOCAL_EP_SIZE:EP_SIZE],
|
||||||
|
name="initial_joiner",
|
||||||
|
)
|
||||||
|
wait_for_http_ready(
|
||||||
|
f"{cls.base_url}/health_generate",
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
process=cls.primary,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _launch(
|
||||||
|
cls,
|
||||||
|
*,
|
||||||
|
node_rank: int,
|
||||||
|
port: int,
|
||||||
|
visible_devices: list[str],
|
||||||
|
name: str,
|
||||||
|
recover: bool = False,
|
||||||
|
) -> subprocess.Popen:
|
||||||
|
log_path = Path(f"/tmp/elastic_ep_recover_{name}_{int(time.time())}.log")
|
||||||
|
log_file = open(log_path, "w")
|
||||||
|
env = os.environ.copy()
|
||||||
|
env["CUDA_VISIBLE_DEVICES"] = ",".join(visible_devices)
|
||||||
|
process = subprocess.Popen(
|
||||||
|
_server_args(node_rank, port, recover),
|
||||||
|
env=env,
|
||||||
|
stdout=log_file,
|
||||||
|
stderr=subprocess.STDOUT,
|
||||||
|
)
|
||||||
|
cls.processes.append(process)
|
||||||
|
cls.log_files.append(log_file)
|
||||||
|
cls.log_paths[name] = log_path
|
||||||
|
print(f"Started {name}; log: {log_path}")
|
||||||
|
return process
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
for process in reversed(getattr(cls, "processes", [])):
|
||||||
|
if process.poll() is None:
|
||||||
|
kill_process_tree(process.pid, wait_timeout=60)
|
||||||
|
for log_file in getattr(cls, "log_files", []):
|
||||||
|
log_file.close()
|
||||||
|
|
||||||
|
def _generate(self, routed_dp_rank: int | None = None) -> requests.Response:
|
||||||
|
payload = {
|
||||||
|
"text": "The capital of France is",
|
||||||
|
"sampling_params": {"max_new_tokens": 4, "temperature": 0.0},
|
||||||
|
}
|
||||||
|
if routed_dp_rank is not None:
|
||||||
|
payload["routed_dp_rank"] = routed_dp_rank
|
||||||
|
return requests.post(f"{self.base_url}/generate", json=payload, timeout=90)
|
||||||
|
|
||||||
|
def _generate_ok(self, description: str, routed_dp_rank: int | None = None) -> None:
|
||||||
|
response = self._generate(routed_dp_rank)
|
||||||
|
self.assertEqual(response.status_code, 200, f"{description}: {response.text}")
|
||||||
|
payload = response.json()
|
||||||
|
generated_text = payload.get("text", "")
|
||||||
|
self.assertIn(
|
||||||
|
"paris",
|
||||||
|
generated_text.casefold(),
|
||||||
|
f"{description}: unexpected generation: {generated_text!r}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _wait_for_recover_capture(self) -> None:
|
||||||
|
deadline = time.monotonic() + RECOVER_TIMEOUT_SECONDS
|
||||||
|
log_path = self.log_paths["recover_joiner"]
|
||||||
|
marker = "Capture target decode CUDA graph end"
|
||||||
|
while time.monotonic() < deadline:
|
||||||
|
self.assertIsNone(
|
||||||
|
self.recover_joiner.poll(),
|
||||||
|
"Recover joiner exited during CUDA graph capture",
|
||||||
|
)
|
||||||
|
if log_path.exists() and log_path.read_text(errors="replace").count(
|
||||||
|
marker
|
||||||
|
) >= (EP_SIZE - LOCAL_EP_SIZE):
|
||||||
|
return
|
||||||
|
time.sleep(2)
|
||||||
|
self.fail(f"Timed out waiting for recover CUDA graph capture: {log_path}")
|
||||||
|
|
||||||
|
def _wait_for_recovered_ranks(self) -> None:
|
||||||
|
self._wait_for_recover_capture()
|
||||||
|
self._generate_ok("recovery trigger")
|
||||||
|
deadline = time.monotonic() + RECOVER_TIMEOUT_SECONDS
|
||||||
|
marker = f"recover ranks {list(range(LOCAL_EP_SIZE, EP_SIZE))} done"
|
||||||
|
primary_log = self.log_paths["primary"]
|
||||||
|
while time.monotonic() < deadline:
|
||||||
|
self.assertIsNone(
|
||||||
|
self.recover_joiner.poll(), "Recover joiner exited before rejoining"
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
primary_log.exists()
|
||||||
|
and primary_log.read_text(errors="replace").count(marker)
|
||||||
|
>= LOCAL_EP_SIZE
|
||||||
|
):
|
||||||
|
for request_index in range(3):
|
||||||
|
self._generate_ok(f"post-recovery request {request_index + 1}")
|
||||||
|
return
|
||||||
|
time.sleep(2)
|
||||||
|
self.fail(f"Timed out waiting for recovery collective: {primary_log}")
|
||||||
|
|
||||||
|
def test_recover_four_ranks(self):
|
||||||
|
self._generate_ok("initial service")
|
||||||
|
|
||||||
|
kill_process_tree(self.initial_joiner.pid, wait_timeout=60)
|
||||||
|
# Give the terminated schedulers time to disappear before fault handling.
|
||||||
|
time.sleep(RECOVER_WAIT_SECONDS)
|
||||||
|
self._generate_ok("degraded service after node1 failure")
|
||||||
|
|
||||||
|
visible_devices = _visible_device_ids()
|
||||||
|
self.recover_joiner = self._launch(
|
||||||
|
node_rank=1,
|
||||||
|
port=JOINER_PORT,
|
||||||
|
visible_devices=visible_devices[LOCAL_EP_SIZE:EP_SIZE],
|
||||||
|
name="recover_joiner",
|
||||||
|
recover=True,
|
||||||
|
)
|
||||||
|
self._wait_for_recovered_ranks()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user