Simplify routed experts test and move base64 encoding to tokenizer manager (#21634)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Lianmin Zheng
2026-03-29 12:44:01 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 2acdda1d85
commit 1d9c8e8c9e
6 changed files with 35 additions and 45 deletions
@@ -21,7 +21,6 @@ from collections import OrderedDict, defaultdict
from typing import Dict, List, Optional, Tuple, Union from typing import Dict, List, Optional, Tuple, Union
import psutil import psutil
import pybase64
import setproctitle import setproctitle
import zmq import zmq
@@ -319,21 +318,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
return output_strs return output_strs
def _extract_routed_experts(
self, recv_obj: BatchTokenIDOutput
) -> list[str | None] | None:
routed_experts = None
if recv_obj.routed_experts is not None:
routed_experts = [
(
pybase64.b64encode(routed_experts.numpy().tobytes()).decode("utf-8")
if routed_experts is not None
else None
)
for routed_experts in recv_obj.routed_experts
]
return routed_experts
def handle_batch_token_id_out(self, recv_obj: BatchTokenIDOutput): def handle_batch_token_id_out(self, recv_obj: BatchTokenIDOutput):
# If handling idle batch, set output_strs to []. # If handling idle batch, set output_strs to [].
output_strs = ( output_strs = (
@@ -341,8 +325,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
if len(recv_obj.rids) > 0 if len(recv_obj.rids) > 0
else [] else []
) )
routed_experts = self._extract_routed_experts(recv_obj)
return BatchStrOutput( return BatchStrOutput(
rids=recv_obj.rids, rids=recv_obj.rids,
http_worker_ipcs=recv_obj.http_worker_ipcs, http_worker_ipcs=recv_obj.http_worker_ipcs,
@@ -370,7 +352,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
output_token_ids_logprobs_idx=recv_obj.output_token_ids_logprobs_idx, output_token_ids_logprobs_idx=recv_obj.output_token_ids_logprobs_idx,
output_token_entropy_val=recv_obj.output_token_entropy_val, output_token_entropy_val=recv_obj.output_token_entropy_val,
output_hidden_states=recv_obj.output_hidden_states, output_hidden_states=recv_obj.output_hidden_states,
routed_experts=routed_experts, routed_experts=recv_obj.routed_experts,
customized_info=recv_obj.customized_info, customized_info=recv_obj.customized_info,
placeholder_tokens_idx=None, placeholder_tokens_idx=None,
placeholder_tokens_val=None, placeholder_tokens_val=None,
@@ -32,6 +32,7 @@ from http import HTTPStatus
from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union
import fastapi import fastapi
import pybase64
import uvloop import uvloop
import zmq import zmq
import zmq.asyncio import zmq.asyncio
@@ -1597,7 +1598,11 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if getattr(recv_obj, "output_hidden_states", None): if getattr(recv_obj, "output_hidden_states", None):
meta_info["hidden_states"] = recv_obj.output_hidden_states[i] meta_info["hidden_states"] = recv_obj.output_hidden_states[i]
if getattr(recv_obj, "routed_experts", None): if getattr(recv_obj, "routed_experts", None):
meta_info["routed_experts"] = recv_obj.routed_experts[i] routed_experts_tensor = recv_obj.routed_experts[i]
if routed_experts_tensor is not None:
meta_info["routed_experts"] = pybase64.b64encode(
routed_experts_tensor.numpy().tobytes()
).decode("utf-8")
if getattr(recv_obj, "customized_info", None): if getattr(recv_obj, "customized_info", None):
for k, v in recv_obj.customized_info.items(): for k, v in recv_obj.customized_info.items():
meta_info[k] = v[i] meta_info[k] = v[i]
+1 -1
View File
@@ -155,7 +155,7 @@ def _is_numa_available() -> bool:
return False return False
if not shutil.which("numactl") and envs.SGLANG_NUMA_BIND_V2.get(): if not shutil.which("numactl") and envs.SGLANG_NUMA_BIND_V2.get():
logger.warning( logger.debug(
"numactl command not found, skipping NUMA node configuration for GPU. Install numactl (e.g., apt-get install numactl) to enable automatic NUMA binding." "numactl command not found, skipping NUMA node configuration for GPU. Install numactl (e.g., apt-get install numactl) to enable automatic NUMA binding."
) )
return False return False
+10 -1
View File
@@ -204,7 +204,16 @@ def encode_image_base64(image_path: Union[str, bytes]):
elif isinstance(image_path, bytes): elif isinstance(image_path, bytes):
return pybase64.b64encode(image_path).decode("utf-8") return pybase64.b64encode(image_path).decode("utf-8")
else: else:
# image_path is PIL.WebPImagePlugin.WebPImageFile import torch
if isinstance(image_path, torch.Tensor):
# Convert GPU-decoded image tensor (C, H, W) uint8 to PIL Image
from PIL import Image
tensor = image_path.cpu() if image_path.device.type != "cpu" else image_path
image_path = Image.fromarray(tensor.permute(1, 2, 0).numpy())
# image_path is a PIL Image
image = image_path image = image_path
buffered = BytesIO() buffered = BytesIO()
image.save(buffered, format="PNG") image.save(buffered, format="PNG")
@@ -32,6 +32,7 @@ class TestSRTBackend(CustomTestCase):
model_path=DEFAULT_MODEL_NAME_FOR_TEST, model_path=DEFAULT_MODEL_NAME_FOR_TEST,
cuda_graph_max_bs=4, cuda_graph_max_bs=4,
mem_fraction_static=0.7, mem_fraction_static=0.7,
log_level="info",
) )
sgl.set_default_backend(cls.backend) sgl.set_default_backend(cls.backend)
@@ -1,13 +1,14 @@
import asyncio import asyncio
import json
import logging import logging
import unittest import unittest
from typing import List from typing import List
import aiohttp import aiohttp
import requests
import torch import torch
from torch.nn.utils.rnn import pad_sequence from torch.nn.utils.rnn import pad_sequence
from sglang.benchmark.utils import download_and_cache_hf_file
from sglang.srt.layers.moe.routed_experts_capturer import ( from sglang.srt.layers.moe.routed_experts_capturer import (
extract_routed_experts_from_meta_info, extract_routed_experts_from_meta_info,
) )
@@ -21,23 +22,18 @@ from sglang.test.test_utils import (
popen_launch_server, popen_launch_server,
) )
register_cuda_ci(est_time=360, suite="stage-c-test-4-gpu-h100") register_cuda_ci(est_time=200, suite="stage-b-test-2-gpu-large")
register_amd_ci( register_amd_ci(
est_time=360, est_time=200,
suite="stage-c-test-4-gpu-amd", suite="stage-b-test-2-gpu-large-amd",
disabled="TP=4 DP=4 routed expert mismatch >15% on AMD; needs TP/DP tuning + concurrency reduction", disabled="TP=2 DP=2 routed expert mismatch >15% on AMD; needs TP/DP tuning + concurrency reduction",
) )
SHAREGPT_URL = ( SHAREGPT_REPO_ID = "anon8231489123/ShareGPT_Vicuna_unfiltered"
"https://huggingface.co/datasets/anon8231489123/" SHAREGPT_FILENAME = "ShareGPT_V3_unfiltered_cleaned_split.json"
"ShareGPT_Vicuna_unfiltered/resolve/main/ShareGPT_V3_unfiltered_cleaned_split.json"
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@unittest.skip(
"Flaky in CI, need to be fixed and re-enabled. See https://github.com/sgl-project/sglang/issues/21266"
)
class TestReturnRoutedExperts(CustomTestCase): class TestReturnRoutedExperts(CustomTestCase):
# modified from test_hicache.py # modified from test_hicache.py
@classmethod @classmethod
@@ -50,31 +46,28 @@ class TestReturnRoutedExperts(CustomTestCase):
"--disable-cuda-graph", "--disable-cuda-graph",
"--disable-radix-cache", "--disable-radix-cache",
"--tp", "--tp",
4, 2,
"--dp", "--dp",
4, 2,
"--enable-dp-attention", "--enable-dp-attention",
] ]
cls.reference_args = [ cls.reference_args = [
"--enable-return-routed-experts", "--enable-return-routed-experts",
"--enable-deterministic-inference", "--enable-deterministic-inference",
"--tp", "--tp",
4, 2,
"--dp", "--dp",
4, 2,
"--enable-dp-attention", "--enable-dp-attention",
] ]
cls.sampling_args = { cls.sampling_args = {
"temperature": 0, "temperature": 0,
} }
# prepare ShareGPT dataset # prepare ShareGPT dataset
try: dataset_path = download_and_cache_hf_file(SHAREGPT_REPO_ID, SHAREGPT_FILENAME)
response = requests.get(SHAREGPT_URL, timeout=60) with open(dataset_path) as f:
response.raise_for_status() data = json.load(f)
data = response.json()
print(f"Dataset size: {len(data)}") print(f"Dataset size: {len(data)}")
except requests.exceptions.RequestException as e:
raise Exception(f"Failed to download ShareGPT dataset: {e}") from e
cls.texts = [] cls.texts = []
for s in data: for s in data:
if "conversations" in s and len(s["conversations"]) > 0: if "conversations" in s and len(s["conversations"]) > 0: