[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,