[RL] Fix crash when the reqs in a batch have a mix of return_routed_experts = True and False. (#26423)
Co-authored-by: root <root@slurm-h200-209-231.slurm-compute.tenant-slurm.svc.cluster.local> Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
This commit is contained in:
co-authored by
root
Cursor
Cheng Wan
parent
804f01a823
commit
6f1c9fc77b
@@ -121,8 +121,21 @@ class SchedulerOutputStreamer:
|
||||
skip_req: Optional[Req] = None,
|
||||
is_idle_batch: bool = False,
|
||||
):
|
||||
return_hidden_states = any(
|
||||
req.return_hidden_states for req in reqs if req is not skip_req
|
||||
)
|
||||
return_routed_experts = any(
|
||||
req.return_routed_experts for req in reqs if req is not skip_req
|
||||
)
|
||||
return_indexer_topk = any(
|
||||
req.return_indexer_topk for req in reqs if req is not skip_req
|
||||
)
|
||||
|
||||
acc = _GenerationStreamAccumulator(
|
||||
return_logprob=return_logprob,
|
||||
return_hidden_states=return_hidden_states,
|
||||
return_routed_experts=return_routed_experts,
|
||||
return_indexer_topk=return_indexer_topk,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
disaggregation_mode=self.disaggregation_mode,
|
||||
default_stream_interval=self.server_args.stream_interval,
|
||||
@@ -229,6 +242,9 @@ class SchedulerOutputStreamer:
|
||||
@dataclass(slots=True, kw_only=True)
|
||||
class _GenerationStreamAccumulator:
|
||||
return_logprob: bool
|
||||
return_hidden_states: bool
|
||||
return_routed_experts: bool
|
||||
return_indexer_topk: bool
|
||||
spec_algorithm: Any
|
||||
disaggregation_mode: DisaggregationMode
|
||||
default_stream_interval: int
|
||||
@@ -256,9 +272,9 @@ class _GenerationStreamAccumulator:
|
||||
spec_num_correct_drafts: list = field(default_factory=list)
|
||||
spec_correct_drafts_histogram: list = field(default_factory=list)
|
||||
retraction_counts: list = field(default_factory=list)
|
||||
output_hidden_states: list = field(default_factory=list)
|
||||
routed_experts: list = field(default_factory=list)
|
||||
indexer_topk: list = field(default_factory=list)
|
||||
output_hidden_states: Optional[list] = None
|
||||
routed_experts: Optional[list] = None
|
||||
indexer_topk: Optional[list] = None
|
||||
customized_info: dict = field(default_factory=dict)
|
||||
time_stats: list = field(default_factory=list)
|
||||
input_token_logprobs_val: Optional[list] = None
|
||||
@@ -275,6 +291,13 @@ class _GenerationStreamAccumulator:
|
||||
output_token_ids_logprobs_idx: Optional[list] = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.return_hidden_states:
|
||||
self.output_hidden_states = []
|
||||
if self.return_routed_experts:
|
||||
self.routed_experts = []
|
||||
if self.return_indexer_topk:
|
||||
self.indexer_topk = []
|
||||
|
||||
if self.return_logprob:
|
||||
self.input_token_logprobs_val = []
|
||||
self.input_token_logprobs_idx = []
|
||||
@@ -434,12 +457,18 @@ class _GenerationStreamAccumulator:
|
||||
self.output_token_ids_logprobs_val.append([])
|
||||
self.output_token_ids_logprobs_idx.append([])
|
||||
|
||||
if req.return_hidden_states:
|
||||
self.output_hidden_states.append(req.hidden_states)
|
||||
if req.return_routed_experts:
|
||||
self.routed_experts.append(req.routed_experts)
|
||||
if req.return_indexer_topk:
|
||||
self.indexer_topk.append(req.indexer_topk)
|
||||
if self.return_hidden_states:
|
||||
self.output_hidden_states.append(
|
||||
req.hidden_states if req.return_hidden_states else None
|
||||
)
|
||||
if self.return_routed_experts:
|
||||
self.routed_experts.append(
|
||||
req.routed_experts if req.return_routed_experts else None
|
||||
)
|
||||
if self.return_indexer_topk:
|
||||
self.indexer_topk.append(
|
||||
req.indexer_topk if req.return_indexer_topk else None
|
||||
)
|
||||
|
||||
if req.customized_info is not None:
|
||||
for k, v in req.customized_info.items():
|
||||
@@ -486,9 +515,9 @@ class _GenerationStreamAccumulator:
|
||||
output_token_ids_logprobs_val=self.output_token_ids_logprobs_val,
|
||||
output_token_ids_logprobs_idx=self.output_token_ids_logprobs_idx,
|
||||
output_token_entropy_val=None,
|
||||
output_hidden_states=self.output_hidden_states or None,
|
||||
routed_experts=self.routed_experts or None,
|
||||
indexer_topk=self.indexer_topk or None,
|
||||
output_hidden_states=self.output_hidden_states,
|
||||
routed_experts=self.routed_experts,
|
||||
indexer_topk=self.indexer_topk,
|
||||
customized_info=self.customized_info,
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
|
||||
@@ -1784,7 +1784,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
]
|
||||
|
||||
if getattr(recv_obj, "output_hidden_states", None):
|
||||
meta_info["hidden_states"] = recv_obj.output_hidden_states[i]
|
||||
hidden_states = recv_obj.output_hidden_states[i]
|
||||
if hidden_states is not None:
|
||||
meta_info["hidden_states"] = hidden_states
|
||||
if getattr(recv_obj, "routed_experts", None):
|
||||
val = recv_obj.routed_experts[i]
|
||||
if val is not None:
|
||||
|
||||
@@ -77,6 +77,32 @@ class TestReturnRoutedExperts(CustomTestCase):
|
||||
]
|
||||
cls.reference_args = common
|
||||
cls.sampling_args = {"temperature": 0}
|
||||
cls.texts = None
|
||||
cls.baseline_results = None
|
||||
cls.reference_results = None
|
||||
cls._endpoints = [
|
||||
(
|
||||
"/generate",
|
||||
cls._build_generate_payload,
|
||||
extract_routed_experts_from_meta_info,
|
||||
),
|
||||
(
|
||||
"/v1/chat/completions",
|
||||
cls._build_chat_payload,
|
||||
extract_routed_experts_from_openai_response,
|
||||
),
|
||||
(
|
||||
"/v1/completions",
|
||||
cls._build_completion_payload,
|
||||
extract_routed_experts_from_openai_response,
|
||||
),
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def _ensure_comparison_results(cls):
|
||||
if cls.baseline_results is not None and cls.reference_results is not None:
|
||||
return
|
||||
|
||||
# prepare ShareGPT dataset
|
||||
dataset_path = download_and_cache_hf_file(SHAREGPT_REPO_ID, SHAREGPT_FILENAME)
|
||||
with open(dataset_path) as f:
|
||||
@@ -96,23 +122,6 @@ class TestReturnRoutedExperts(CustomTestCase):
|
||||
if not cls.texts:
|
||||
raise ValueError("No valid texts found in the dataset")
|
||||
cls.texts = cls.texts[:100]
|
||||
cls._endpoints = [
|
||||
(
|
||||
"/generate",
|
||||
cls._build_generate_payload,
|
||||
extract_routed_experts_from_meta_info,
|
||||
),
|
||||
(
|
||||
"/v1/chat/completions",
|
||||
cls._build_chat_payload,
|
||||
extract_routed_experts_from_openai_response,
|
||||
),
|
||||
(
|
||||
"/v1/completions",
|
||||
cls._build_completion_payload,
|
||||
extract_routed_experts_from_openai_response,
|
||||
),
|
||||
]
|
||||
cls.baseline_results = cls._collect_results(cls.baseline_args)
|
||||
cls.reference_results = cls._collect_results(cls.reference_args)
|
||||
|
||||
@@ -128,8 +137,13 @@ class TestReturnRoutedExperts(CustomTestCase):
|
||||
def test_return_routed_experts_completions(cls):
|
||||
cls._run_endpoint_test("/v1/completions")
|
||||
|
||||
def test_mixed_return_routed_experts_batch_alignment(self):
|
||||
self._run_mixed_batch_alignment_case([])
|
||||
self._run_mixed_batch_alignment_case(["--tokenizer-worker-num", 2])
|
||||
|
||||
@classmethod
|
||||
def _run_endpoint_test(cls, endpoint):
|
||||
cls._ensure_comparison_results()
|
||||
captured_baseline_experts = cls.baseline_results[endpoint]
|
||||
captured_reference_experts = cls.reference_results[endpoint]
|
||||
|
||||
@@ -171,6 +185,72 @@ class TestReturnRoutedExperts(CustomTestCase):
|
||||
finally:
|
||||
kill_process_tree(process.pid)
|
||||
|
||||
@classmethod
|
||||
def _run_mixed_batch_alignment_case(cls, other_args):
|
||||
process = popen_launch_server(
|
||||
DEFAULT_ENABLE_ROUTED_EXPERTS_MODEL_NAME_FOR_TEST,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=[
|
||||
"--tp",
|
||||
2,
|
||||
"--enable-return-routed-experts",
|
||||
"--disable-cuda-graph",
|
||||
"--disable-piecewise-cuda-graph",
|
||||
*other_args,
|
||||
],
|
||||
)
|
||||
try:
|
||||
responses = asyncio.run(cls._send_mixed_batch())
|
||||
cls._assert_mixed_batch_result(responses)
|
||||
finally:
|
||||
kill_process_tree(process.pid)
|
||||
|
||||
@classmethod
|
||||
async def _send_mixed_batch(cls):
|
||||
payload_no_rr = {
|
||||
"text": "The quick brown fox jumps over the lazy dog.",
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": 16,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
"return_routed_experts": False,
|
||||
}
|
||||
payload_with_rr = {
|
||||
"text": "The quick brown fox jumps over the lazy dog.",
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": 16,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
"return_routed_experts": True,
|
||||
}
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
return await asyncio.gather(
|
||||
cls._post_generate(session, payload_no_rr),
|
||||
cls._post_generate(session, payload_with_rr),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _post_generate(session, payload):
|
||||
async with session.post(
|
||||
f"{DEFAULT_URL_FOR_TEST}/generate", json=payload
|
||||
) as response:
|
||||
body = await response.json()
|
||||
if response.status != 200:
|
||||
raise AssertionError(f"HTTP {response.status}: {body}")
|
||||
if "error" in body:
|
||||
raise AssertionError(f"generate returned error: {body['error']}")
|
||||
return body
|
||||
|
||||
@classmethod
|
||||
def _assert_mixed_batch_result(cls, responses):
|
||||
no_rr, with_rr = responses
|
||||
assert "routed_experts" not in no_rr.get("meta_info", {})
|
||||
assert "routed_experts" in with_rr.get("meta_info", {})
|
||||
|
||||
@classmethod
|
||||
async def _collect_results_async(cls):
|
||||
results = {}
|
||||
|
||||
Reference in New Issue
Block a user