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:
co-authored by
Claude Opus 4.6
parent
2acdda1d85
commit
1d9c8e8c9e
@@ -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]
|
||||||
|
|||||||
@@ -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
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user