[Feature] Support return_hidden_states="last" (#30177)

Co-authored-by: litao.dream <litao.dream@bytedance.com>
This commit is contained in:
Tao Li
2026-08-02 15:09:33 +08:00
committed by GitHub
co-authored by litao.dream
parent 1685d29f21
commit a0b7bcf592
30 changed files with 1096 additions and 268 deletions
+91 -2
View File
@@ -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")