diff --git a/miles/backends/sglang_utils/arguments.py b/miles/backends/sglang_utils/arguments.py index 40853c8d39..4a5fbc12be 100644 --- a/miles/backends/sglang_utils/arguments.py +++ b/miles/backends/sglang_utils/arguments.py @@ -143,10 +143,6 @@ def validate_args(args): if args.true_on_policy_mode: args.sglang_enable_deterministic_inference = True - if getattr(args, "recompute_logprobs_via_prefill", False): - args.sglang_enable_prefill_only_deterministic_inference = True - args.sglang_enable_deterministic_inference = True - if args.sglang_dp_size > 1: assert args.sglang_enable_dp_attention diff --git a/miles/rollout/generate_utils/prefill_logprobs.py b/miles/rollout/generate_utils/prefill_logprobs.py deleted file mode 100644 index 680cf35378..0000000000 --- a/miles/rollout/generate_utils/prefill_logprobs.py +++ /dev/null @@ -1,176 +0,0 @@ -from __future__ import annotations - -from collections import defaultdict -from collections.abc import Mapping -from typing import Any - -from miles.backends.megatron_utils.lora_utils import LORA_ADAPTER_NAME, is_lora_enabled -from miles.utils.http_utils import post -from miles.utils.processing_utils import encode_image_for_rollout_engine -from miles.utils.types import Sample - - -def _build_prefill_scoring_payload( - args: Any, - sample: Sample, - sampling_params: Mapping[str, Any], -) -> dict[str, Any]: - prompt_len = len(sample.tokens) - sample.response_length - if prompt_len <= 0: - raise ValueError( - "Cannot recompute rollout logprobs via prefill without at least one prompt token: " - f"tokens={len(sample.tokens)}, response_length={sample.response_length}" - ) - - payload = { - "input_ids": sample.tokens, - "sampling_params": { - **dict(sampling_params), - "max_new_tokens": 0, - "temperature": 0, - "skip_special_tokens": False, - }, - "return_logprob": True, - # SGLang returns input_token_logprobs aligned to tokens from logprob_start_len, - # with the first value None. Start one token before the response so the - # returned tail contains every response-token logprob. - "logprob_start_len": prompt_len - 1, - } - - if is_lora_enabled(args): - payload["lora_path"] = LORA_ADAPTER_NAME - - if sample.multimodal_inputs and sample.multimodal_inputs.get("images"): - image_data = sample.multimodal_inputs["images"] - payload["image_data"] = [encode_image_for_rollout_engine(image) for image in image_data] - - return payload - - -def _can_batch_prefill_score(args: Any, samples: list[Sample]) -> bool: - if getattr(args, "sglang_router_policy", None) == "consistent_hashing": - return False - return not any(sample.multimodal_inputs and sample.multimodal_inputs.get("images") for sample in samples) - - -def _build_batch_prefill_scoring_payload( - args: Any, - samples: list[Sample], - sampling_params: Mapping[str, Any], -) -> dict[str, Any]: - payloads = [_build_prefill_scoring_payload(args, sample, sampling_params) for sample in samples] - logprob_start_len = payloads[0]["logprob_start_len"] - if any(payload["logprob_start_len"] != logprob_start_len for payload in payloads): - raise ValueError("Batched SGLang prefill scoring requires a shared logprob_start_len") - - batch_payload: dict[str, Any] = { - "input_ids": [payload["input_ids"] for payload in payloads], - "sampling_params": payloads[0]["sampling_params"], - "return_logprob": True, - "logprob_start_len": logprob_start_len, - } - if "lora_path" in payloads[0]: - batch_payload["lora_path"] = payloads[0]["lora_path"] - return batch_payload - - -def _extract_response_logprobs(sample: Sample, meta_info: Mapping[str, Any]) -> list[float]: - input_token_logprobs = meta_info.get("input_token_logprobs") - if not input_token_logprobs: - raise ValueError("SGLang prefill scoring response did not include input_token_logprobs") - - response_items = input_token_logprobs[-sample.response_length :] - response_tokens = sample.tokens[-sample.response_length :] - scored_tokens = [item[1] for item in response_items] - if scored_tokens != response_tokens: - raise ValueError( - "SGLang prefill scoring token alignment mismatch: " - f"expected response tail {response_tokens[:8]}... len={len(response_tokens)}, " - f"got {scored_tokens[:8]}... len={len(scored_tokens)}" - ) - - response_logprobs = [item[0] for item in response_items] - if any(logprob is None for logprob in response_logprobs): - raise ValueError("SGLang prefill scoring returned None for a response-token logprob") - - return response_logprobs - - -async def recompute_rollout_logprobs_via_prefill( - args: Any, - sample: Sample, - *, - url: str, - sampling_params: Mapping[str, Any], - headers: Mapping[str, str] | None = None, -) -> None: - if not getattr(args, "recompute_logprobs_via_prefill", False): - return - if sample.response_length == 0: - sample.rollout_log_probs = [] - return - if sample.status == Sample.Status.ABORTED: - return - - payload = _build_prefill_scoring_payload(args, sample, sampling_params) - output = await post(url, payload, headers=headers) - sample.rollout_log_probs = _extract_response_logprobs(sample, output["meta_info"]) - sample.metadata["rollout_log_probs_source"] = "sglang_prefill_recompute" - - -async def recompute_samples_rollout_logprobs_via_prefill( - args: Any, - samples: list[Sample], - *, - url: str, - sampling_params: Mapping[str, Any], -) -> None: - if not getattr(args, "recompute_logprobs_via_prefill", False): - return - - samples_to_score = [ - sample for sample in samples if sample.response_length != 0 and sample.status != Sample.Status.ABORTED - ] - if not samples_to_score: - return - - flush_url = url.rsplit("/", 1)[0] + "/flush_cache" - - if _can_batch_prefill_score(args, samples_to_score): - samples_by_logprob_start_len: dict[int, list[Sample]] = defaultdict(list) - for sample in samples_to_score: - prompt_len = len(sample.tokens) - sample.response_length - samples_by_logprob_start_len[prompt_len - 1].append(sample) - - for batch_samples in samples_by_logprob_start_len.values(): - # SGLang can serve scoring requests from radix/KV cache. Flush before - # each scoring group so every group uses the same clean-prefill path. - await post(flush_url, {}) - payload = _build_batch_prefill_scoring_payload(args, batch_samples, sampling_params) - outputs = await post(url, payload) - if not isinstance(outputs, list): - raise ValueError(f"SGLang batch prefill scoring returned {type(outputs).__name__}, expected list") - if len(outputs) != len(batch_samples): - raise ValueError( - "SGLang batch prefill scoring output count mismatch: " - f"expected {len(batch_samples)}, got {len(outputs)}" - ) - for sample, output in zip(batch_samples, outputs, strict=True): - sample.rollout_log_probs = _extract_response_logprobs(sample, output["meta_info"]) - sample.metadata["rollout_log_probs_source"] = "sglang_prefill_recompute" - return - - for sample in samples_to_score: - headers = None - uses_consistent_hashing = getattr(args, "sglang_router_policy", None) == "consistent_hashing" - if uses_consistent_hashing and sample.session_id: - headers = {"X-SMG-Routing-Key": sample.session_id} - - await post(flush_url, {}, headers=headers) - await recompute_rollout_logprobs_via_prefill( - args, - sample, - url=url, - sampling_params=sampling_params, - headers=headers, - ) diff --git a/miles/rollout/inference_rollout/inference_rollout_train.py b/miles/rollout/inference_rollout/inference_rollout_train.py index eed15557bd..fbf909b374 100644 --- a/miles/rollout/inference_rollout/inference_rollout_train.py +++ b/miles/rollout/inference_rollout/inference_rollout_train.py @@ -9,7 +9,6 @@ from miles.rollout.base_types import RolloutFnTrainOutput from miles.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter -from miles.rollout.generate_utils.prefill_logprobs import recompute_samples_rollout_logprobs_via_prefill from miles.rollout.inference_rollout.inference_rollout_common import GenerateState, generate_and_rm_group from miles.utils import dumper_utils from miles.utils.http_utils import get, post @@ -153,11 +152,4 @@ async def generate_rollout_async( if f := load_function(args.rollout_all_samples_process_path): f(args, all_samples, data_source) - await recompute_samples_rollout_logprobs_via_prefill( - args, - [sample for group in data for sample in group], - url=f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate", - sampling_params=state.sampling_params, - ) - return RolloutFnTrainOutput(samples=data, metrics=metric_gatherer.collect()), aborted_samples diff --git a/miles/rollout/sglang_rollout.py b/miles/rollout/sglang_rollout.py index c001665087..461ebdf314 100644 --- a/miles/rollout/sglang_rollout.py +++ b/miles/rollout/sglang_rollout.py @@ -31,7 +31,6 @@ ) from miles.utils.types import Sample -from .generate_utils.prefill_logprobs import recompute_samples_rollout_logprobs_via_prefill from .rm_hub import async_rm, batched_async_rm __all__ = ["generate_rollout", "get_model_url"] @@ -468,13 +467,6 @@ async def generate_rollout_async( process_func = load_function(args.rollout_all_samples_process_path) process_func(args, all_samples, data_source) - await recompute_samples_rollout_logprobs_via_prefill( - args, - [sample for group in data for sample in group], - url=get_model_url(args, "default"), - sampling_params=state.sampling_params, - ) - return RolloutFnTrainOutput(samples=data, metrics=metric_gatherer.collect()), aborted_samples diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index df8317ef5b..bf4ebdedb4 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -145,15 +145,6 @@ def add_train_arguments(parser): default=False, help="Whether to enable true-on-policy mode.", ) - parser.add_argument( - "--recompute-logprobs-via-prefill", - action="store_true", - default=False, - help=( - "Recompute rollout logprobs via SGLang prefill instead of decode kernels. " - "Only needed for models whose prefill and decode paths are not numerically identical." - ), - ) parser.add_argument( "--train-env-vars", type=json.loads, @@ -1882,9 +1873,6 @@ def _resolve_eval_datasets(args) -> list[EvalDatasetConfig]: def miles_validate_args(args): args.eval_datasets = _resolve_eval_datasets(args) - if args.recompute_logprobs_via_prefill: - assert args.true_on_policy_mode, "--recompute-logprobs-via-prefill requires --true-on-policy-mode" - # Normalize --tito-allowed-append-roles: lowercase + deduplicate. raw_roles = getattr(args, "tito_allowed_append_roles", ["tool"]) args.tito_allowed_append_roles = sorted(set(r.lower() for r in raw_roles)) diff --git a/tests/fast/rollout/generate_utils/test_prefill_logprobs.py b/tests/fast/rollout/generate_utils/test_prefill_logprobs.py deleted file mode 100644 index aa736a8194..0000000000 --- a/tests/fast/rollout/generate_utils/test_prefill_logprobs.py +++ /dev/null @@ -1,169 +0,0 @@ -from types import SimpleNamespace - -import pytest - -from miles.rollout.generate_utils import prefill_logprobs -from miles.utils.types import Sample - - -@pytest.mark.asyncio -async def test_recompute_rollout_logprobs_via_prefill_uses_response_tail(monkeypatch): - sample = Sample( - tokens=[10, 11, 12, 20, 21, 22], - response_length=3, - rollout_log_probs=[-9.0, -9.0, -9.0], - status=Sample.Status.COMPLETED, - ) - args = SimpleNamespace(recompute_logprobs_via_prefill=True, sglang_enable_lora=False) - seen = {} - - async def fake_post(url, payload, headers=None): - seen["url"] = url - seen["payload"] = payload - seen["headers"] = headers - return { - "meta_info": { - "input_token_logprobs": [ - (None, 12), - (-0.1, 20), - (-0.2, 21), - (-0.3, 22), - ] - } - } - - monkeypatch.setattr(prefill_logprobs, "post", fake_post) - - await prefill_logprobs.recompute_rollout_logprobs_via_prefill( - args, - sample, - url="http://localhost/generate", - sampling_params={"temperature": 1, "max_new_tokens": 128}, - headers={"X-Test": "1"}, - ) - - assert sample.rollout_log_probs == [-0.1, -0.2, -0.3] - assert sample.metadata["rollout_log_probs_source"] == "sglang_prefill_recompute" - assert seen["url"] == "http://localhost/generate" - assert seen["headers"] == {"X-Test": "1"} - assert seen["payload"]["input_ids"] == sample.tokens - assert seen["payload"]["return_logprob"] is True - assert seen["payload"]["logprob_start_len"] == 2 - assert seen["payload"]["sampling_params"]["max_new_tokens"] == 0 - assert seen["payload"]["sampling_params"]["temperature"] == 0 - - -@pytest.mark.asyncio -async def test_recompute_rollout_logprobs_via_prefill_checks_token_alignment(monkeypatch): - sample = Sample(tokens=[10, 11, 20], response_length=1, status=Sample.Status.COMPLETED) - args = SimpleNamespace(recompute_logprobs_via_prefill=True, sglang_enable_lora=False) - - async def fake_post(url, payload, headers=None): - return {"meta_info": {"input_token_logprobs": [(None, 11), (-0.1, 999)]}} - - monkeypatch.setattr(prefill_logprobs, "post", fake_post) - - with pytest.raises(ValueError, match="token alignment mismatch"): - await prefill_logprobs.recompute_rollout_logprobs_via_prefill( - args, - sample, - url="http://localhost/generate", - sampling_params={}, - ) - - -@pytest.mark.asyncio -async def test_recompute_samples_flushes_each_batch_and_batches_prefill_score(monkeypatch): - samples = [ - Sample(tokens=[10, 11, 20], response_length=1, status=Sample.Status.COMPLETED), - Sample(tokens=[10, 11, 21], response_length=1, status=Sample.Status.COMPLETED), - ] - args = SimpleNamespace( - recompute_logprobs_via_prefill=True, - sglang_enable_lora=False, - sglang_router_policy="round_robin", - ) - calls = [] - - async def fake_post(url, payload, action="post", headers=None): - calls.append((url, payload, action, headers)) - if url.endswith("/flush_cache"): - return {} - return [ - {"meta_info": {"input_token_logprobs": [(None, 11), (-float(tokens[-1]), tokens[-1])]}} - for tokens in payload["input_ids"] - ] - - monkeypatch.setattr(prefill_logprobs, "post", fake_post) - - await prefill_logprobs.recompute_samples_rollout_logprobs_via_prefill( - args, - samples, - url="http://localhost/generate", - sampling_params={"max_new_tokens": 32}, - ) - - assert [sample.rollout_log_probs for sample in samples] == [[-20.0], [-21.0]] - assert [call[0] for call in calls] == [ - "http://localhost/flush_cache", - "http://localhost/generate", - ] - assert [call[2] for call in calls] == ["post", "post"] - assert calls[1][1]["input_ids"] == [[10, 11, 20], [10, 11, 21]] - assert calls[1][1]["logprob_start_len"] == 1 - - -@pytest.mark.asyncio -async def test_recompute_samples_batches_by_logprob_start_len(monkeypatch): - samples = [ - Sample(tokens=[10, 11, 20], response_length=1, status=Sample.Status.COMPLETED), - Sample(tokens=[10, 11, 12, 21], response_length=1, status=Sample.Status.COMPLETED), - Sample(tokens=[10, 11, 22], response_length=1, status=Sample.Status.COMPLETED), - ] - args = SimpleNamespace( - recompute_logprobs_via_prefill=True, - sglang_enable_lora=False, - sglang_router_policy="round_robin", - ) - calls = [] - - async def fake_post(url, payload, action="post", headers=None): - calls.append((url, payload, action, headers)) - if url.endswith("/flush_cache"): - return {} - return [ - { - "meta_info": { - "input_token_logprobs": [ - (None, tokens[-2]), - (-float(tokens[-1]), tokens[-1]), - ] - } - } - for tokens in payload["input_ids"] - ] - - monkeypatch.setattr(prefill_logprobs, "post", fake_post) - - await prefill_logprobs.recompute_samples_rollout_logprobs_via_prefill( - args, - samples, - url="http://localhost/generate", - sampling_params={"max_new_tokens": 32}, - ) - - assert [sample.rollout_log_probs for sample in samples] == [ - [-20.0], - [-21.0], - [-22.0], - ] - assert [call[0] for call in calls] == [ - "http://localhost/flush_cache", - "http://localhost/generate", - "http://localhost/flush_cache", - "http://localhost/generate", - ] - assert calls[1][1]["logprob_start_len"] == 1 - assert calls[1][1]["input_ids"] == [[10, 11, 20], [10, 11, 22]] - assert calls[3][1]["logprob_start_len"] == 2 - assert calls[3][1]["input_ids"] == [[10, 11, 12, 21]] diff --git a/tests/fast/true_on_policy/test_config.py b/tests/fast/true_on_policy/test_config.py index 46562896ad..a89df18440 100644 --- a/tests/fast/true_on_policy/test_config.py +++ b/tests/fast/true_on_policy/test_config.py @@ -354,7 +354,6 @@ def test_megatron_true_on_policy_keeps_sequence_parallel_and_enables_backend_fla assert "--use-sglang" not in plan.train_args assert "--true-on-policy-contract qwen3_dense_true_on_policy_v1" in plan.train_args assert "--sglang-true-on-policy-contract qwen3_dense_true_on_policy_v1" in plan.train_args - assert "--recompute-logprobs-via-prefill" not in plan.train_args assert "--batch-invariant-mode" in plan.train_args assert "--no-rope-fusion" in plan.train_args assert "ROW_LINEAR_ENABLE_INV" not in plan.env_vars diff --git a/tests/fast/true_on_policy/test_run_qwen3_30b_a3b.py b/tests/fast/true_on_policy/test_run_qwen3_30b_a3b.py index afdce91aa4..022f7c692f 100644 --- a/tests/fast/true_on_policy/test_run_qwen3_30b_a3b.py +++ b/tests/fast/true_on_policy/test_run_qwen3_30b_a3b.py @@ -49,7 +49,6 @@ def fake_execute_train(**kwargs): assert "--sglang-enable-dp-attention" not in train_args assert "--sglang-true-on-policy-contract qwen3_moe_true_on_policy_v1" in train_args assert "--true-on-policy-contract qwen3_moe_true_on_policy_v1" in train_args - assert "--recompute-logprobs-via-prefill" not in train_args assert "--sequence-parallel" not in train_args assert "--no-gradient-accumulation-fusion" not in train_args assert "--use-sglang" not in train_args diff --git a/tests/fast/true_on_policy/test_run_qwen3_4b.py b/tests/fast/true_on_policy/test_run_qwen3_4b.py index 2268654bce..d78157ab6f 100644 --- a/tests/fast/true_on_policy/test_run_qwen3_4b.py +++ b/tests/fast/true_on_policy/test_run_qwen3_4b.py @@ -32,7 +32,6 @@ def fake_execute_train(**kwargs): assert "--sglang-true-on-policy-contract qwen3_dense_true_on_policy_v1" in train_args assert "--sglang-attention-backend fa3" in train_args assert "--true-on-policy-contract qwen3_dense_true_on_policy_v1" in train_args - assert "--recompute-logprobs-via-prefill" not in train_args assert "--load /root/models/Qwen3-4B_torch_dist" in train_args assert "--save /root/shared_data/unit-test/checkpoints" in train_args assert "--use-sglang" not in train_args @@ -83,7 +82,6 @@ def fake_execute_train(**kwargs): assert "--save /root/shared_data/unit-test-tp2-cp4/checkpoints" in train_args assert "--sglang-true-on-policy-contract qwen3_dense_true_on_policy_v1" in train_args assert "--sglang-attention-backend fa3" in train_args - assert "--recompute-logprobs-via-prefill" not in train_args assert "--use-sglang" not in train_args assert "--batch-invariant-mode" in train_args assert "--no-bias-swiglu-fusion" in train_args @@ -122,6 +120,5 @@ def fake_execute_train(**kwargs): assert "--true-on-policy-mode" not in train_args assert "--sglang-true-on-policy-contract" not in train_args assert "--true-on-policy-contract" not in train_args - assert "--recompute-logprobs-via-prefill" not in train_args assert "--use-sglang" not in train_args assert "ROW_LINEAR_ENABLE_INV" not in env_vars diff --git a/tests/fast/utils/test_arguments.py b/tests/fast/utils/test_arguments.py index 9f1ebc8eff..6e37080638 100644 --- a/tests/fast/utils/test_arguments.py +++ b/tests/fast/utils/test_arguments.py @@ -140,15 +140,6 @@ def test_respects_start_rollout_id(self) -> None: assert args.num_rollout == 6 -def test_recompute_logprobs_via_prefill_flag_is_parsed(): - parser = argparse.ArgumentParser() - get_miles_extra_args_provider()(parser) - - args = parser.parse_args(["--recompute-logprobs-via-prefill"] + REQUIRED_ARGS) - - assert args.recompute_logprobs_via_prefill is True - - @pytest.mark.parametrize( ( "rollout_num_gpus_per_engine", @@ -174,7 +165,6 @@ def test_true_on_policy_args_propagate_to_sglang_server_args( sglang_router_policy=None, sglang_router_ip=None, true_on_policy_mode=True, - recompute_logprobs_via_prefill=False, sglang_true_on_policy_contract="qwen3_dense_true_on_policy_v1", sglang_enable_deterministic_inference=False, sglang_enable_prefill_only_deterministic_inference=False,