diff --git a/slime/backends/vllm_utils/vllm_engine.py b/slime/backends/vllm_utils/vllm_engine.py index 47ad04e46..09d3df49b 100644 --- a/slime/backends/vllm_utils/vllm_engine.py +++ b/slime/backends/vllm_utils/vllm_engine.py @@ -5,6 +5,10 @@ import multiprocessing import os import time +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + import numpy as np from urllib.parse import quote import requests @@ -304,9 +308,6 @@ def launch_server_process( # unless the user already passed --vllm-max-model-len explicitly. if args.rollout_max_context_len is not None and getattr(args, "vllm_max_model_len", None) is None: cmd += ["--max-model-len", str(args.rollout_max_context_len)] - if getattr(args, "use_rollout_routing_replay", False): - cmd += ["--enable-return-routed-experts"] - # vime-preferred defaults — must be explicitly forwarded because the vllm-side # default would otherwise apply (the generic forwarder skips values that equal # action.default). @@ -332,6 +333,22 @@ def _user_overrode(dest: str) -> bool: _, action = entry return getattr(args, dest, action.default) != action.default + # MoE routing replay: routed-experts return + sync scheduling (incompatible with + # async scheduling on vLLM 0.21.x). Expert parallel is opt-in via + # ``--vllm-enable-expert-parallel``. + if getattr(args, "use_rollout_routing_replay", False): + cmd += ["--enable-return-routed-experts"] + if not _user_overrode("vllm_async_scheduling"): + cmd += ["--no-async-scheduling"] + # Prefix cache hits skip prefill MoE forwards; routed-experts capture then + # only covers decode (~gen_len-1 rows) and prompt_routed_experts is missing. + if not _user_overrode("vllm_enable_prefix_caching"): + cmd += ["--no-enable-prefix-caching"] + if getattr(args, "vllm_enable_expert_parallel", False): + cmd += ["--enable-expert-parallel"] + if not _user_overrode("vllm_expert_placement_strategy"): + cmd += ["--expert-placement-strategy", "linear"] + # 1) gpu_memory_utilization: vllm default 0.92 OOMs in colocate training; vime ships 0.55. if _user_overrode("vllm_gpu_memory_utilization"): gpu_mem = args.vllm_gpu_memory_utilization @@ -387,6 +404,76 @@ def _redact_cmd_for_log(cmd: list[str]) -> str: return " ".join(parts) +def _routing_rows_from_http_payload(value: Any) -> np.ndarray | None: + """Decode vLLM routed-experts HTTP field (base64 npy or nested list).""" + import base64 + import io + + import numpy as np + + if value is None: + return None + if isinstance(value, str): + return np.load(io.BytesIO(base64.b64decode(value)), allow_pickle=False) + if isinstance(value, list): + return np.asarray(value, dtype=np.int32) + return None + + +def _verify_generate_routed_experts(base_url: str, model: str, timeout_s: float = 120.0) -> None: + """Smoke-check MoE routing via ``/v1/completions`` (matches R3 rollout on vLLM 0.21.x).""" + import numpy as np + + base = base_url.rstrip("/") + payload = { + "model": model, + "prompt": "Routing replay smoke test.", + "max_tokens": 8, + "temperature": 0.0, + "logprobs": 1, + "return_token_ids": True, + "stream": False, + } + response = requests.post(f"{base}/v1/completions", json=payload, timeout=timeout_s) + response.raise_for_status() + body = response.json() + choice = (body.get("choices") or [{}])[0] + pre = _routing_rows_from_http_payload(body.get("prompt_routed_experts")) + gen = _routing_rows_from_http_payload(choice.get("routed_experts")) + if pre is None and gen is None: + raise RuntimeError( + "vLLM /v1/completions returned no routed-experts fields. " + "Ensure the server was started with --enable-return-routed-experts " + "(use_rollout_routing_replay)." + ) + if pre is None or gen is None: + raise RuntimeError( + "vLLM /v1/completions must return both prompt_routed_experts and " + "choices[].routed_experts for routing replay. " + "Ensure --enable-return-routed-experts and --no-async-scheduling on the vLLM cmdline." + ) + + out_ids = choice.get("token_ids") or [] + usage = body.get("usage") or {} + num_prompt = int(usage.get("prompt_tokens") or 0) + num_gen = int(usage.get("completion_tokens") or len(out_ids)) + expected_rows = num_prompt + num_gen - 1 if num_prompt > 0 and num_gen > 0 else 0 + + merged = np.concatenate([pre, gen], axis=0) + n_rows = int(merged.shape[0]) + if expected_rows > 0 and n_rows not in (expected_rows, expected_rows + 1): + raise RuntimeError( + f"vLLM routing replay smoke check: merged routing rows {n_rows} != " + f"expected {expected_rows} (prompt+gen len(tokens)-1). " + "Ensure --enable-return-routed-experts and --no-async-scheduling." + ) + logger.info( + "vLLM routing replay smoke check OK (/v1/completions): prompt+gen routing rows=%s " "(expected %s)", + n_rows, + expected_rows, + ) + + def _wait_server_healthy(base_url: str, process: multiprocessing.Process | None, timeout_s: float = 300.0) -> None: """Wait until the vLLM server responds on ``GET /health`` (SGLang stacks typically use ``GET /health_generate``).""" start = time.time() @@ -550,7 +637,10 @@ def _init_normal(self) -> None: visible_devices=visible_devices, model_path=self.model_path, ) - _wait_server_healthy(self._http_base(), process=self.process) + base = self._http_base() + _wait_server_healthy(base, process=self.process) + if getattr(self.args, "use_rollout_routing_replay", False): + _verify_generate_routed_experts(base, self.model_path) def _post_json(self, endpoint: str, payload: dict, timeout: float) -> requests.Response: url = f"{self._http_base()}/{endpoint.lstrip('/')}" diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py index 3dd576fe5..06737ffe3 100644 --- a/slime/ray/rollout.py +++ b/slime/ray/rollout.py @@ -774,7 +774,17 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl if samples[0].rollout_log_probs is not None: train_data["rollout_log_probs"] = [sample.rollout_log_probs for sample in samples] - if samples[0].rollout_routed_experts is not None: + if getattr(self.args, "use_rollout_routing_replay", False): + routed = [sample.rollout_routed_experts for sample in samples] + missing = [i for i, r in enumerate(routed) if r is None] + if missing: + raise ValueError( + f"use_rollout_routing_replay: {len(missing)}/{len(samples)} samples missing " + "rollout_routed_experts (see rollout logs for vLLM routing replay errors). " + "Ensure vLLM serves with --enable-return-routed-experts and --no-async-scheduling." + ) + train_data["rollout_routed_experts"] = routed + elif samples[0].rollout_routed_experts is not None: train_data["rollout_routed_experts"] = [sample.rollout_routed_experts for sample in samples] if samples[0].train_metadata is not None: diff --git a/slime/rollout/vllm_rollout.py b/slime/rollout/vllm_rollout.py index 7d5705548..9014a422b 100644 --- a/slime/rollout/vllm_rollout.py +++ b/slime/rollout/vllm_rollout.py @@ -161,40 +161,101 @@ def _vllm_meta_from_generate_choice(args: Namespace, choice: dict, usage: dict | def _decode_vllm_routed_experts(value: str) -> np.ndarray: + """Decode vLLM routed-experts field when returned as base64 ``.npy`` (optional).""" raw = base64.b64decode(value.encode("ascii"), validate=True) return np.load(io.BytesIO(raw), allow_pickle=False) -def _apply_vllm_routed_experts( - args: Namespace, - sample: Sample, - _output: dict, +def _routing_array_from_payload(value: Any) -> np.ndarray | None: + if value is None: + return None + if isinstance(value, str): + return _decode_vllm_routed_experts(value) + if isinstance(value, list): + return np.asarray(value, dtype=np.int32) + logger.warning("routed_experts payload must be nested list or base64 npy from vLLM HTTP API") + return None + + +def _merge_generate_routed_experts( + output: dict, choice: dict, -) -> None: - """Populate ``sample.rollout_routed_experts`` from vLLM ``choices[].routed_experts`` when enabled. + gen_token_count: int | None, + prompt_token_count: int | None = None, +) -> np.ndarray | None: + """Merge vLLM generate routing (same semantics as ``/v1/completions``). - vLLM ``/inference/v1/generate`` returns routed experts as a base64 encoded - ``.npy`` payload on each response choice when the server is launched with - ``--enable-return-routed-experts``. + SGLang returns one buffer reshaped to ``(len(tokens) - 1, num_layers, top_k)``. """ - if not getattr(args, "use_rollout_routing_replay", False): - return + parts: list[np.ndarray] = [] + prompt_re = output.get("prompt_routed_experts") + if prompt_re is not None: + prompt_arr = _routing_array_from_payload(prompt_re) + if prompt_arr is not None and prompt_arr.ndim == 3: + parts.append(prompt_arr) + elif prompt_arr is not None: + logger.warning(f"Unexpected prompt_routed_experts ndim={prompt_arr.ndim}") + gen_re = choice.get("routed_experts") - if gen_re is None: - return - arr = _decode_vllm_routed_experts(gen_re) - n_tok = len(sample.tokens) - expected_rows = max(0, n_tok - 1) + if gen_re is not None: + gen_arr = _routing_array_from_payload(gen_re) + if gen_arr is not None and gen_arr.ndim == 3: + if ( + prompt_re is None + and prompt_token_count is not None + and prompt_token_count > 0 + and gen_token_count is not None + and gen_token_count > 0 + and ( + gen_arr.shape[0] > gen_token_count or gen_arr.shape[0] >= prompt_token_count + gen_token_count - 1 + ) + ): + prompt_arr = gen_arr[:prompt_token_count] + gen_arr = gen_arr[prompt_token_count:] + if prompt_arr.size > 0: + parts.append(prompt_arr) + if gen_token_count is not None and gen_token_count > 0 and gen_arr.shape[0] > gen_token_count: + gen_arr = gen_arr[:gen_token_count] + if gen_arr.size > 0: + parts.append(gen_arr) + elif gen_arr is not None: + logger.warning(f"Unexpected routed_experts ndim={gen_arr.ndim}") + + if not parts: + return None + if len(parts) == 1: + return parts[0] + return np.concatenate(parts, axis=0) + + +def _align_routed_experts_rows(arr: np.ndarray, expected_rows: int) -> np.ndarray | None: + """Resize merged routing to ``len(tokens) - 1`` rows (SGLang / Megatron layout).""" if arr.ndim != 3: logger.warning(f"Unexpected routed_experts ndim={arr.ndim} shape={arr.shape}") - return - if arr.shape[0] == n_tok: - arr = arr[:-1] - elif arr.shape[0] != expected_rows: + return None + n_rows = arr.shape[0] + if n_rows == expected_rows: + return arr + if n_rows == expected_rows + 1: + return arr[:-1] + if n_rows > expected_rows + 1: logger.warning( - f"routed_experts row count {arr.shape[0]} not in {{{expected_rows}, {n_tok}}}; " - "skipping rollout_routed_experts assign", + f"routed_experts row count {n_rows} > expected {expected_rows}; trimming to {expected_rows}", ) + return arr[:expected_rows] + logger.warning( + f"routed_experts row count {n_rows} < expected {expected_rows}; " + "missing prompt_routed_experts on generate response?", + ) + return None + + +def _assign_rollout_routed_experts(args: Namespace, sample: Sample, arr: np.ndarray | None) -> None: + if arr is None: + return + expected_rows = max(0, len(sample.tokens) - 1) + arr = _align_routed_experts_rows(arr, expected_rows) + if arr is None: return nl = getattr(args, "num_layers", None) mtk = getattr(args, "moe_router_topk", None) @@ -203,7 +264,37 @@ def _apply_vllm_routed_experts( f"routed_experts shape {arr.shape} does not match args (num_layers={nl}, moe_router_topk={mtk})", ) return - sample.rollout_routed_experts = arr + sample.rollout_routed_experts = np.ascontiguousarray(arr.astype(np.int32, copy=True)) + + +def _apply_vllm_routed_experts( + args: Namespace, + sample: Sample, + output: dict, + choice: dict, + gen_token_count: int | None = None, + prompt_token_count: int | None = None, +) -> None: + """Populate ``sample.rollout_routed_experts`` from ``/inference/v1/generate`` when R3 is enabled.""" + if not getattr(args, "use_rollout_routing_replay", False): + return + arr = _merge_generate_routed_experts(output, choice, gen_token_count, prompt_token_count=prompt_token_count) + _assign_rollout_routed_experts(args, sample, arr) + if sample.rollout_routed_experts is not None: + return + if sample.status == Sample.Status.ABORTED and sample.response_length == 0: + return + pre = output.get("prompt_routed_experts") + gen = choice.get("routed_experts") + raise RuntimeError( + "vLLM routing replay: failed to set sample.rollout_routed_experts. " + f"prompt_routed_experts in response={pre is not None}, " + f"choices[0].routed_experts={gen is not None}, " + f"tokens={len(sample.tokens)}, gen_tokens={gen_token_count}. " + "Check vLLM was launched with --enable-return-routed-experts, " + "--no-enable-prefix-caching, and --enforce-eager. " + "Rollout uses /v1/completions for R3 (not /inference/v1/generate)." + ) def _inference_generate_tokens_and_logprobs(choice: dict[str, Any]) -> tuple[list[int], list[float]]: @@ -322,6 +413,65 @@ def submit_generate_tasks(self, samples: list[list[Sample]]) -> None: self.remaining_batch_size += len(samples) +def _use_vllm_completions_for_r3(args: Namespace, sample: Sample, *, has_images: bool) -> bool: + """Use ``/v1/completions`` for MoE routing replay (v0.21.x returns prompt+gen routing there). + + ``/inference/v1/generate`` often exposes only decode rows on ``choices[].routed_experts``. + Partial continuation must keep the token-id generate API. + """ + if not getattr(args, "use_rollout_routing_replay", False): + return False + if has_images: + return False + if len(sample.response) > 0: + return False + return True + + +def _build_completion_request_body( + model: str, + prompt: str, + sampling_params: dict[str, Any], +) -> dict[str, Any]: + """Map rollout ``sampling_params`` to vLLM ``/v1/completions`` request body.""" + body: dict[str, Any] = { + "model": model, + "prompt": prompt, + "max_tokens": sampling_params["max_new_tokens"], + "temperature": sampling_params["temperature"], + "top_p": sampling_params["top_p"], + "logprobs": 1, + "return_token_ids": True, + "stream": False, + } + tk = sampling_params.get("top_k") + if tk is not None and tk > 0: + body["top_k"] = tk + if sampling_params.get("stop"): + body["stop"] = sampling_params["stop"] + if sampling_params.get("stop_token_ids"): + body["stop_token_ids"] = sampling_params["stop_token_ids"] + if sampling_params.get("seed") is not None: + body["seed"] = sampling_params["seed"] + return body + + +def _completion_tokens_and_logprobs(choice: dict[str, Any]) -> tuple[list[int], list[float]]: + """Parse ``token_ids`` and ``logprobs`` from a vLLM ``/v1/completions`` choice.""" + tids_raw = choice.get("token_ids") + if not isinstance(tids_raw, list) or not tids_raw: + return [], [] + tids = [int(x) for x in tids_raw] + lp = choice.get("logprobs") + if not isinstance(lp, dict): + return tids, [0.0] * len(tids) + token_logprobs = lp.get("token_logprobs") + if isinstance(token_logprobs, list) and len(token_logprobs) >= len(tids): + lps = [float(x) if x is not None else 0.0 for x in token_logprobs[-len(tids) :]] + return tids, lps + return tids, [0.0] * len(tids) + + def _build_inference_sampling_params(sampling_params: dict[str, Any]) -> dict[str, Any]: """Map rollout ``sampling_params`` to vLLM ``/inference/v1/generate`` ``sampling_params`` body.""" sp: dict[str, Any] = { @@ -421,6 +571,7 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A if getattr(args, "router_policy", None) == "consistent_hashing": headers = {"X-SMG-Routing-Key": sample.session_id} + used_completions_r3 = False if images: # Disaggregated MM flow: render (preprocess) then tokens-only generate — see vLLM docs # ``examples/online_serving/disaggregated_serving`` (``/v1/chat/completions/render`` + @@ -441,6 +592,18 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A gen_url = f"{base}/inference/v1/generate" with trace_span(sample, "vllm_mm_generate", attrs={"max_tokens": params["max_new_tokens"]}): output = await post(gen_url, generate_body, headers=headers) + request_prompt_len = len(generate_body.get("token_ids") or []) + elif _use_vllm_completions_for_r3(args, sample, has_images=False): + used_completions_r3 = True + completion_url = f"{base}/v1/completions" + completion_body = _build_completion_request_body( + args.hf_checkpoint, + sample.prompt, + params, + ) + with trace_span(sample, "vllm_completion", attrs={"max_tokens": params["max_new_tokens"]}): + output = await post(completion_url, completion_body, headers=headers) + request_prompt_len = len(prompt_ids) else: url = f"{base}/inference/v1/generate" # vLLM disaggregated ``/inference/v1/generate`` is token-only. On partial continuation, send the @@ -454,6 +617,7 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A "token_ids": token_ids, "sampling_params": inference_sampling_params, } + request_prompt_len = len(token_ids) with trace_span(sample, "vllm_inference_generate", attrs={"max_new_tokens": params["max_new_tokens"]}): output = await post(url, payload, headers=headers) @@ -461,13 +625,19 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A skip_sp = params.get("skip_special_tokens") skip_decode = True if skip_sp is None else bool(skip_sp) out_ids = choice.get("token_ids") or [] - text = ( - state.tokenizer.decode(out_ids, skip_special_tokens=skip_decode) - if isinstance(out_ids, list) and out_ids - else "" - ) + if used_completions_r3 and choice.get("text") is not None: + text = str(choice.get("text") or "") + else: + text = ( + state.tokenizer.decode(out_ids, skip_special_tokens=skip_decode) + if isinstance(out_ids, list) and out_ids + else "" + ) meta = _vllm_meta_from_generate_choice(args, choice, output.get("usage")) - new_response_tokens, new_response_log_probs = _inference_generate_tokens_and_logprobs(choice) + if used_completions_r3: + new_response_tokens, new_response_log_probs = _completion_tokens_and_logprobs(choice) + else: + new_response_tokens, new_response_log_probs = _inference_generate_tokens_and_logprobs(choice) new_response_tokens, new_response_log_probs = _align_engine_tokens_and_logprobs( new_response_tokens, new_response_log_probs ) @@ -491,7 +661,14 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A sample.rollout_log_probs = [] sample.rollout_log_probs += new_response_log_probs - _apply_vllm_routed_experts(args, sample, output, choice) + _apply_vllm_routed_experts( + args, + sample, + output, + choice, + gen_token_count=len(new_response_tokens) if new_response_tokens else None, + prompt_token_count=request_prompt_len, + ) sample.update_from_meta_info(args, meta) return sample diff --git a/tests/test_vllm_generate_endpoint.py b/tests/test_vllm_generate_endpoint.py index 36ab90bc8..32f8db1c5 100644 --- a/tests/test_vllm_generate_endpoint.py +++ b/tests/test_vllm_generate_endpoint.py @@ -181,7 +181,13 @@ def _execute_case(case: VLLMGenerateCase): assert len(sample.rollout_log_probs) == sample.response_length assert sample.status in (Sample.Status.COMPLETED, Sample.Status.TRUNCATED) if case.use_rollout_routing_replay: - assert sample.rollout_routed_experts is not None + re = sample.rollout_routed_experts + assert re is not None + assert re.ndim == 3 + expected_rows = len(sample.tokens) - 1 + assert ( + re.shape[0] == expected_rows + ), f"rollout_routed_experts rows {re.shape[0]} != len(tokens)-1 ({expected_rows})" finally: _stop_process_tree(process) diff --git a/tests/unit/rollout/test_vllm_rollout.py b/tests/unit/rollout/test_vllm_rollout.py index 0204122f0..c4e9149f0 100644 --- a/tests/unit/rollout/test_vllm_rollout.py +++ b/tests/unit/rollout/test_vllm_rollout.py @@ -226,30 +226,56 @@ def test_decode_vllm_routed_experts_roundtrip(): np.testing.assert_array_equal(decoded, arr) -@pytest.mark.unit -def test_apply_vllm_routed_experts_assigns_when_shape_matches(): - arr = np.zeros((3, 2, 1), dtype=np.int32) +def _encode_routed_npy(arr: np.ndarray) -> str: buf = io.BytesIO() np.save(buf, arr) - encoded = base64.b64encode(buf.getvalue()).decode("ascii") + return base64.b64encode(buf.getvalue()).decode("ascii") + + +@pytest.mark.unit +def test_merge_generate_routed_experts_trims_extra_gen_rows(): + prompt = np.ones((2, 2, 1), dtype=np.int32) + gen = np.full((4, 2, 1), 2, dtype=np.int32) + merged = mod._merge_generate_routed_experts( + {"prompt_routed_experts": _encode_routed_npy(prompt)}, + {"routed_experts": _encode_routed_npy(gen)}, + gen_token_count=2, + ) + np.testing.assert_array_equal(merged, np.concatenate([prompt, gen[:2]], axis=0)) + - sample = Sample(tokens=[1, 2, 3, 4]) +@pytest.mark.unit +def test_apply_vllm_routed_experts_merged_matches_sglang_layout(): + prompt = np.zeros((2, 2, 1), dtype=np.int32) + gen = np.zeros((2, 2, 1), dtype=np.int32) + sample = Sample(tokens=[1, 2, 50, 51]) args = Namespace(use_rollout_routing_replay=True, num_layers=2, moe_router_topk=1) - mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": encoded}) - np.testing.assert_array_equal(sample.rollout_routed_experts, arr) + mod._apply_vllm_routed_experts( + args, + sample, + {"prompt_routed_experts": _encode_routed_npy(prompt)}, + {"routed_experts": _encode_routed_npy(gen)}, + gen_token_count=2, + ) + np.testing.assert_array_equal(sample.rollout_routed_experts, np.zeros((3, 2, 1), dtype=np.int32)) @pytest.mark.unit -def test_apply_vllm_routed_experts_strips_prompt_row_when_n_tok_rows(): - arr = np.zeros((4, 2, 1), dtype=np.int32) - buf = io.BytesIO() - np.save(buf, arr) - encoded = base64.b64encode(buf.getvalue()).decode("ascii") +def test_apply_vllm_routed_experts_merges_prompt_and_gen_nested_lists(): + prompt_part = np.ones((2, 2, 1), dtype=np.int32) + gen_part = np.full((2, 2, 1), 2, dtype=np.int32) + expected = np.concatenate([prompt_part, gen_part], axis=0) - sample = Sample(tokens=[1, 2, 3, 4]) + sample = Sample(tokens=[10, 20, 30, 40, 50]) args = Namespace(use_rollout_routing_replay=True, num_layers=2, moe_router_topk=1) - mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": encoded}) - np.testing.assert_array_equal(sample.rollout_routed_experts, arr[:-1]) + mod._apply_vllm_routed_experts( + args, + sample, + {"prompt_routed_experts": prompt_part.tolist()}, + {"routed_experts": gen_part.tolist()}, + gen_token_count=2, + ) + np.testing.assert_array_equal(sample.rollout_routed_experts, expected) @pytest.mark.unit @@ -431,8 +457,8 @@ def test_apply_vllm_routed_experts_disabled_or_missing(): sample = Sample(tokens=[1, 2, 3]) mod._apply_vllm_routed_experts(Namespace(use_rollout_routing_replay=False), sample, {}, {}) assert sample.rollout_routed_experts is None - mod._apply_vllm_routed_experts(Namespace(use_rollout_routing_replay=True), sample, {}, {}) - assert sample.rollout_routed_experts is None + with pytest.raises(RuntimeError, match="routing replay"): + mod._apply_vllm_routed_experts(Namespace(use_rollout_routing_replay=True), sample, {}, {}) @pytest.mark.unit @@ -440,17 +466,27 @@ def test_apply_vllm_routed_experts_skips_bad_shape(): arr = np.zeros((2, 2), dtype=np.int32) sample = Sample(tokens=[1, 2, 3]) args = Namespace(use_rollout_routing_replay=True) - mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": _encode_routed(arr)}) - assert sample.rollout_routed_experts is None + with pytest.raises(RuntimeError, match="routing replay"): + mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": _encode_routed(arr)}) @pytest.mark.unit -def test_apply_vllm_routed_experts_skips_row_mismatch(): +def test_apply_vllm_routed_experts_trims_when_too_many_rows(): arr = np.zeros((9, 2, 1), dtype=np.int32) sample = Sample(tokens=[1, 2, 3]) args = Namespace(use_rollout_routing_replay=True, num_layers=2, moe_router_topk=1) mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": _encode_routed(arr)}) - assert sample.rollout_routed_experts is None + assert sample.rollout_routed_experts is not None + assert sample.rollout_routed_experts.shape == (2, 2, 1) + + +@pytest.mark.unit +def test_apply_vllm_routed_experts_raises_when_too_few_rows(): + arr = np.zeros((1, 2, 1), dtype=np.int32) + sample = Sample(tokens=[1, 2, 3]) + args = Namespace(use_rollout_routing_replay=True, num_layers=2, moe_router_topk=1) + with pytest.raises(RuntimeError, match="routing replay"): + mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": _encode_routed(arr)}) @pytest.mark.unit @@ -458,8 +494,8 @@ def test_apply_vllm_routed_experts_skips_layer_topk_mismatch(): arr = np.zeros((2, 3, 4), dtype=np.int32) sample = Sample(tokens=[1, 2, 3]) args = Namespace(use_rollout_routing_replay=True, num_layers=2, moe_router_topk=1) - mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": _encode_routed(arr)}) - assert sample.rollout_routed_experts is None + with pytest.raises(RuntimeError, match="routing replay"): + mod._apply_vllm_routed_experts(args, sample, {}, {"routed_experts": _encode_routed(arr)}) @pytest.mark.unit @@ -588,15 +624,19 @@ async def fake_post(url, payload, headers=None, **kwargs): @pytest.mark.unit def test_generate_applies_routed_experts(patch_generate_state, monkeypatch): - # After generate: 2 prompt + 2 response tokens => expected_rows = 3 - arr = np.zeros((3, 2, 1), dtype=np.int32) + # Fake tokenizer yields 3 prompt ids; +2 response => 5 tokens, 4 routing rows. + prompt_rows = np.ones((2, 2, 1), dtype=np.int32) + gen_rows = np.full((2, 2, 1), 2, dtype=np.int32) + expected = np.concatenate([prompt_rows, gen_rows], axis=0) + post_mock = AsyncMock( return_value={ + "prompt_routed_experts": _encode_routed_npy(prompt_rows), "choices": [ { "token_ids": [50, 51], "finish_reason": "stop", - "routed_experts": _encode_routed(arr), + "routed_experts": _encode_routed_npy(gen_rows), "logprobs": {"content": [{}, {}]}, } ], @@ -605,7 +645,8 @@ def test_generate_applies_routed_experts(patch_generate_state, monkeypatch): ) monkeypatch.setattr(mod, "post", post_mock) - sample = Sample(index=0, prompt="ab") + # _FakeTokenizer encodes up to 3 chars => 3 prompt ids + 2 response = 5 tokens. + sample = Sample(index=0, prompt="abc") asyncio.run( mod.generate( _rollout_args(use_rollout_routing_replay=True, num_layers=2, moe_router_topk=1), @@ -613,7 +654,9 @@ def test_generate_applies_routed_experts(patch_generate_state, monkeypatch): _default_sampling_params(max_new_tokens=4), ) ) - np.testing.assert_array_equal(sample.rollout_routed_experts, arr) + np.testing.assert_array_equal(sample.rollout_routed_experts, expected) + assert len(sample.tokens) == 5 + assert sample.rollout_routed_experts.shape[0] == len(sample.tokens) - 1 @pytest.mark.unit