[Feature] Support return_hidden_states="last" (#30177)
Co-authored-by: litao.dream <litao.dream@bytedance.com>
This commit is contained in:
@@ -32,7 +32,7 @@ class TestHiddenState(CustomTestCase):
|
||||
model_path=cls.model_path,
|
||||
random_seed=42,
|
||||
skip_tokenizer_init=True,
|
||||
enable_return_hidden_states=True,
|
||||
return_hidden_states_mode="full",
|
||||
mem_fraction_static=0.7,
|
||||
)
|
||||
|
||||
@@ -53,8 +53,12 @@ class TestHiddenState(CustomTestCase):
|
||||
return_hidden_states=True,
|
||||
)
|
||||
|
||||
expected_num_hidden_states = self.sampling_params["max_new_tokens"]
|
||||
for output in outputs:
|
||||
self.assertEqual(len(output["meta_info"]["hidden_states"]), 8)
|
||||
self.assertEqual(
|
||||
len(output["meta_info"]["hidden_states"]),
|
||||
expected_num_hidden_states,
|
||||
)
|
||||
for i in range(len(output["meta_info"]["hidden_states"])):
|
||||
assert isinstance(output["meta_info"]["hidden_states"][i], list)
|
||||
output["meta_info"]["hidden_states"][i] = torch.tensor(
|
||||
@@ -103,6 +107,91 @@ class TestHiddenState(CustomTestCase):
|
||||
)
|
||||
)
|
||||
|
||||
def test_return_last_hidden_state(self):
|
||||
outputs = self.engine.generate(
|
||||
input_ids=self.input_ids,
|
||||
sampling_params=self.sampling_params,
|
||||
return_hidden_states="last",
|
||||
)
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
self.model_path, torch_dtype=torch.bfloat16, device_map=get_device()
|
||||
)
|
||||
|
||||
for input_id, output in zip(self.input_ids, outputs):
|
||||
sg_hidden_state = torch.tensor(
|
||||
output["meta_info"]["hidden_states"], dtype=torch.bfloat16
|
||||
).to(get_device())
|
||||
self.assertEqual(sg_hidden_state.dim(), 1)
|
||||
|
||||
with torch.inference_mode():
|
||||
hf_out = model(
|
||||
torch.tensor(
|
||||
[input_id + output["output_ids"][:-1]], device=model.device
|
||||
),
|
||||
output_hidden_states=True,
|
||||
)
|
||||
hf_last_hidden_state = hf_out["hidden_states"][-1][0, -1]
|
||||
|
||||
atol = 0.8
|
||||
self.assertTrue(
|
||||
torch.allclose(
|
||||
hf_last_hidden_state,
|
||||
sg_hidden_state,
|
||||
atol=atol,
|
||||
rtol=0,
|
||||
)
|
||||
)
|
||||
|
||||
def test_mixed_return_hidden_states_modes(self):
|
||||
outputs = self.engine.generate(
|
||||
input_ids=self.input_ids + [self.input_ids[0]],
|
||||
sampling_params=self.sampling_params,
|
||||
return_hidden_states=[False, True, "last"],
|
||||
)
|
||||
|
||||
self.assertNotIn("hidden_states", outputs[0]["meta_info"])
|
||||
|
||||
full_hidden_states = outputs[1]["meta_info"]["hidden_states"]
|
||||
last_hidden_state = outputs[2]["meta_info"]["hidden_states"]
|
||||
|
||||
self.assertIsInstance(full_hidden_states, list)
|
||||
self.assertEqual(
|
||||
len(full_hidden_states), self.sampling_params["max_new_tokens"]
|
||||
)
|
||||
self.assertEqual(torch.tensor(full_hidden_states[0]).dim(), 2)
|
||||
|
||||
last_hidden_state = torch.tensor(last_hidden_state)
|
||||
self.assertEqual(last_hidden_state.dim(), 1)
|
||||
|
||||
def test_mixed_return_hidden_states_modes_with_warm_cache(self):
|
||||
# Prime the radix cache so each repeated prompt only extends its
|
||||
# uncached suffix during the mixed-mode prefill.
|
||||
self.engine.generate(
|
||||
input_ids=self.input_ids,
|
||||
sampling_params={"temperature": 0, "max_new_tokens": 1},
|
||||
return_hidden_states=False,
|
||||
)
|
||||
|
||||
outputs = self.engine.generate(
|
||||
input_ids=self.input_ids + [self.input_ids[1]],
|
||||
sampling_params={"temperature": 0, "max_new_tokens": 1},
|
||||
return_hidden_states=[True, True, "last"],
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
torch.tensor(outputs[0]["meta_info"]["hidden_states"][0]).dim(),
|
||||
2,
|
||||
)
|
||||
self.assertEqual(
|
||||
torch.tensor(outputs[2]["meta_info"]["hidden_states"]).dim(),
|
||||
1,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
torch.tensor(outputs[1]["meta_info"]["hidden_states"][0])[-1],
|
||||
torch.tensor(outputs[2]["meta_info"]["hidden_states"]),
|
||||
)
|
||||
|
||||
def test_repeatedly_changes_hidden_states(self):
|
||||
outputs_completion_first_round = self.engine.generate(
|
||||
input_ids=self.input_ids,
|
||||
|
||||
@@ -100,7 +100,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
|
||||
)
|
||||
|
||||
for choice in response.choices:
|
||||
assert hasattr(choice, "hidden_states") == return_hidden_states
|
||||
assert hasattr(choice, "hidden_states") == bool(return_hidden_states)
|
||||
if return_hidden_states:
|
||||
assert choice.hidden_states is not None, "hidden_states was None"
|
||||
|
||||
@@ -139,7 +139,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
|
||||
usage = response.usage
|
||||
for choice in response.choices:
|
||||
if hasattr(choice, "hidden_states"):
|
||||
assert return_hidden_states
|
||||
assert bool(return_hidden_states)
|
||||
assert choice.hidden_states is not None
|
||||
hidden_states_list.append(choice.hidden_states)
|
||||
|
||||
@@ -169,7 +169,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
|
||||
)
|
||||
|
||||
for choice in response.choices:
|
||||
assert hasattr(choice, "hidden_states") == return_hidden_states
|
||||
assert hasattr(choice, "hidden_states") == bool(return_hidden_states)
|
||||
if return_hidden_states:
|
||||
assert choice.hidden_states is not None, "hidden_states was None"
|
||||
|
||||
@@ -196,7 +196,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
|
||||
for response in generator:
|
||||
for choice in response.choices:
|
||||
if hasattr(choice.delta, "hidden_states"):
|
||||
assert return_hidden_states
|
||||
assert bool(return_hidden_states)
|
||||
assert choice.delta.hidden_states is not None
|
||||
hidden_states_list.append(choice.delta.hidden_states)
|
||||
|
||||
@@ -227,7 +227,7 @@ class TestOpenAIServerWithHiddenStatesEnabled(
|
||||
)
|
||||
cls.base_url += "/v1"
|
||||
cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
|
||||
cls.return_hidden_states = [False, True]
|
||||
cls.return_hidden_states = [False, True, "last"]
|
||||
cls.use_list_input = [True, False]
|
||||
cls.parallel_sample_nums = [1, 2]
|
||||
|
||||
@@ -253,7 +253,7 @@ class TestOpenAIServerWithHiddenStatesEnabledAndCUDAGraphDisabled(
|
||||
)
|
||||
cls.base_url += "/v1"
|
||||
cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
|
||||
cls.return_hidden_states = [False, True]
|
||||
cls.return_hidden_states = [False, True, "last"]
|
||||
cls.use_list_input = [True, False]
|
||||
cls.parallel_sample_nums = [1]
|
||||
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
||||
SchedulerBatchResultProcessor,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor:
|
||||
metrics_reporter = Mock()
|
||||
metrics_reporter.num_generated_tokens = 0
|
||||
metrics_reporter.forward_ct_decode = 0
|
||||
return SchedulerBatchResultProcessor(
|
||||
is_generation=True,
|
||||
disaggregation_mode=None,
|
||||
enable_overlap=False,
|
||||
enable_overlap_mlx=False,
|
||||
server_args=SimpleNamespace(
|
||||
enable_metrics=False,
|
||||
enable_hisparse=False,
|
||||
enable_return_hidden_states=True,
|
||||
return_hidden_states_mode=server_mode,
|
||||
),
|
||||
model_config=SimpleNamespace(think_end_id=None),
|
||||
token_to_kv_pool_allocator=Mock(),
|
||||
tree_cache=None,
|
||||
hisparse_coordinator=None,
|
||||
req_to_token_pool=None,
|
||||
decode_offload_manager=None,
|
||||
metrics_collector=None,
|
||||
metrics_reporter=metrics_reporter,
|
||||
draft_worker=None,
|
||||
model_worker=Mock(),
|
||||
logprob_result_processor=None,
|
||||
output_streamer=Mock(),
|
||||
abort_request=lambda *args, **kwargs: None,
|
||||
)
|
||||
|
||||
|
||||
class _PrefillReq:
|
||||
def __init__(self, *, rid: str, inflight_middle_chunks: int, return_hidden_states):
|
||||
self.rid = rid
|
||||
self.inflight_middle_chunks = inflight_middle_chunks
|
||||
self.return_hidden_states = return_hidden_states
|
||||
self.hidden_states = []
|
||||
self.is_retracted = False
|
||||
self.output_ids = []
|
||||
self.time_stats = Mock()
|
||||
self.return_logprob = False
|
||||
self.return_sampling_mask = False
|
||||
self.grammar = None
|
||||
self.require_reasoning = False
|
||||
self.customized_info = None
|
||||
|
||||
def finished(self):
|
||||
return False
|
||||
|
||||
def update_finish_state(self):
|
||||
return None
|
||||
|
||||
|
||||
class _DecodeReq:
|
||||
def __init__(self):
|
||||
self.return_hidden_states = "last"
|
||||
self.hidden_states = []
|
||||
self.output_ids = []
|
||||
self.finished_len = None
|
||||
self.is_retracted = False
|
||||
self.return_logprob = False
|
||||
self.return_sampling_mask = False
|
||||
self.grammar = None
|
||||
self.time_stats = Mock()
|
||||
|
||||
def finished(self):
|
||||
return self.finished_len is not None
|
||||
|
||||
def update_finish_state(self, new_accept_len):
|
||||
if len(self.output_ids) >= 6:
|
||||
self.finished_len = 5
|
||||
|
||||
|
||||
class TestPrefillHiddenStateOffsets(CustomTestCase):
|
||||
def test_active_middle_chunk_advances_before_new_last_request(self):
|
||||
cases = (
|
||||
(
|
||||
"full",
|
||||
CaptureHiddenMode.FULL,
|
||||
torch.tensor([[10.0], [11.0], [20.0], [21.0], [22.0]]),
|
||||
),
|
||||
(
|
||||
"last",
|
||||
CaptureHiddenMode.LAST,
|
||||
torch.tensor([[11.0], [22.0]]),
|
||||
),
|
||||
)
|
||||
|
||||
for server_mode, capture_mode, hidden_states in cases:
|
||||
with self.subTest(server_mode=server_mode):
|
||||
middle = _PrefillReq(
|
||||
rid="middle",
|
||||
inflight_middle_chunks=1,
|
||||
return_hidden_states=False,
|
||||
)
|
||||
last = _PrefillReq(
|
||||
rid="last",
|
||||
inflight_middle_chunks=0,
|
||||
return_hidden_states="last",
|
||||
)
|
||||
batch = SimpleNamespace(
|
||||
reqs=[middle, last],
|
||||
decoding_reqs=[],
|
||||
return_logprob=False,
|
||||
return_hidden_states=True,
|
||||
return_hidden_states_mode=capture_mode,
|
||||
spec_info=None,
|
||||
prefill_stats=None,
|
||||
dp_cooperation_info=None,
|
||||
)
|
||||
result = SimpleNamespace(
|
||||
copy_done=None,
|
||||
routed_experts_output=None,
|
||||
indexer_topk_output=None,
|
||||
logits_output=SimpleNamespace(
|
||||
hidden_states=hidden_states,
|
||||
customized_info=None,
|
||||
),
|
||||
next_token_ids=torch.tensor([0, 1]),
|
||||
extend_input_len_per_req=[2, 3],
|
||||
extend_logprob_start_len_per_req=None,
|
||||
grammar_advanced=False,
|
||||
can_run_cuda_graph=False,
|
||||
skipped_output_comm=False,
|
||||
)
|
||||
processor = _make_processor(server_mode)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components."
|
||||
"batch_result_processor.maybe_cache_unfinished_req"
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components."
|
||||
"batch_result_processor.get_memory",
|
||||
return_value=SimpleNamespace(enable_hisparse=False),
|
||||
),
|
||||
):
|
||||
processor.process_batch_result_prefill(batch, result)
|
||||
|
||||
self.assertEqual(middle.hidden_states, [])
|
||||
self.assertEqual(last.hidden_states, [[22.0]])
|
||||
|
||||
|
||||
class TestDecodeHiddenStateRetention(CustomTestCase):
|
||||
def test_last_mode_multi_step_storage_stays_bounded(self):
|
||||
processor = _make_processor()
|
||||
req = _DecodeReq()
|
||||
batch = SimpleNamespace(
|
||||
reqs=[req],
|
||||
return_logprob=False,
|
||||
spec_algorithm=SimpleNamespace(is_none=lambda: False),
|
||||
batch_size=lambda: 1,
|
||||
)
|
||||
first_step = torch.arange(8, dtype=torch.float32).view(4, 2)
|
||||
second_step = torch.arange(16, dtype=torch.float32).view(8, 2)[4:]
|
||||
|
||||
def result(hidden_states):
|
||||
return SimpleNamespace(
|
||||
copy_done=None,
|
||||
routed_experts_output=None,
|
||||
indexer_topk_output=None,
|
||||
logits_output=SimpleNamespace(hidden_states=hidden_states),
|
||||
next_token_ids=None,
|
||||
can_run_cuda_graph=False,
|
||||
num_correct_drafts=0,
|
||||
num_block_accept_tokens=0,
|
||||
num_cap_tokens=0,
|
||||
speculative_num_draft_tokens=4,
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
SchedulerBatchResultProcessor,
|
||||
"_normalize_decode_outputs",
|
||||
side_effect=[
|
||||
([[1, 2, 3]], None),
|
||||
([[4, 5, 6]], None),
|
||||
],
|
||||
),
|
||||
patch.object(
|
||||
SchedulerBatchResultProcessor,
|
||||
"_maybe_update_reasoning_tokens",
|
||||
),
|
||||
patch.object(
|
||||
SchedulerBatchResultProcessor,
|
||||
"_handle_finish_state_updated_req",
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.managers.scheduler_components."
|
||||
"batch_result_processor.get_observability",
|
||||
return_value=SimpleNamespace(enable_metrics=False),
|
||||
),
|
||||
):
|
||||
processor.process_batch_result_decode(batch, result(first_step))
|
||||
|
||||
self.assertEqual(req.hidden_states, [first_step[2].tolist()])
|
||||
self.assertEqual(len(req.hidden_states), 1)
|
||||
|
||||
# Only the first two accepted tokens are valid because the request
|
||||
# stops inside this speculative verify step.
|
||||
processor.process_batch_result_decode(batch, result(second_step))
|
||||
|
||||
self.assertEqual(req.hidden_states, [second_step[1].tolist()])
|
||||
self.assertEqual(len(req.hidden_states), 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,72 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class TestHiddenStateServerMode(CustomTestCase):
|
||||
@staticmethod
|
||||
def _make_tokenizer_manager(mode):
|
||||
manager = TokenizerManager.__new__(TokenizerManager)
|
||||
manager.context_len = 128
|
||||
manager.num_reserved_tokens = 0
|
||||
manager.allow_auto_truncate = False
|
||||
manager.validate_total_tokens = False
|
||||
manager.is_generation = True
|
||||
manager.server_args = SimpleNamespace(
|
||||
enable_return_hidden_states=mode is not None,
|
||||
return_hidden_states_mode=mode,
|
||||
enable_custom_logit_processor=False,
|
||||
)
|
||||
manager._validate_token_ids_logprob = Mock()
|
||||
return manager
|
||||
|
||||
@staticmethod
|
||||
def _make_request(return_hidden_states):
|
||||
return GenerateReqInput(
|
||||
input_ids=[1, 2, 3],
|
||||
sampling_params={},
|
||||
return_hidden_states=return_hidden_states,
|
||||
)
|
||||
|
||||
def test_last_server_accepts_false_and_last(self):
|
||||
manager = self._make_tokenizer_manager("last")
|
||||
|
||||
for mode in (False, "last"):
|
||||
with self.subTest(mode=mode):
|
||||
manager._validate_one_request(
|
||||
self._make_request(mode),
|
||||
[1, 2, 3],
|
||||
)
|
||||
|
||||
def test_last_server_rejects_full(self):
|
||||
manager = self._make_tokenizer_manager("last")
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"server maximum `last`",
|
||||
):
|
||||
manager._validate_one_request(
|
||||
self._make_request(True),
|
||||
[1, 2, 3],
|
||||
)
|
||||
|
||||
def test_full_server_accepts_all_request_modes(self):
|
||||
manager = self._make_tokenizer_manager("full")
|
||||
|
||||
for mode in (False, "last", True):
|
||||
with self.subTest(mode=mode):
|
||||
manager._validate_one_request(
|
||||
self._make_request(mode),
|
||||
[1, 2, 3],
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -164,6 +164,48 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
# Check text expansion
|
||||
self.assertEqual(req.text, expected_text)
|
||||
|
||||
def test_return_hidden_states_expands_with_parallel_sampling(self):
|
||||
req = GenerateReqInput(
|
||||
text=["Prompt 1", "Prompt 2"],
|
||||
sampling_params={"n": 2},
|
||||
return_hidden_states=[False, "last"],
|
||||
)
|
||||
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
self.assertEqual(
|
||||
req.return_hidden_states,
|
||||
[False, "last", False, "last"],
|
||||
)
|
||||
self.assertEqual(
|
||||
[req[i].return_hidden_states for i in range(4)],
|
||||
[False, "last", False, "last"],
|
||||
)
|
||||
|
||||
def test_return_hidden_states_batch_length_is_validated(self):
|
||||
req = GenerateReqInput(
|
||||
text=["Prompt 1", "Prompt 2"],
|
||||
return_hidden_states=["last"],
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"return_hidden_states should be equal to the batch size",
|
||||
):
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
def test_return_hidden_states_batch_modes_are_validated(self):
|
||||
req = GenerateReqInput(
|
||||
text=["Prompt 1", "Prompt 2"],
|
||||
return_hidden_states=[False, "invalid"],
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"return_hidden_states must be a boolean or the string literal 'last'",
|
||||
):
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
def test_mixed_none_and_images_with_parallel_samples(self):
|
||||
"""Test that when some batch items have images and others None, parallel expansion works correctly."""
|
||||
req = copy.deepcopy(self.base_req)
|
||||
@@ -480,14 +522,14 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
logprob_start_len=[10, 5],
|
||||
top_logprobs_num=[5, 3],
|
||||
token_ids_logprob=[[7, 8, 9], [4, 5, 6]],
|
||||
return_hidden_states=[False, False, True],
|
||||
return_hidden_states=[False, True],
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
self.assertEqual(req.return_logprob, [True, False])
|
||||
self.assertEqual(req.logprob_start_len, [10, 5])
|
||||
self.assertEqual(req.top_logprobs_num, [5, 3])
|
||||
self.assertEqual(req.token_ids_logprob, [[7, 8, 9], [4, 5, 6]])
|
||||
self.assertEqual(req.return_hidden_states, [False, False, True])
|
||||
self.assertEqual(req.return_hidden_states, [False, True])
|
||||
|
||||
def test_custom_logit_processor_normalization(self):
|
||||
"""Test normalization of custom_logit_processor."""
|
||||
@@ -559,7 +601,7 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
modalities=["image", "image"],
|
||||
lora_path=["path1", "path2"],
|
||||
custom_logit_processor=["processor1", "processor2"],
|
||||
return_hidden_states=True,
|
||||
return_hidden_states=[True, "last"],
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
@@ -580,6 +622,7 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
self.assertEqual(item0.lora_path, "path1")
|
||||
self.assertEqual(item0.custom_logit_processor, "processor1")
|
||||
self.assertEqual(item0.return_hidden_states, True)
|
||||
self.assertEqual(req[1].return_hidden_states, "last")
|
||||
|
||||
def test_getitem_preserves_return_prompt_token_ids(self):
|
||||
"""Batch subrequests must keep the prompt-token-id return flag."""
|
||||
|
||||
@@ -11,6 +11,7 @@ maybe_stub_sgl_kernel()
|
||||
|
||||
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator # noqa: E402
|
||||
from sglang.srt.managers.scheduler import Scheduler # noqa: E402
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode # noqa: E402
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
@@ -23,6 +24,7 @@ def _make_req(req_pool_idx, origin_input_ids, output_ids):
|
||||
return_logprob=False,
|
||||
grammar=None,
|
||||
return_hidden_states=False,
|
||||
return_hidden_states_mode=CaptureHiddenMode.NULL,
|
||||
is_prefill_only=False,
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardMode,
|
||||
get_server_return_hidden_states_mode,
|
||||
)
|
||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||
DecodeCudaGraphRunner,
|
||||
)
|
||||
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
|
||||
PrefillCudaGraphRunner,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class TestHiddenStateGraphRecapture(CustomTestCase):
|
||||
def test_server_mode_sets_graph_capture_ceiling(self):
|
||||
disabled = SimpleNamespace(
|
||||
enable_return_hidden_states=False,
|
||||
return_hidden_states_mode=None,
|
||||
)
|
||||
last = SimpleNamespace(
|
||||
enable_return_hidden_states=True,
|
||||
return_hidden_states_mode="last",
|
||||
)
|
||||
full = SimpleNamespace(
|
||||
enable_return_hidden_states=True,
|
||||
return_hidden_states_mode="full",
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
get_server_return_hidden_states_mode(disabled),
|
||||
CaptureHiddenMode.NULL,
|
||||
)
|
||||
self.assertEqual(
|
||||
get_server_return_hidden_states_mode(last),
|
||||
CaptureHiddenMode.LAST,
|
||||
)
|
||||
self.assertEqual(
|
||||
get_server_return_hidden_states_mode(full),
|
||||
CaptureHiddenMode.FULL,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _make_runner(runner_cls, capture_hidden_mode):
|
||||
runner = runner_cls.__new__(runner_cls)
|
||||
runner.capture_hidden_mode = capture_hidden_mode
|
||||
runner.backend = Mock()
|
||||
runner.capture = Mock()
|
||||
return runner
|
||||
|
||||
@staticmethod
|
||||
def _make_forward_batch(capture_hidden_mode):
|
||||
return SimpleNamespace(
|
||||
capture_hidden_mode=capture_hidden_mode,
|
||||
spec_info=None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _make_prefill_runner_for_can_run(capture_hidden_mode):
|
||||
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
|
||||
runner._is_full_backend = False
|
||||
runner.prefill_backend_name = Backend.BREAKABLE
|
||||
runner.has_mha_companion_layers = False
|
||||
runner.enable_lora = False
|
||||
runner._capture_chunked_prefix = False
|
||||
runner.capture_hidden_mode = capture_hidden_mode
|
||||
runner.capture_num_tokens = [4]
|
||||
runner.max_num_tokens = 4
|
||||
return runner
|
||||
|
||||
@staticmethod
|
||||
def _make_prefill_forward_batch(capture_hidden_mode, spec_capture_hidden_mode):
|
||||
return SimpleNamespace(
|
||||
batch_size=1,
|
||||
input_embeds=None,
|
||||
replace_embeds=None,
|
||||
forward_mode=ForwardMode.EXTEND,
|
||||
capture_hidden_mode=capture_hidden_mode,
|
||||
spec_info=SimpleNamespace(capture_hidden_mode=spec_capture_hidden_mode),
|
||||
global_num_tokens_cpu=None,
|
||||
return_logprob=False,
|
||||
extend_prefix_lens_cpu=None,
|
||||
input_ids=list(range(4)),
|
||||
)
|
||||
|
||||
def test_stronger_graph_is_reused_for_weaker_modes(self):
|
||||
runner = self._make_runner(DecodeCudaGraphRunner, CaptureHiddenMode.FULL)
|
||||
|
||||
for required_mode in (
|
||||
CaptureHiddenMode.FULL,
|
||||
CaptureHiddenMode.NULL,
|
||||
CaptureHiddenMode.LAST,
|
||||
CaptureHiddenMode.FULL,
|
||||
CaptureHiddenMode.NULL,
|
||||
):
|
||||
with self.subTest(required_mode=required_mode):
|
||||
runner._validate_capture_hidden_mode(
|
||||
self._make_forward_batch(required_mode)
|
||||
)
|
||||
|
||||
self.assertEqual(runner.capture_hidden_mode, CaptureHiddenMode.FULL)
|
||||
runner.backend.cleanup.assert_not_called()
|
||||
runner.capture.assert_not_called()
|
||||
|
||||
def test_graph_does_not_recapture_above_fixed_server_mode(self):
|
||||
for runner_cls in (
|
||||
DecodeCudaGraphRunner,
|
||||
PrefillCudaGraphRunner,
|
||||
CPUGraphRunner,
|
||||
):
|
||||
runner = self._make_runner(runner_cls, CaptureHiddenMode.NULL)
|
||||
|
||||
with self.subTest(runner_cls=runner_cls), self.assertRaisesRegex(
|
||||
RuntimeError,
|
||||
"exceeds the fixed (CUDA|CPU) graph capture mode",
|
||||
):
|
||||
runner._validate_capture_hidden_mode(
|
||||
self._make_forward_batch(CaptureHiddenMode.LAST)
|
||||
)
|
||||
|
||||
self.assertEqual(runner.capture_hidden_mode, CaptureHiddenMode.NULL)
|
||||
runner.backend.cleanup.assert_not_called()
|
||||
runner.capture.assert_not_called()
|
||||
|
||||
def test_spec_worker_override_is_the_effective_runtime_mode(self):
|
||||
runner = self._make_prefill_runner_for_can_run(CaptureHiddenMode.LAST)
|
||||
forward_batch = self._make_prefill_forward_batch(
|
||||
CaptureHiddenMode.LAST,
|
||||
CaptureHiddenMode.FULL,
|
||||
)
|
||||
|
||||
self.assertTrue(runner.can_run_graph(forward_batch))
|
||||
for runner_cls in (
|
||||
DecodeCudaGraphRunner,
|
||||
PrefillCudaGraphRunner,
|
||||
CPUGraphRunner,
|
||||
):
|
||||
graph_runner = self._make_runner(runner_cls, CaptureHiddenMode.LAST)
|
||||
with self.subTest(runner_cls=runner_cls):
|
||||
graph_runner._validate_capture_hidden_mode(forward_batch)
|
||||
|
||||
def test_prefill_graph_falls_back_for_stronger_effective_mode(self):
|
||||
runner = self._make_prefill_runner_for_can_run(CaptureHiddenMode.LAST)
|
||||
forward_batch = self._make_prefill_forward_batch(
|
||||
CaptureHiddenMode.FULL,
|
||||
CaptureHiddenMode.LAST,
|
||||
)
|
||||
|
||||
self.assertFalse(runner.can_run_graph(forward_batch))
|
||||
|
||||
def test_prefill_graph_accepts_weaker_spec_mode(self):
|
||||
runner = self._make_prefill_runner_for_can_run(CaptureHiddenMode.FULL)
|
||||
forward_batch = self._make_prefill_forward_batch(
|
||||
CaptureHiddenMode.NULL,
|
||||
CaptureHiddenMode.LAST,
|
||||
)
|
||||
|
||||
self.assertTrue(runner.can_run_graph(forward_batch))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -6,8 +6,13 @@ from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
import sglang.srt.model_executor.model_runner_components.cuda_graph_setup as graph_setup
|
||||
import sglang.srt.model_executor.runner.prefill_cuda_graph_runner as runner_module
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
||||
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
|
||||
capture_prefill_graph,
|
||||
)
|
||||
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
|
||||
PrefillCudaGraphRunner,
|
||||
)
|
||||
@@ -56,6 +61,29 @@ class _FakeKVIndexKernel:
|
||||
|
||||
|
||||
class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
||||
def test_eagle_target_tc_piecewise_skips_last_mode_capture(self):
|
||||
eager_runner = object()
|
||||
model_runner = SimpleNamespace(
|
||||
is_draft_worker=False,
|
||||
spec_algorithm=SimpleNamespace(is_eagle=lambda: True),
|
||||
server_args=SimpleNamespace(
|
||||
enable_return_hidden_states=True,
|
||||
return_hidden_states_mode="last",
|
||||
),
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
graph_setup,
|
||||
"check_cuda_graph_backend",
|
||||
return_value=False,
|
||||
):
|
||||
runner = capture_prefill_graph(
|
||||
model_runner=model_runner,
|
||||
eager_runner=eager_runner,
|
||||
)
|
||||
|
||||
self.assertIs(runner, eager_runner)
|
||||
|
||||
def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
|
||||
model_runner = SimpleNamespace(
|
||||
server_args=SimpleNamespace(
|
||||
@@ -212,7 +240,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
||||
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
|
||||
runner._capture_req_slots = 4
|
||||
runner.enable_lora = False
|
||||
runner.capture_hidden_mode = None
|
||||
runner.capture_hidden_mode = CaptureHiddenMode.NULL
|
||||
runner.max_num_tokens = 32
|
||||
runner.capture_num_tokens = [4]
|
||||
runner.backend = SimpleNamespace()
|
||||
@@ -227,7 +255,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
||||
input_embeds=None,
|
||||
replace_embeds=None,
|
||||
forward_mode=SimpleNamespace(is_target_verify=lambda: False),
|
||||
capture_hidden_mode=None,
|
||||
capture_hidden_mode=CaptureHiddenMode.NULL,
|
||||
global_num_tokens_cpu=None,
|
||||
return_logprob=False,
|
||||
extend_prefix_lens_cpu=[8],
|
||||
|
||||
@@ -42,6 +42,45 @@ _mock_device.start()
|
||||
|
||||
|
||||
class TestPrepareServerArgs(CustomTestCase):
|
||||
def test_return_hidden_states_mode_configuration(self):
|
||||
disabled = ServerArgs(model_path="dummy")
|
||||
self.assertFalse(disabled.enable_return_hidden_states)
|
||||
self.assertIsNone(disabled.return_hidden_states_mode)
|
||||
|
||||
last = ServerArgs(
|
||||
model_path="dummy",
|
||||
return_hidden_states_mode="last",
|
||||
)
|
||||
self.assertTrue(last.enable_return_hidden_states)
|
||||
self.assertEqual(last.return_hidden_states_mode, "last")
|
||||
|
||||
legacy_full = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_return_hidden_states=True,
|
||||
)
|
||||
self.assertTrue(legacy_full.enable_return_hidden_states)
|
||||
self.assertEqual(legacy_full.return_hidden_states_mode, "full")
|
||||
|
||||
parsed_last = prepare_server_args(
|
||||
[
|
||||
"--model-path",
|
||||
"dummy",
|
||||
"--return-hidden-states-mode",
|
||||
"last",
|
||||
]
|
||||
)
|
||||
self.assertTrue(parsed_last.enable_return_hidden_states)
|
||||
self.assertEqual(parsed_last.return_hidden_states_mode, "last")
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"return_hidden_states_mode must be one of",
|
||||
):
|
||||
ServerArgs(
|
||||
model_path="dummy",
|
||||
return_hidden_states_mode="lst",
|
||||
)
|
||||
|
||||
def test_config_nested_dict_args_are_json(self):
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
|
||||
f.write("mm-process-config:\n image:\n resize: 128\n")
|
||||
|
||||
Reference in New Issue
Block a user