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:
Lianmin Zheng
2026-05-31 20:50:29 -07:00
committed by GitHub
co-authored by Jaewon
parent 11411aa49d
commit a779791b3f
4 changed files with 59 additions and 9 deletions
+1
View File
@@ -244,6 +244,7 @@ class Envs:
SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1)
SGLANG_SIMULATE_ACC_METHOD = EnvStr("match-expected")
SGLANG_SIMULATE_UNIFORM_EXPERTS = EnvBool(False)
SGLANG_SIMULATE_ROUND_ROBIN_EXPERTS = EnvBool(False)
SGLANG_TORCH_PROFILER_DIR = EnvStr("/tmp")
SGLANG_OTLP_EXPORTER_SCHEDULE_DELAY_MILLIS = EnvInt(500)
SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE = EnvInt(64)
+50 -6
View File
@@ -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 ---------------------------------------
@@ -1533,17 +1552,42 @@ def select_experts(
renormalize=renormalize,
)
if envs.SGLANG_SIMULATE_UNIFORM_EXPERTS.get():
# Benchmark-only: override gating with uniform round-robin expert assignment
simulate_uniform_experts = envs.SGLANG_SIMULATE_UNIFORM_EXPERTS.get()
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.
num_tokens, k = topk_ids.shape
num_experts = router_logits.shape[1]
offsets = torch.randint(0, num_experts, (num_tokens, 1), device=topk_ids.device)
steps = torch.arange(k, device=topk_ids.device).unsqueeze(0)
topk_ids = ((offsets + steps * (num_experts // k)) % num_experts).to(
topk_ids.dtype
if k > 0:
offsets = torch.randint(
0, num_experts, (num_tokens, 1), device=topk_ids.device
)
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,
)
if k > 0:
topk_weights = torch.full_like(topk_weights, 1.0 / k)
topk_ids, topk_weights = _post_process_topk_ids(
topk_ids=topk_ids,
+5
View File
@@ -1208,6 +1208,11 @@ def kill_process_tree(
`parent_pid == os.getpid()` branch calls `sys.exit(0)` and cannot wait
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:
parent_pid = os.getpid()
include_parent = False
@@ -190,7 +190,7 @@ class BenchArgs:
"--dataset-name",
type=str,
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.",
)
parser.add_argument(
@@ -517,7 +517,7 @@ def run_one_case(
_flush_cache_with_retry(url, "/flush_cache")
# 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:
raise ValueError(
f"Unsupported dataset for batch benchmark: {dataset_name}. "