Add random-ids dataset, round-robin expert simulation, and kill_process_tree logging (#26862)
Co-authored-by: Jaewon <52840625+jaewonlee-fb@users.noreply.github.com>
This commit is contained in:
@@ -244,6 +244,7 @@ class Envs:
|
|||||||
SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1)
|
SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1)
|
||||||
SGLANG_SIMULATE_ACC_METHOD = EnvStr("match-expected")
|
SGLANG_SIMULATE_ACC_METHOD = EnvStr("match-expected")
|
||||||
SGLANG_SIMULATE_UNIFORM_EXPERTS = EnvBool(False)
|
SGLANG_SIMULATE_UNIFORM_EXPERTS = EnvBool(False)
|
||||||
|
SGLANG_SIMULATE_ROUND_ROBIN_EXPERTS = EnvBool(False)
|
||||||
SGLANG_TORCH_PROFILER_DIR = EnvStr("/tmp")
|
SGLANG_TORCH_PROFILER_DIR = EnvStr("/tmp")
|
||||||
SGLANG_OTLP_EXPORTER_SCHEDULE_DELAY_MILLIS = EnvInt(500)
|
SGLANG_OTLP_EXPORTER_SCHEDULE_DELAY_MILLIS = EnvInt(500)
|
||||||
SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE = EnvInt(64)
|
SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE = EnvInt(64)
|
||||||
|
|||||||
@@ -299,6 +299,25 @@ class BypassedTopKOutput(NamedTuple):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_round_robin_expert_ids(
|
||||||
|
num_tokens: int,
|
||||||
|
topk: int,
|
||||||
|
num_experts: int,
|
||||||
|
*,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
layer_id: Optional[int] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if topk == 0:
|
||||||
|
return torch.empty((num_tokens, 0), device=device, dtype=dtype)
|
||||||
|
|
||||||
|
step = max(num_experts // topk, 1)
|
||||||
|
layer_offset = 0 if layer_id is None else layer_id
|
||||||
|
offsets = torch.arange(num_tokens, device=device, dtype=dtype).unsqueeze(1)
|
||||||
|
steps = torch.arange(topk, device=device, dtype=dtype).unsqueeze(0) * step
|
||||||
|
return (offsets + layer_offset + steps) % num_experts
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------- TopK ---------------------------------------
|
# -------------------------------- TopK ---------------------------------------
|
||||||
|
|
||||||
|
|
||||||
@@ -1533,17 +1552,42 @@ def select_experts(
|
|||||||
renormalize=renormalize,
|
renormalize=renormalize,
|
||||||
)
|
)
|
||||||
|
|
||||||
if envs.SGLANG_SIMULATE_UNIFORM_EXPERTS.get():
|
simulate_uniform_experts = envs.SGLANG_SIMULATE_UNIFORM_EXPERTS.get()
|
||||||
# Benchmark-only: override gating with uniform round-robin expert assignment
|
simulate_round_robin_experts = envs.SGLANG_SIMULATE_ROUND_ROBIN_EXPERTS.get()
|
||||||
|
if simulate_uniform_experts and simulate_round_robin_experts:
|
||||||
|
raise ValueError(
|
||||||
|
"SGLANG_SIMULATE_UNIFORM_EXPERTS and "
|
||||||
|
"SGLANG_SIMULATE_ROUND_ROBIN_EXPERTS are mutually exclusive"
|
||||||
|
)
|
||||||
|
|
||||||
|
if simulate_uniform_experts:
|
||||||
|
# Benchmark-only: override gating with random-offset uniform expert assignment
|
||||||
# to avoid expert imbalance from dummy/random weights. Do NOT use in production.
|
# to avoid expert imbalance from dummy/random weights. Do NOT use in production.
|
||||||
num_tokens, k = topk_ids.shape
|
num_tokens, k = topk_ids.shape
|
||||||
num_experts = router_logits.shape[1]
|
num_experts = router_logits.shape[1]
|
||||||
offsets = torch.randint(0, num_experts, (num_tokens, 1), device=topk_ids.device)
|
if k > 0:
|
||||||
steps = torch.arange(k, device=topk_ids.device).unsqueeze(0)
|
offsets = torch.randint(
|
||||||
topk_ids = ((offsets + steps * (num_experts // k)) % num_experts).to(
|
0, num_experts, (num_tokens, 1), device=topk_ids.device
|
||||||
topk_ids.dtype
|
)
|
||||||
|
steps = torch.arange(k, device=topk_ids.device).unsqueeze(0)
|
||||||
|
step = max(num_experts // k, 1)
|
||||||
|
topk_ids = ((offsets + steps * step) % num_experts).to(topk_ids.dtype)
|
||||||
|
topk_weights = torch.ones_like(topk_weights) / k
|
||||||
|
elif simulate_round_robin_experts:
|
||||||
|
# Benchmark-only: override gating with deterministic expert assignment
|
||||||
|
# to avoid routing noise from dummy/random weights. Do NOT use in production.
|
||||||
|
num_tokens, k = topk_ids.shape
|
||||||
|
num_experts = router_logits.shape[1]
|
||||||
|
topk_ids = _make_round_robin_expert_ids(
|
||||||
|
num_tokens,
|
||||||
|
k,
|
||||||
|
num_experts,
|
||||||
|
device=topk_ids.device,
|
||||||
|
dtype=topk_ids.dtype,
|
||||||
|
layer_id=layer_id,
|
||||||
)
|
)
|
||||||
topk_weights = torch.ones_like(topk_weights) / k
|
if k > 0:
|
||||||
|
topk_weights = torch.full_like(topk_weights, 1.0 / k)
|
||||||
|
|
||||||
topk_ids, topk_weights = _post_process_topk_ids(
|
topk_ids, topk_weights = _post_process_topk_ids(
|
||||||
topk_ids=topk_ids,
|
topk_ids=topk_ids,
|
||||||
|
|||||||
@@ -1208,6 +1208,11 @@ def kill_process_tree(
|
|||||||
`parent_pid == os.getpid()` branch calls `sys.exit(0)` and cannot wait
|
`parent_pid == os.getpid()` branch calls `sys.exit(0)` and cannot wait
|
||||||
for itself -- use `include_parent=False` if child reap must finish first.
|
for itself -- use `include_parent=False` if child reap must finish first.
|
||||||
"""
|
"""
|
||||||
|
logger.info(
|
||||||
|
f"kill_process_tree called: parent_pid={parent_pid}, "
|
||||||
|
f"include_parent={include_parent}, pid={os.getpid()}"
|
||||||
|
)
|
||||||
|
|
||||||
if parent_pid is None:
|
if parent_pid is None:
|
||||||
parent_pid = os.getpid()
|
parent_pid = os.getpid()
|
||||||
include_parent = False
|
include_parent = False
|
||||||
|
|||||||
@@ -190,7 +190,7 @@ class BenchArgs:
|
|||||||
"--dataset-name",
|
"--dataset-name",
|
||||||
type=str,
|
type=str,
|
||||||
default=BenchArgs.dataset_name,
|
default=BenchArgs.dataset_name,
|
||||||
choices=["mmmu", "random", "generated-shared-prefix"],
|
choices=["mmmu", "random", "random-ids", "generated-shared-prefix"],
|
||||||
help="Name of the dataset to benchmark on.",
|
help="Name of the dataset to benchmark on.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -517,7 +517,7 @@ def run_one_case(
|
|||||||
_flush_cache_with_retry(url, "/flush_cache")
|
_flush_cache_with_retry(url, "/flush_cache")
|
||||||
|
|
||||||
# Load input token ids via bench_serving.get_dataset
|
# Load input token ids via bench_serving.get_dataset
|
||||||
supported_datasets = ("random", "mmmu", "generated-shared-prefix")
|
supported_datasets = ("random", "random-ids", "mmmu", "generated-shared-prefix")
|
||||||
if dataset_name not in supported_datasets:
|
if dataset_name not in supported_datasets:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported dataset for batch benchmark: {dataset_name}. "
|
f"Unsupported dataset for batch benchmark: {dataset_name}. "
|
||||||
|
|||||||
Reference in New Issue
Block a user