Skip to content
Merged
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
10 changes: 4 additions & 6 deletions miles/backends/sglang_utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -147,12 +147,10 @@ def validate_args(args):
if args.sglang_dp_size > 1:
assert args.sglang_enable_dp_attention

if args.sglang_router_policy:
from miles.utils.environ import enable_experimental_rollout_refactor

assert (
not enable_experimental_rollout_refactor()
), "--sglang-router-policy is not supported with MILES_EXPERIMENTAL_ROLLOUT_REFACTOR=1"
if args.sglang_router_policy is None and args.use_session_server:
args.sglang_router_policy = "manual"
if args.router_assignment_mode == "random":
args.router_assignment_mode = "min_load"

if getattr(args, "sglang_router_ip", None):
args.sglang_router_ip = _wrap_ipv6(args.sglang_router_ip)
3 changes: 2 additions & 1 deletion miles/rollout/generate_hub/multi_turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from miles.rollout.generate_utils.generate_endpoint_utils import (
compute_prompt_ids_from_sample,
compute_request_payload,
compute_routing_headers,
update_sample_from_response,
)
from miles.rollout.generate_utils.tool_call_utils import (
Expand Down Expand Up @@ -56,7 +57,7 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput:
if args.generate_multi_samples:
sample = deepcopy(input.sample)

output = await post(url, payload)
output = await post(url, payload, headers=compute_routing_headers(args, sample))
await update_sample_from_response(args, sample, payload=payload, output=output, update_loss_mask=True)

if args.generate_multi_samples:
Expand Down
3 changes: 2 additions & 1 deletion miles/rollout/generate_hub/single_turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from miles.rollout.generate_utils.generate_endpoint_utils import (
compute_prompt_ids_from_sample,
compute_request_payload,
compute_routing_headers,
update_sample_from_response,
)
from miles.utils.http_utils import post
Expand Down Expand Up @@ -40,7 +41,7 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput:
sample.status = halt_status
return GenerateFnOutput(samples=sample)

output = await post(url, payload)
output = await post(url, payload, headers=compute_routing_headers(args, sample))
await update_sample_from_response(args, sample, payload=payload, output=output)

return GenerateFnOutput(samples=sample)
15 changes: 15 additions & 0 deletions miles/rollout/generate_utils/generate_endpoint_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,21 @@ def compute_prompt_ids_from_sample(state, sample, tools=None):
return state.tokenizer.encode(prompt, add_special_tokens=False)


def policy_uses_routing_key(args) -> bool:
return args.sglang_router_policy in ("consistent_hashing", "manual")


def compute_routing_headers(args, sample: Sample) -> dict[str, str] | None:
if policy_uses_routing_key(args) and not sample.routing_key:
raise ValueError(
f"router policy {args.sglang_router_policy} routes by X-SMG-Routing-Key, "
f"but sample (index={sample.index}) has no routing_key set"
)
if sample.routing_key:
return {"X-SMG-Routing-Key": sample.routing_key}
return None


def compute_request_payload(
args,
input_ids: list[int],
Expand Down
8 changes: 3 additions & 5 deletions miles/rollout/generate_utils/prefill_logprobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from collections.abc import Mapping
from typing import Any

from miles.rollout.generate_utils.generate_endpoint_utils import compute_routing_headers, policy_uses_routing_key
from miles.utils.http_utils import post
from miles.utils.lora import LORA_ADAPTER_NAME, is_lora_enabled
from miles.utils.processing_utils import encode_image_for_rollout_engine
Expand Down Expand Up @@ -48,7 +49,7 @@ def _build_prefill_scoring_payload(


def _can_batch_prefill_score(args: Any, samples: list[Sample]) -> bool:
if getattr(args, "sglang_router_policy", None) == "consistent_hashing":
if policy_uses_routing_key(args):
return False
return not any(sample.multimodal_inputs and sample.multimodal_inputs.get("images") for sample in samples)

Expand Down Expand Up @@ -161,10 +162,7 @@ async def recompute_samples_rollout_logprobs_via_prefill(
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}
headers = compute_routing_headers(args, sample)

await post(flush_url, {}, headers=headers)
await recompute_rollout_logprobs_via_prefill(
Expand Down
2 changes: 1 addition & 1 deletion miles/rollout/generate_utils/sample_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,7 +142,7 @@ def _merge_metadata():
metadata=_merge_metadata(),
generate_function_path=_merge_equal_value("generate_function_path"),
train_metadata=_merge_equal_value("train_metadata"),
session_id=_merge_equal_value("session_id"),
routing_key=_merge_equal_value("routing_key"),
non_generation_time=_merge_equal_value("non_generation_time"),
spec_info=_merge_spec_info(a.spec_info, b.spec_info),
prefix_cache_info=_merge_prefix_cache_info(a.prefix_cache_info, b.prefix_cache_info),
Expand Down
7 changes: 7 additions & 0 deletions miles/rollout/inference_rollout/inference_rollout_common.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import asyncio
import logging
import uuid
from argparse import Namespace
from copy import deepcopy
from typing import Any
Expand All @@ -15,6 +16,7 @@
RolloutFnTrainOutput,
)
from miles.rollout.generate_hub.single_turn import generate
from miles.rollout.generate_utils.generate_endpoint_utils import policy_uses_routing_key
from miles.rollout.inference_rollout.compatibility import load_generate_function
from miles.rollout.rm_hub import async_rm, batched_async_rm
from miles.utils.processing_utils import load_processor, load_tokenizer
Expand Down Expand Up @@ -125,6 +127,11 @@ async def generate_and_rm_group(
if state.aborted:
return group

if policy_uses_routing_key(args):
for sample in group:
if sample.routing_key is None:
sample.routing_key = str(uuid.uuid4())

log_prefix = f"[group indices={[getattr(s, 'index', '?') for s in group]}]"
logger.debug(f"{log_prefix} Starting group with {len(group)} samples")
tasks = []
Expand Down
4 changes: 4 additions & 0 deletions miles/rollout/inference_rollout/inference_rollout_eval.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,12 @@
import asyncio
import copy
import logging
import uuid
from typing import Any

from tqdm import tqdm

from miles.rollout.generate_utils.generate_endpoint_utils import policy_uses_routing_key
from miles.rollout.inference_rollout.inference_rollout_common import (
GenerateState,
compute_sampling_params,
Expand Down Expand Up @@ -66,6 +68,8 @@ async def eval_rollout_single_dataset(
sample.index = sample_index
sample_index += 1
sample.metadata = dataset_cfg.inject_metadata(getattr(sample, "metadata", None))
if policy_uses_routing_key(args):
sample.routing_key = str(uuid.uuid4())
sampling_params = base_sampling_params
if getattr(args, "sglang_enable_deterministic_inference", False):
sampling_params = base_sampling_params.copy()
Expand Down
21 changes: 12 additions & 9 deletions miles/rollout/sglang_rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,11 @@
)
from miles.utils.types import Sample

from .generate_utils.generate_endpoint_utils import get_indexer_topk_from_response
from .generate_utils.generate_endpoint_utils import (
compute_routing_headers,
get_indexer_topk_from_response,
policy_uses_routing_key,
)
from .generate_utils.prefill_logprobs import recompute_samples_rollout_logprobs_via_prefill
from .rm_hub import async_rm, batched_async_rm

Expand Down Expand Up @@ -192,10 +196,7 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A
if not sample.tokens: # Initialize sample.tokens for the first turn
sample.tokens = prompt_ids

# Use session_id for consistent hashing routing if router uses consistent_hashing policy
headers = None
if args.sglang_router_policy == "consistent_hashing" and sample.session_id:
headers = {"X-SMG-Routing-Key": sample.session_id}
headers = compute_routing_headers(args, sample)

output = await post(url, payload, headers=headers)
if getattr(args, "use_opd", False) and opd_top_k > 0 and opd_top_k_strategy != "only-teacher":
Expand Down Expand Up @@ -313,11 +314,11 @@ async def generate_and_rm_group(
if state.aborted:
return group

# Generate a unique session_id for each sample in the group (consistent hashing only)
if args.sglang_router_policy == "consistent_hashing":
# Generate a unique routing_key for each sample in the group (routing-key policies only)
if policy_uses_routing_key(args):
for sample in group:
if sample.session_id is None:
sample.session_id = str(uuid.uuid4())
if sample.routing_key is None:
sample.routing_key = str(uuid.uuid4())

tasks = []
for idx, sample in enumerate(group):
Expand Down Expand Up @@ -565,6 +566,8 @@ async def eval_rollout_single_dataset(
sample_index += 1
sample.metadata = dataset_cfg.inject_metadata(getattr(sample, "metadata", None))
sample.generate_function_path = getattr(dataset_cfg, "custom_generate_function_path", None)
if policy_uses_routing_key(args):
sample.routing_key = str(uuid.uuid4())
sampling_params = base_sampling_params
if getattr(args, "sglang_enable_deterministic_inference", False):
sampling_params = base_sampling_params.copy()
Expand Down
7 changes: 3 additions & 4 deletions miles/utils/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,9 +52,8 @@ class Status(Enum):
# metadata used during training, e.g., what loss to use for this sample.
train_metadata: dict | None = None

# Session ID for consistent hashing routing (used when router policy is consistent_hashing)
# TODO: Its definition needs to merge with the session server's session id in the new rollout function.
session_id: str | None = None
# Per-sample routing key for the router's consistent_hashing policy (sent as X-SMG-Routing-Key)
routing_key: str | None = None

non_generation_time: float = 0.0 # time spent in non-generation steps

Expand Down Expand Up @@ -217,7 +216,7 @@ def reset_for_retry(self) -> None:
"""Reset generated outputs so the original prompt can be re-sampled.

Keeps identity / prompt fields (group_index, index, prompt, label,
multimodal_inputs, metadata, generate_function_path, session_id) and
multimodal_inputs, metadata, generate_function_path, routing_key) and
restores everything else to dataclass defaults.
"""
self.tokens = []
Expand Down
Loading