Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 0 additions & 4 deletions miles/backends/sglang_utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
176 changes: 0 additions & 176 deletions miles/rollout/generate_utils/prefill_logprobs.py

This file was deleted.

8 changes: 0 additions & 8 deletions miles/rollout/inference_rollout/inference_rollout_train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
8 changes: 0 additions & 8 deletions miles/rollout/sglang_rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down Expand Up @@ -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


Expand Down
12 changes: 0 additions & 12 deletions miles/utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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))
Expand Down
Loading
Loading