TP/PP Consensus checker (#34406)
This commit is contained in:
@@ -1068,6 +1068,11 @@ SGLang supports various environment variables that can be used to configure its
|
|||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Detect and report ranks that fall behind during collective ops.</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Detect and report ranks that fall behind during collective ops.</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_ENABLE_RANK_CONSENSUS_CHECKER</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Check for PP/TP-rank divergence. Kill the server when divergence occurs. Helpful for trouble-shooting server hangs issues.</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>False</code></td>
|
||||||
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FORCE_SHUTDOWN</code></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FORCE_SHUTDOWN</code></td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Force an immediate process-group shutdown on exit.</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Force an immediate process-group shutdown on exit.</td>
|
||||||
|
|||||||
@@ -329,6 +329,7 @@ class Envs:
|
|||||||
SGLANG_LOG_REQUEST_HEADERS = EnvTuple(tuple())
|
SGLANG_LOG_REQUEST_HEADERS = EnvTuple(tuple())
|
||||||
SGLANG_LOG_SCHEDULER_STATUS_TARGET = EnvStr("")
|
SGLANG_LOG_SCHEDULER_STATUS_TARGET = EnvStr("")
|
||||||
SGLANG_LOG_SCHEDULER_STATUS_INTERVAL = EnvFloat(60.0)
|
SGLANG_LOG_SCHEDULER_STATUS_INTERVAL = EnvFloat(60.0)
|
||||||
|
SGLANG_ENABLE_RANK_CONSENSUS_CHECKER = EnvBool(False)
|
||||||
|
|
||||||
# ===================================================================
|
# ===================================================================
|
||||||
# IPC, broadcasters, and ports
|
# IPC, broadcasters, and ports
|
||||||
|
|||||||
@@ -309,6 +309,7 @@ from sglang.srt.utils import (
|
|||||||
is_hip,
|
is_hip,
|
||||||
is_mps,
|
is_mps,
|
||||||
kill_itself_when_parent_died,
|
kill_itself_when_parent_died,
|
||||||
|
rank_consensus_checker,
|
||||||
require_mlp_sync,
|
require_mlp_sync,
|
||||||
set_gpu_proc_affinity,
|
set_gpu_proc_affinity,
|
||||||
set_random_seed,
|
set_random_seed,
|
||||||
@@ -655,6 +656,8 @@ class Scheduler(
|
|||||||
|
|
||||||
self.init_batch_result_processor()
|
self.init_batch_result_processor()
|
||||||
|
|
||||||
|
self.init_rank_consensus_checker()
|
||||||
|
|
||||||
self.is_initializing = False
|
self.is_initializing = False
|
||||||
self.init_startup_timing_summary()
|
self.init_startup_timing_summary()
|
||||||
|
|
||||||
@@ -1688,6 +1691,8 @@ class Scheduler(
|
|||||||
if self.decode_offload_manager is not None:
|
if self.decode_offload_manager is not None:
|
||||||
self.decode_offload_manager.release_host_resources()
|
self.decode_offload_manager.release_host_resources()
|
||||||
|
|
||||||
|
rank_consensus_checker.shutdown()
|
||||||
|
|
||||||
def run_event_loop(self) -> None:
|
def run_event_loop(self) -> None:
|
||||||
"""Run the scheduler's event loop.
|
"""Run the scheduler's event loop.
|
||||||
|
|
||||||
@@ -2114,6 +2119,16 @@ class Scheduler(
|
|||||||
get_running_batch=lambda: self.running_batch,
|
get_running_batch=lambda: self.running_batch,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def init_rank_consensus_checker(self) -> None:
|
||||||
|
groups = []
|
||||||
|
if self.attn_cp_group is not None and self.attn_tp_group is not None:
|
||||||
|
groups += [self.attn_cp_group, self.attn_tp_group]
|
||||||
|
else:
|
||||||
|
groups += [self.tp_group]
|
||||||
|
if self.pp_group is not None:
|
||||||
|
groups += [self.pp_group]
|
||||||
|
rank_consensus_checker.configure(groups)
|
||||||
|
|
||||||
def init_kv_events_publisher(self) -> None:
|
def init_kv_events_publisher(self) -> None:
|
||||||
self.kv_events_publisher = SchedulerKvEventsPublisher(
|
self.kv_events_publisher = SchedulerKvEventsPublisher(
|
||||||
kv_events_config=get_observability().kv_events_config,
|
kv_events_config=get_observability().kv_events_config,
|
||||||
|
|||||||
@@ -95,6 +95,7 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.mem_cache.memory_pool_host import PoolEntry
|
from sglang.srt.mem_cache.memory_pool_host import PoolEntry
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
from sglang.srt.utils.rank_consensus_checker import rank_consensus
|
||||||
|
|
||||||
T = TypeVar("T")
|
T = TypeVar("T")
|
||||||
|
|
||||||
@@ -491,6 +492,10 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
if self.host_pool_group is not None:
|
if self.host_pool_group is not None:
|
||||||
self.host_pool_group.destroy()
|
self.host_pool_group.destroy()
|
||||||
|
|
||||||
|
@rank_consensus(
|
||||||
|
same_params=["params"],
|
||||||
|
same_results=["result.full_kv_hit_length", "result.swa_host_hit_length"],
|
||||||
|
)
|
||||||
def match_prefix(self, params: MatchPrefixParams) -> MatchResult:
|
def match_prefix(self, params: MatchPrefixParams) -> MatchResult:
|
||||||
result = self.session.try_match_prefix(params)
|
result = self.session.try_match_prefix(params)
|
||||||
if result is not None:
|
if result is not None:
|
||||||
@@ -1700,6 +1705,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
operation_terminated = states[1].item() == 1
|
operation_terminated = states[1].item() == 1
|
||||||
return can_terminate or operation_terminated
|
return can_terminate or operation_terminated
|
||||||
|
|
||||||
|
@rank_consensus(same_params=True, same_results=True)
|
||||||
def check_prefetch_progress(self, req_id: str) -> bool:
|
def check_prefetch_progress(self, req_id: str) -> bool:
|
||||||
if req_id not in self.ongoing_prefetch:
|
if req_id not in self.ongoing_prefetch:
|
||||||
return True
|
return True
|
||||||
@@ -1917,6 +1923,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
return 0
|
return 0
|
||||||
return self.buffer_pipeline.staged_prefetch_swa_tokens(req_id)
|
return self.buffer_pipeline.staged_prefetch_swa_tokens(req_id)
|
||||||
|
|
||||||
|
@rank_consensus(same_params=True)
|
||||||
def release_aborted_request(self, rid: str) -> None:
|
def release_aborted_request(self, rid: str) -> None:
|
||||||
self.prefetch_loaded_tokens_by_reqid.pop(rid, None)
|
self.prefetch_loaded_tokens_by_reqid.pop(rid, None)
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -0,0 +1,432 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import functools
|
||||||
|
import hashlib
|
||||||
|
import inspect
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import queue
|
||||||
|
import threading
|
||||||
|
from typing import TYPE_CHECKING, Any, Callable, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_sync_groups: List[dist.ProcessGroup] = [] # Dedicated gloo groups (one per rank-set).
|
||||||
|
_q: Optional[queue.Queue[str]] = None
|
||||||
|
_worker_thread: Optional[threading.Thread] = None
|
||||||
|
_scheduler_thread: Optional[threading.Thread] = None
|
||||||
|
|
||||||
|
|
||||||
|
def rank_consensus(func=None, *, same_params=None, same_results=None, **kwargs):
|
||||||
|
"""
|
||||||
|
Mark a function that should be consensus in PP and TP ranks. Here consensus means,
|
||||||
|
the same order of calling, same parameters and return values optionally.
|
||||||
|
|
||||||
|
The function must be called in the scheduler thread.
|
||||||
|
|
||||||
|
Usages:
|
||||||
|
|
||||||
|
* Assert that the function is called by all ranks. The parameters or results may not be same.
|
||||||
|
@rank_consensus
|
||||||
|
def foo():
|
||||||
|
pass
|
||||||
|
|
||||||
|
* Assert that all parameters are same in all ranks.
|
||||||
|
@rank_consensus(same_params = True)
|
||||||
|
def foo(a, b):
|
||||||
|
pass
|
||||||
|
|
||||||
|
* Assert that some parameters are same in all ranks.
|
||||||
|
@rank_consensus(same_params = ["a", "c"])
|
||||||
|
def foo(a, b, c):
|
||||||
|
pass
|
||||||
|
|
||||||
|
* Assert that part of the parameters are same in all ranks.
|
||||||
|
@rank_consensus(same_params = ["a.req_id"])
|
||||||
|
def foo(a):
|
||||||
|
pass
|
||||||
|
|
||||||
|
* Assert that results are same in all ranks.
|
||||||
|
@rank_consensus(same_results = True)
|
||||||
|
def foo():
|
||||||
|
return 1
|
||||||
|
|
||||||
|
* Assert for part of the results are same.
|
||||||
|
@rank_consensus(same_results = ["result.some_field"]
|
||||||
|
def foo():
|
||||||
|
return SomeObject()
|
||||||
|
|
||||||
|
@rank_consensus(same_results = ["result.field", "len(result.field2)"]
|
||||||
|
def foo():
|
||||||
|
return SomeObject()
|
||||||
|
|
||||||
|
* Assert the function is called by all ranks and all parameters and results are the same.
|
||||||
|
@rank_consensus(same_params = True, same_results = True)
|
||||||
|
def foo():
|
||||||
|
return 1
|
||||||
|
"""
|
||||||
|
if kwargs:
|
||||||
|
raise TypeError(
|
||||||
|
f"rank_consensus() got unexpected keyword argument(s): " f"{list(kwargs)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
params_selector = _normalize_selector(same_params, "same_params")
|
||||||
|
results_selector = _normalize_selector(same_results, "same_results")
|
||||||
|
|
||||||
|
def decorator(func: Callable) -> Callable:
|
||||||
|
# This decorator function called at import time. So it should be zero runtime overhead
|
||||||
|
# when the consensus checker is disabled.
|
||||||
|
if not envs.SGLANG_ENABLE_RANK_CONSENSUS_CHECKER.get():
|
||||||
|
return func
|
||||||
|
|
||||||
|
# Unwrap static/class-method descriptors so we always operate on the
|
||||||
|
# raw function. We remember the descriptor type so we can re-wrap the
|
||||||
|
# result and the class-body descriptor protocol keeps working.
|
||||||
|
if isinstance(func, (classmethod, staticmethod)):
|
||||||
|
raw_func = func.__func__
|
||||||
|
descriptor_type = type(func)
|
||||||
|
else:
|
||||||
|
raw_func = func
|
||||||
|
descriptor_type = None
|
||||||
|
sig = inspect.signature(raw_func)
|
||||||
|
|
||||||
|
# When calling class method or object method with "same_params=True",
|
||||||
|
# skip the first "cls" or "self", as the text format for that
|
||||||
|
# may include memory addresses, which are considered divergence.
|
||||||
|
skip_name: Optional[str] = None
|
||||||
|
if _is_method_with_receiver(func) and len(sig.parameters) > 0:
|
||||||
|
skip_name = next(iter(sig.parameters))
|
||||||
|
|
||||||
|
@functools.wraps(raw_func)
|
||||||
|
def wrapper(*args: Any, **kwargs: Any) -> Any:
|
||||||
|
params_payload = "<no check>"
|
||||||
|
if params_selector is not None:
|
||||||
|
# Bind once and apply defaults so that name-based selectors work
|
||||||
|
# regardless of whether the caller passed positionally or by kw.
|
||||||
|
bound = sig.bind(*args, **kwargs)
|
||||||
|
bound.apply_defaults()
|
||||||
|
arguments = dict(bound.arguments)
|
||||||
|
params_payload = _build_payload(
|
||||||
|
"call", params_selector, arguments, skip_name
|
||||||
|
)
|
||||||
|
assert_same("%s called params=%s", raw_func.__name__, params_payload)
|
||||||
|
|
||||||
|
result = raw_func(*args, **kwargs)
|
||||||
|
|
||||||
|
result_payload = "<no check>"
|
||||||
|
if results_selector is not None:
|
||||||
|
result_scope = {"result": result}
|
||||||
|
result_payload = _build_payload(
|
||||||
|
"return",
|
||||||
|
results_selector,
|
||||||
|
result_scope,
|
||||||
|
)
|
||||||
|
assert_same("%s returns result=%s", raw_func.__name__, result_payload)
|
||||||
|
return result
|
||||||
|
|
||||||
|
# Re-wrap into the original descriptor type so class-body access
|
||||||
|
# (C.method / instance.method) still binds correctly.
|
||||||
|
if descriptor_type is staticmethod:
|
||||||
|
return staticmethod(wrapper)
|
||||||
|
if descriptor_type is classmethod:
|
||||||
|
return classmethod(wrapper)
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
if func is not None:
|
||||||
|
# Bare `@rank_consensus` form.
|
||||||
|
return decorator(func)
|
||||||
|
else:
|
||||||
|
# `@rank_consensus(same_params=True, same_results=True)` form.
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_selector(
|
||||||
|
value: None | bool | str | list[str], name: str
|
||||||
|
) -> None | bool | list[str]:
|
||||||
|
"""Normalize a selector argument to one of:
|
||||||
|
``None`` (skip), ``True`` (compare everything), or ``list[str]`` (the
|
||||||
|
expressions to evaluate). ``False`` is treated as ``None``.
|
||||||
|
"""
|
||||||
|
if value is None or value is False:
|
||||||
|
return None
|
||||||
|
if value is True:
|
||||||
|
return True
|
||||||
|
if isinstance(value, str):
|
||||||
|
return [value]
|
||||||
|
if isinstance(value, list) and all(isinstance(s, str) for s in value):
|
||||||
|
return list(value)
|
||||||
|
raise TypeError(f"{name} must be True / False / str / list[str], got {value!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def _is_method_with_receiver(func: Any) -> bool:
|
||||||
|
"""Return True iff ``func`` is a method whose first parameter is a
|
||||||
|
receiver (instance for instance-methods, class for class-methods) that
|
||||||
|
should be dropped from the ``same_params=True`` payload.
|
||||||
|
|
||||||
|
Distinguishes:
|
||||||
|
* ``staticmethod`` object -> False (no receiver)
|
||||||
|
* ``classmethod`` object -> True (receiver is the class)
|
||||||
|
* plain ``def`` defined inside a class body (``__qualname__`` has a
|
||||||
|
dot before the final segment and is not a ``<locals>`` closure) ->
|
||||||
|
True (instance method)
|
||||||
|
* anything else (module-level function, nested function, lambda) ->
|
||||||
|
False
|
||||||
|
"""
|
||||||
|
if isinstance(func, staticmethod):
|
||||||
|
return False
|
||||||
|
if isinstance(func, classmethod):
|
||||||
|
return True
|
||||||
|
if inspect.isfunction(func):
|
||||||
|
qualname = getattr(func, "__qualname__", "")
|
||||||
|
# ``C.m`` -> True; ``m`` -> False; ``outer.<locals>.m`` -> False
|
||||||
|
# (closures aren't class-body methods).
|
||||||
|
if "." in qualname and "<locals>" not in qualname:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _build_payload(
|
||||||
|
tag: str,
|
||||||
|
selector: bool | list[str],
|
||||||
|
scope: dict[str, Any],
|
||||||
|
skip_name: Optional[str] = None,
|
||||||
|
) -> str:
|
||||||
|
"""Serialize the selected values into a single comparable string.
|
||||||
|
|
||||||
|
``skip_name`` only applies to the ``True`` (whole-scope) form and is used
|
||||||
|
to drop the receiver (``self`` / ``cls``) from method payloads; explicit
|
||||||
|
``list[str]`` selectors honor exactly what the user listed.
|
||||||
|
"""
|
||||||
|
if selector is True:
|
||||||
|
# Whole scope is the payload. For the call checkpoint, the scope is
|
||||||
|
# the arguments dict; for the return checkpoint, the caller wrapped
|
||||||
|
# result into the scope, so we repr ``result`` directly.
|
||||||
|
if tag == "call":
|
||||||
|
if skip_name is not None:
|
||||||
|
scope = {k: v for k, v in scope.items() if k != skip_name}
|
||||||
|
return repr(scope)
|
||||||
|
return repr(scope["result"])
|
||||||
|
parts: list[str] = []
|
||||||
|
for expr in selector:
|
||||||
|
value = _eval_selector(expr, scope)
|
||||||
|
parts.append(f"{expr}={value!r}")
|
||||||
|
return " | ".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _eval_selector(expr: str, scope: dict[str, Any]) -> Any:
|
||||||
|
"""Evaluate a selector expression in a restricted scope.
|
||||||
|
|
||||||
|
Errors (unknown parameter name, missing attribute, bad syntax) propagate
|
||||||
|
-- they are caller bugs and must not be silently swallowed or confused
|
||||||
|
with cross-rank divergence.
|
||||||
|
"""
|
||||||
|
safe_builtins = {
|
||||||
|
"len": len,
|
||||||
|
"int": int,
|
||||||
|
"str": str,
|
||||||
|
"bool": bool,
|
||||||
|
"float": float,
|
||||||
|
"tuple": tuple,
|
||||||
|
"list": list,
|
||||||
|
"dict": dict,
|
||||||
|
"set": set,
|
||||||
|
"sorted": sorted,
|
||||||
|
"min": min,
|
||||||
|
"max": max,
|
||||||
|
"sum": sum,
|
||||||
|
}
|
||||||
|
return eval(expr, {"__builtins__": safe_builtins}, dict(scope))
|
||||||
|
|
||||||
|
|
||||||
|
def enabled() -> bool:
|
||||||
|
"""Test that the checker has been enabled and configure() is called."""
|
||||||
|
return _q is not None
|
||||||
|
|
||||||
|
|
||||||
|
def assert_same(msg_fmt: str, *args: Any) -> None:
|
||||||
|
"""Record a decision that every TP/PP rank must make identically.
|
||||||
|
|
||||||
|
Must be called from the scheduler thread. If the env var is set and the
|
||||||
|
checker is configured, an assertion guards that the caller is on the
|
||||||
|
scheduler thread recorded at configure() time — events from other threads
|
||||||
|
would interleave out of order with peer ranks and corrupt the lock-step
|
||||||
|
drain.
|
||||||
|
|
||||||
|
When the divergence checker is disabled, this is a zero-overhead no-op.
|
||||||
|
|
||||||
|
Example:
|
||||||
|
assert_same("my decision: %s %d", "foo", 100)
|
||||||
|
|
||||||
|
Prefer `@rank_consensus` over this function for code-cleanliness.
|
||||||
|
"""
|
||||||
|
if not enabled():
|
||||||
|
return
|
||||||
|
# Sanity check: only the scheduler thread is allowed to enqueue. Other
|
||||||
|
# callers would race with the worker's min-length drain and desynchronize
|
||||||
|
# ranks, since their events would not exist on peer ranks.
|
||||||
|
if threading.current_thread() is not _scheduler_thread:
|
||||||
|
raise RuntimeError("rdc.assert_same must be called from the scheduler thread")
|
||||||
|
# Format eagerly: args may reference mutable state that mutates
|
||||||
|
# between now and when the worker thread drains the queue.
|
||||||
|
_q.put(msg_fmt % args)
|
||||||
|
|
||||||
|
|
||||||
|
def configure(groups: List[GroupCoordinator]) -> None:
|
||||||
|
"""Initialize the checker. No-op if SGLANG_ENABLE_RANK_CONSENSUS_CHECKER is not set."""
|
||||||
|
global _sync_groups, _q, _worker_thread, _scheduler_thread
|
||||||
|
|
||||||
|
if not envs.SGLANG_ENABLE_RANK_CONSENSUS_CHECKER.get():
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.warning(
|
||||||
|
"Rank consensus checker is enabled. The server will suicide if rank divergence detected."
|
||||||
|
)
|
||||||
|
# Build a dedicated sync group. So our synchronization work will not affect
|
||||||
|
# the scheduler thread at all.
|
||||||
|
_sync_groups = _create_sync_groups(groups)
|
||||||
|
_q = queue.Queue()
|
||||||
|
# Assume the calling thread is the schedule thread.
|
||||||
|
# We will check assert_same() must be called by the scheduler thread.
|
||||||
|
_scheduler_thread = threading.current_thread()
|
||||||
|
_worker_thread = threading.Thread(
|
||||||
|
target=_worker_loop, name="rank_consensus_checker", daemon=True
|
||||||
|
)
|
||||||
|
_worker_thread.start()
|
||||||
|
|
||||||
|
|
||||||
|
def _create_sync_groups(
|
||||||
|
groups: List[GroupCoordinator],
|
||||||
|
) -> List[dist.ProcessGroup]:
|
||||||
|
"""Create duplicated groups, used for background thread"""
|
||||||
|
from sglang.srt.distributed.parallel_state import create_custom_parallel_group
|
||||||
|
|
||||||
|
dedicated: List[dist.ProcessGroup] = []
|
||||||
|
seen_rank_sets: set[tuple[int, ...]] = set()
|
||||||
|
for group in groups:
|
||||||
|
if group is None:
|
||||||
|
continue
|
||||||
|
# Skip single-rank groups: nothing to compare against.
|
||||||
|
if torch.distributed.get_world_size(group=group.cpu_group) == 1:
|
||||||
|
continue
|
||||||
|
group_ranks = tuple(torch.distributed.get_process_group_ranks(group.cpu_group))
|
||||||
|
if group_ranks in seen_rank_sets:
|
||||||
|
continue
|
||||||
|
seen_rank_sets.add(group_ranks)
|
||||||
|
pg = create_custom_parallel_group(group_ranks=list(group_ranks), backend="gloo")
|
||||||
|
if pg is not None:
|
||||||
|
dedicated.append(pg)
|
||||||
|
return dedicated
|
||||||
|
|
||||||
|
|
||||||
|
def _destroy_dedicated_groups() -> None:
|
||||||
|
for pg in _sync_groups:
|
||||||
|
try:
|
||||||
|
torch.distributed.destroy_process_group(pg)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def shutdown() -> None:
|
||||||
|
"""Flush the queue, stop the worker thread, and disable assert_same."""
|
||||||
|
global _q, _worker_thread, _sync_groups, _scheduler_thread
|
||||||
|
|
||||||
|
q = _q
|
||||||
|
if q is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Put a sentinel value to wake the worker if it is blocked on _q.get().
|
||||||
|
q.put(None)
|
||||||
|
if _worker_thread is not None:
|
||||||
|
_worker_thread.join()
|
||||||
|
_worker_thread = None
|
||||||
|
# Tear down the dedicated gloo groups BEFORE clearing _groups so the
|
||||||
|
# destroy helper can see them. Worker thread is already joined, so there
|
||||||
|
# is no concurrent all_reduce on these groups.
|
||||||
|
_destroy_dedicated_groups()
|
||||||
|
_q = None
|
||||||
|
_sync_groups = []
|
||||||
|
_scheduler_thread = None
|
||||||
|
|
||||||
|
|
||||||
|
def _worker_loop() -> None:
|
||||||
|
"""Consume events in lock-step with peer ranks via gloo all-reduce.
|
||||||
|
|
||||||
|
Each iteration:
|
||||||
|
1. Determine the items available in _q.
|
||||||
|
2. Drain exactly the minimum number of items in all ranks.
|
||||||
|
3. Compare all events are identical across ranks.
|
||||||
|
"""
|
||||||
|
while _q is not None:
|
||||||
|
# Drain first. Block waiting for the first event.
|
||||||
|
first = _q.get()
|
||||||
|
if first is None:
|
||||||
|
# shutdown() is called.
|
||||||
|
return
|
||||||
|
|
||||||
|
# Drain more whenever available.
|
||||||
|
# Every rank should drain the same number.
|
||||||
|
count = _all_reduce_min_int(_q.qsize())
|
||||||
|
events: List[str] = [first]
|
||||||
|
shutdown_signaled = False
|
||||||
|
for _ in range(count):
|
||||||
|
event = _q.get()
|
||||||
|
if event is None:
|
||||||
|
# shutdown() sentinel arrived mid-batch: stop draining but
|
||||||
|
# still check the events we already hold — they are real
|
||||||
|
# decisions every rank must agree on. Then exit, since the
|
||||||
|
# sentinel means shutdown() is waiting on worker_thread.join().
|
||||||
|
shutdown_signaled = True
|
||||||
|
break
|
||||||
|
events.append(event)
|
||||||
|
|
||||||
|
# Cross-rank check.
|
||||||
|
_check_for_consensus(events)
|
||||||
|
|
||||||
|
if shutdown_signaled:
|
||||||
|
return
|
||||||
|
|
||||||
|
|
||||||
|
def _all_reduce_min_int(value: int) -> int:
|
||||||
|
"""Reduce `value` to its global minimum across every configured group."""
|
||||||
|
tensor = torch.tensor([value], dtype=torch.int64)
|
||||||
|
for group in _sync_groups:
|
||||||
|
dist.all_reduce(tensor, op=dist.ReduceOp.MIN, group=group)
|
||||||
|
return int(tensor.item())
|
||||||
|
|
||||||
|
|
||||||
|
def _check_for_consensus(events: list[str]) -> None:
|
||||||
|
# Compute sha1 of concatenation of all msgs.
|
||||||
|
hasher = hashlib.sha1()
|
||||||
|
for msg in events:
|
||||||
|
hasher.update(msg.encode("utf-8"))
|
||||||
|
|
||||||
|
# Determine if some rank has a different value.
|
||||||
|
value_bytes = hasher.digest()
|
||||||
|
min_value = torch.tensor(list(hasher.digest()), dtype=torch.uint8)
|
||||||
|
max_value = min_value.clone()
|
||||||
|
for group in _sync_groups:
|
||||||
|
dist.all_reduce(min_value, op=dist.ReduceOp.MIN, group=group)
|
||||||
|
dist.all_reduce(max_value, op=dist.ReduceOp.MAX, group=group)
|
||||||
|
if not torch.equal(min_value, max_value):
|
||||||
|
# When divergence, all rank should output the following log.
|
||||||
|
logger.critical(
|
||||||
|
f"Found rank divergence for {len(events)} events(s)! local hash: {value_bytes.hex()}, events = {events}"
|
||||||
|
)
|
||||||
|
for handler in logger.handlers:
|
||||||
|
handler.flush()
|
||||||
|
|
||||||
|
# os._exit instead of sys.exit: this runs in a background thread, where
|
||||||
|
# SystemExit would only kill the thread, not the process. os._exit tears
|
||||||
|
# down the whole scheduler process so a TP/PP mismatch can never
|
||||||
|
# silently keep serving.
|
||||||
|
os._exit(1)
|
||||||
|
|
||||||
|
logger.debug(f"Consensus check passed for {len(events)} event(s).")
|
||||||
@@ -0,0 +1,586 @@
|
|||||||
|
import os
|
||||||
|
import queue
|
||||||
|
import threading
|
||||||
|
import traceback
|
||||||
|
import unittest
|
||||||
|
from multiprocessing import Process
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch.distributed as dist
|
||||||
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
|
from sglang.srt.distributed import parallel_state as ps
|
||||||
|
from sglang.srt.distributed.parallel_state import (
|
||||||
|
get_pp_group,
|
||||||
|
get_tp_group,
|
||||||
|
init_distributed_environment,
|
||||||
|
initialize_model_parallel,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils.rank_consensus_checker import (
|
||||||
|
assert_same,
|
||||||
|
configure,
|
||||||
|
rank_consensus,
|
||||||
|
shutdown,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase, find_available_port
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=30, suite="base-b-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def run_distributed_test(
|
||||||
|
rank: int,
|
||||||
|
world_size: int,
|
||||||
|
pp_size: int,
|
||||||
|
tp_size: int,
|
||||||
|
master_port: int,
|
||||||
|
fn,
|
||||||
|
) -> None:
|
||||||
|
"""Child-process entry point: set up gloo, then run fn.
|
||||||
|
|
||||||
|
Exit codes:
|
||||||
|
* 0 -> fn finished cleanly
|
||||||
|
* 1 -> rdc detected divergence and called os._exit(1) from its worker
|
||||||
|
* 2 -> fn raised (test setup/scenario bug)
|
||||||
|
"""
|
||||||
|
# CUDA_VISIBLE_DEVICES is set to "99" (a non-existent device) by the parent
|
||||||
|
# in _spawn() before this process starts, so by the time the test module
|
||||||
|
# (and torch) is re-imported here, is_cuda_alike() returns False and
|
||||||
|
# GroupCoordinator picks device="cpu". That keeps this test CPU-only and
|
||||||
|
# lets world_size exceed the host's physical GPU count.
|
||||||
|
|
||||||
|
# The CUDA-only communicators (pynccl, custom allreduce) cannot be built
|
||||||
|
# without a GPU -- PyNcclCommunicator calls torch.cuda.device(device).
|
||||||
|
# initialize_model_parallel has no flag to disable pynccl, so patch
|
||||||
|
# init_model_parallel_group to force use_pynccl=False (and clear the
|
||||||
|
# module-level custom-allreduce default via its public setter). patch.object
|
||||||
|
# auto-restores on exit, including the os._exit(2) path below.
|
||||||
|
ps.set_custom_all_reduce(False)
|
||||||
|
|
||||||
|
def _cpu_init_model_parallel_group(
|
||||||
|
*args, _orig=ps.init_model_parallel_group, **kwargs
|
||||||
|
):
|
||||||
|
kwargs.setdefault("use_pynccl", False)
|
||||||
|
kwargs.setdefault("use_custom_allreduce", False)
|
||||||
|
return _orig(*args, **kwargs)
|
||||||
|
|
||||||
|
with patch.object(ps, "init_model_parallel_group", _cpu_init_model_parallel_group):
|
||||||
|
try:
|
||||||
|
os.environ["RANK"] = str(rank)
|
||||||
|
os.environ["WORLD_SIZE"] = str(world_size)
|
||||||
|
os.environ["MASTER_ADDR"] = "localhost"
|
||||||
|
os.environ["MASTER_PORT"] = str(master_port)
|
||||||
|
os.environ["LOCAL_SIZE"] = str(world_size)
|
||||||
|
|
||||||
|
init_distributed_environment(
|
||||||
|
world_size=world_size,
|
||||||
|
rank=rank,
|
||||||
|
distributed_init_method="env://",
|
||||||
|
local_rank=rank,
|
||||||
|
backend="gloo",
|
||||||
|
)
|
||||||
|
|
||||||
|
initialize_model_parallel(
|
||||||
|
tensor_model_parallel_size=tp_size,
|
||||||
|
pipeline_model_parallel_size=pp_size,
|
||||||
|
backend="gloo",
|
||||||
|
)
|
||||||
|
|
||||||
|
fn()
|
||||||
|
except Exception as e:
|
||||||
|
print(f"subprocess[{rank=}] has error: {e}", flush=True)
|
||||||
|
traceback.print_exc()
|
||||||
|
os._exit(2)
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
if dist.is_initialized():
|
||||||
|
dist.destroy_process_group()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyClass:
|
||||||
|
def __init__(self, a: int = None, b: int = None):
|
||||||
|
self.a = a
|
||||||
|
self.b = b
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
return f"DummyClass(a={self.a}, b={self.b})"
|
||||||
|
|
||||||
|
|
||||||
|
class _MethodHost:
|
||||||
|
@rank_consensus(same_params=True)
|
||||||
|
def instance_method(obj, a, b):
|
||||||
|
return a + b
|
||||||
|
|
||||||
|
@rank_consensus(same_params=True)
|
||||||
|
@classmethod
|
||||||
|
def class_method(klass, a):
|
||||||
|
return a + 1
|
||||||
|
|
||||||
|
@rank_consensus(same_params=True)
|
||||||
|
@staticmethod
|
||||||
|
def static_method(a, b):
|
||||||
|
return a * b
|
||||||
|
|
||||||
|
|
||||||
|
class RankConsensusCheckerTestCase(CustomTestCase):
|
||||||
|
def _spawn(self, fn, pp_size: int = 1, tp_size: int = 1, enable_env: bool = True):
|
||||||
|
"""Run fn in world_size spawned gloo children. Returns True iff every
|
||||||
|
child exited with code 0. A detected divergence makes rdc call
|
||||||
|
os._exit(1) from its worker thread; an exception inside fn makes
|
||||||
|
run_distributed_test call os._exit(2). Either way _spawn returns
|
||||||
|
False for that child."""
|
||||||
|
mp.set_start_method("spawn", force=True)
|
||||||
|
master_port = find_available_port(23456)
|
||||||
|
|
||||||
|
old_env = os.getenv("SGLANG_ENABLE_RANK_CONSENSUS_CHECKER")
|
||||||
|
os.environ["SGLANG_ENABLE_RANK_CONSENSUS_CHECKER"] = str(enable_env)
|
||||||
|
|
||||||
|
world_size = pp_size * tp_size
|
||||||
|
processes = []
|
||||||
|
for rank in range(world_size):
|
||||||
|
p = Process(
|
||||||
|
target=run_distributed_test,
|
||||||
|
kwargs=dict(
|
||||||
|
rank=rank,
|
||||||
|
world_size=world_size,
|
||||||
|
pp_size=pp_size,
|
||||||
|
tp_size=tp_size,
|
||||||
|
master_port=master_port,
|
||||||
|
fn=fn,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
p.start()
|
||||||
|
processes.append(p)
|
||||||
|
|
||||||
|
for p in processes:
|
||||||
|
p.join()
|
||||||
|
|
||||||
|
if old_env is None:
|
||||||
|
os.environ.pop("SGLANG_ENABLE_RANK_CONSENSUS_CHECKER")
|
||||||
|
else:
|
||||||
|
os.environ["SGLANG_ENABLE_RANK_CONSENSUS_CHECKER"] = old_env
|
||||||
|
|
||||||
|
return all(p.exitcode == 0 for p in processes)
|
||||||
|
|
||||||
|
|
||||||
|
class TestAssertSame(RankConsensusCheckerTestCase):
|
||||||
|
@staticmethod
|
||||||
|
def same_fn():
|
||||||
|
configure([get_tp_group()])
|
||||||
|
assert_same("same %d", 10)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_same(self):
|
||||||
|
"""Same args on every rank -> no divergence, clean exit."""
|
||||||
|
self.assertTrue(self._spawn(TestAssertSame.same_fn, tp_size=2))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def divergence_fn():
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
assert_same("diverge %d", tp_group.rank_in_group)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_divergence(self):
|
||||||
|
"""Different args on different ranks -> rdc calls os._exit(1) -> child
|
||||||
|
exit code is 1 -> _spawn returns False."""
|
||||||
|
self.assertFalse(self._spawn(TestAssertSame.divergence_fn, tp_size=2))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def divergent_multi_group_fn():
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
pp_group = get_pp_group()
|
||||||
|
configure([tp_group, pp_group])
|
||||||
|
assert_same("diverge %d", tp_group.rank_in_group)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_divergence_detected_multi_group(self):
|
||||||
|
"""Passing the same group twice must still surface the divergence."""
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(
|
||||||
|
TestAssertSame.divergent_multi_group_fn,
|
||||||
|
pp_size=2,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def wrong_thread_fn():
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
|
||||||
|
err_box: queue.Queue = queue.Queue()
|
||||||
|
|
||||||
|
def _other_thread():
|
||||||
|
try:
|
||||||
|
assert_same("from other thread")
|
||||||
|
err_box.put(None)
|
||||||
|
except Exception as e: # noqa: BLE001
|
||||||
|
err_box.put(e)
|
||||||
|
|
||||||
|
t = threading.Thread(target=_other_thread)
|
||||||
|
t.start()
|
||||||
|
t.join()
|
||||||
|
|
||||||
|
err = err_box.get()
|
||||||
|
shutdown()
|
||||||
|
assert isinstance(
|
||||||
|
err, RuntimeError
|
||||||
|
), f"Expected RuntimeError from stray-thread assert_same, got {err!r}"
|
||||||
|
|
||||||
|
def test_assert_same_rejects_non_scheduler_thread(self):
|
||||||
|
"""Check that assert_same() must be called in the scheduler thread. Otherwise report error."""
|
||||||
|
self.assertTrue(self._spawn(TestAssertSame.wrong_thread_fn, tp_size=2))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def disabled_fn():
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
assert_same("diverge %d", tp_group.rank_in_group)
|
||||||
|
|
||||||
|
def test_disabled_is_noop(self):
|
||||||
|
"""Test that when SGLANG_ENABLE_RANK_CONSENSUS_CHECKER=false, assert_same is no-op."""
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(TestAssertSame.disabled_fn, tp_size=2, enable_env=False)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRankConsensusDecorator(RankConsensusCheckerTestCase):
|
||||||
|
@staticmethod
|
||||||
|
def consensus_bare_diverge_fn():
|
||||||
|
@rank_consensus
|
||||||
|
def foo(a: int) -> int:
|
||||||
|
return a
|
||||||
|
|
||||||
|
# Bare decorator only checks "was called", not args; even with diverging
|
||||||
|
# args this must exit clean (no rank divergence).
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
foo(tp_group.rank_in_group)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_bare_decorator_clean_with_diverging_args(self):
|
||||||
|
"""Bare decorator only checks that every rank calls the function;
|
||||||
|
diverging args must NOT be flagged."""
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(TestRankConsensusDecorator.consensus_bare_diverge_fn, tp_size=2)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_all_params_same_fn():
|
||||||
|
@rank_consensus(same_params=True)
|
||||||
|
def foo(a: int, b: int) -> int:
|
||||||
|
return a + b
|
||||||
|
|
||||||
|
configure([get_tp_group()])
|
||||||
|
foo(1, 2)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_all_params_same(self):
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_all_params_same_fn, tp_size=2
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_all_params_diverge_fn():
|
||||||
|
@rank_consensus(same_params=True)
|
||||||
|
def foo(a, b):
|
||||||
|
return a + b
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
# The second argument differs on rank. Expect divergence.
|
||||||
|
foo(1, tp_group.rank_in_group)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_all_params_diverge(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_all_params_diverge_fn, tp_size=2
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_named_params_same_fn():
|
||||||
|
@rank_consensus(same_params=["a", "c"])
|
||||||
|
def foo(a: int, b: int, c: int) -> int:
|
||||||
|
return a + b + c
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
# b diverges but is NOT in the selector list. Expect good.
|
||||||
|
foo(1, tp_group.rank_in_group, 3)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_named_params_ignores_unselected_divergence(self):
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_named_params_same_fn, tp_size=2
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_named_params_diverge_fn():
|
||||||
|
@rank_consensus(same_params=["a", "c"])
|
||||||
|
def foo(a: int, b: int, c: int) -> int:
|
||||||
|
return a + b + c
|
||||||
|
|
||||||
|
# c diverges and IS in the selector list. Expect divergence.
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
foo(1, 2, tp_group.rank_in_group)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_named_params_flags_selected_divergence(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_named_params_diverge_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_dotted_param_same_fn():
|
||||||
|
@rank_consensus(same_params=["_a.a"])
|
||||||
|
def foo(_a: _DummyClass) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
dummy = _DummyClass(a=10, b=tp_group.rank_in_group)
|
||||||
|
foo(dummy)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_dotted_param_same(self):
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_dotted_param_same_fn, tp_size=2
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_dotted_param_diverge_fn():
|
||||||
|
@rank_consensus(same_params=["_a.a"])
|
||||||
|
def foo(_a: _DummyClass) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
dummy = _DummyClass(a=tp_group.rank_in_group, b=10)
|
||||||
|
foo(dummy)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_dotted_param_diverge(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_dotted_param_diverge_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_full_result_same_fn():
|
||||||
|
@rank_consensus(same_results=True)
|
||||||
|
def foo(value: int) -> _DummyClass:
|
||||||
|
return _DummyClass(a=value, b=value * 2)
|
||||||
|
|
||||||
|
configure([get_tp_group()])
|
||||||
|
foo(5)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_full_result_same(self):
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_full_result_same_fn, tp_size=2
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_full_result_diverge_fn():
|
||||||
|
@rank_consensus(same_results=True)
|
||||||
|
def foo(value: int) -> _DummyClass:
|
||||||
|
return _DummyClass(a=value, b=value * 2)
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
foo(tp_group.rank_in_group)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_full_result_diverge(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_full_result_diverge_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_partial_result_same_fn():
|
||||||
|
@rank_consensus(same_results=["result.x", "len(result.y)"])
|
||||||
|
def foo(x, y_list):
|
||||||
|
class _R:
|
||||||
|
pass
|
||||||
|
|
||||||
|
r = _R()
|
||||||
|
r.x = x
|
||||||
|
r.y = y_list
|
||||||
|
return r
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
# x and len(y) both equal across ranks; y contents differ but are not selected. Expect good.
|
||||||
|
foo(x=3, y_list=[tp_group.rank_in_group] * 4)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_partial_result_same(self):
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_partial_result_same_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_partial_result_diverge_fn():
|
||||||
|
@rank_consensus(same_results=["result.x", "len(result.y)"])
|
||||||
|
def foo(x, y_list):
|
||||||
|
class _R:
|
||||||
|
pass
|
||||||
|
|
||||||
|
r = _R()
|
||||||
|
r.x = x
|
||||||
|
r.y = y_list
|
||||||
|
return r
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
# x diverges and IS selected. Expect divergence.
|
||||||
|
foo(x=tp_group.rank_in_group, y_list=[1, 2, 3])
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_partial_result_diverge(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_partial_result_diverge_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_both_same_fn():
|
||||||
|
@rank_consensus(same_params=True, same_results=True)
|
||||||
|
def foo(a: int) -> int:
|
||||||
|
return a * 2
|
||||||
|
|
||||||
|
configure([get_tp_group()])
|
||||||
|
foo(7)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_both_same(self):
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(TestRankConsensusDecorator.consensus_both_same_fn, tp_size=2)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_both_diverge_fn():
|
||||||
|
@rank_consensus(same_params=True, same_results=True)
|
||||||
|
def foo(a: int) -> int:
|
||||||
|
return a * 2
|
||||||
|
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
foo(tp_group.rank_in_group)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_both_diverge(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(TestRankConsensusDecorator.consensus_both_diverge_fn, tp_size=2)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_instance_method_same_fn():
|
||||||
|
configure([get_tp_group()])
|
||||||
|
_MethodHost().instance_method(1, 2)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_instance_method_receiver_dropped(self):
|
||||||
|
# Two ranks build two different _MethodHost instances; without the
|
||||||
|
# receiver-skip the per-rank address would diverge. Clean exit
|
||||||
|
# confirms the receiver is dropped.
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_instance_method_same_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_class_method_same_fn():
|
||||||
|
configure([get_tp_group()])
|
||||||
|
_MethodHost.class_method(5)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_class_method_receiver_dropped(self):
|
||||||
|
# First param is named ``klass`` (not cls); detection must still work.
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_class_method_same_fn, tp_size=2
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_class_method_via_instance_same_fn():
|
||||||
|
configure([get_tp_group()])
|
||||||
|
_MethodHost().class_method(5)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_class_method_via_instance_receiver_dropped(self):
|
||||||
|
# Accessing the classmethod through an instance still binds the class
|
||||||
|
# as the receiver; verify it is still dropped.
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_class_method_via_instance_same_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_static_method_same_fn():
|
||||||
|
configure([get_tp_group()])
|
||||||
|
_MethodHost.static_method(3, 4)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_static_method_no_receiver(self):
|
||||||
|
# Static method: no receiver, equal args -> clean.
|
||||||
|
self.assertTrue(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_static_method_same_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def consensus_static_method_diverge_fn():
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
configure([tp_group])
|
||||||
|
# Static method: no receiver to drop, so a rank-dependent arg diverges.
|
||||||
|
_MethodHost.static_method(tp_group.rank_in_group, 4)
|
||||||
|
shutdown()
|
||||||
|
|
||||||
|
def test_static_method_flags_diverging_arg(self):
|
||||||
|
# Static method: no receiver to drop, so a rank-dependent arg must
|
||||||
|
# still be flagged. Confirms we did not over-skip for static methods.
|
||||||
|
self.assertFalse(
|
||||||
|
self._spawn(
|
||||||
|
TestRankConsensusDecorator.consensus_static_method_diverge_fn,
|
||||||
|
tp_size=2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user