diff --git a/areal/api/cli_args.py b/areal/api/cli_args.py index 210d5babbf..ccdf9cde1c 100644 --- a/areal/api/cli_args.py +++ b/areal/api/cli_args.py @@ -222,6 +222,10 @@ class GenerationHyperparameters: ) }, ) + seed: int | None = field( + default=None, + metadata={"help": "Per-request sampling seed sent to the inference backend."}, + ) lora_name: str = field( default="default_lora", metadata={"help": "Lora name to be used for this generation."}, @@ -2065,6 +2069,11 @@ def build_cmd( return vLLMConfig.build_cmd_from_args(args) +# Keep this list aligned with SGLang's deterministic inference documentation: +# https://docs.sglang.ai/advanced_features/deterministic_inference.html +_SGLANG_DETERMINISTIC_ATTENTION_BACKENDS = frozenset({"flashinfer", "fa3", "triton"}) + + @dataclass class SGLangConfig: """Configuration for SGLang runtime. Refer to: @@ -2098,6 +2107,7 @@ class SGLangConfig: enable_memory_saver: bool = False allow_auto_truncate: bool = False attention_backend: str | None = "fa3" + enable_deterministic_inference: bool = False enable_multimodal: bool = False sampling_backend: str | None = None context_length: int | None = 32768 @@ -2181,6 +2191,19 @@ def build_args( node_rank: int = 0, pp_size: int = 1, ): + attention_backend = sglang_config.attention_backend + if ( + sglang_config.enable_deterministic_inference + and attention_backend is not None + and attention_backend.lower() + not in _SGLANG_DETERMINISTIC_ATTENTION_BACKENDS + ): + logger.warning( + "SGLang deterministic inference is only documented for attention " + "backends %s; configured attention_backend=%r may be non-deterministic.", + sorted(_SGLANG_DETERMINISTIC_ATTENTION_BACKENDS), + attention_backend, + ) # Map "all-linear" to "all" args: dict = conf_as_dict(sglang_config) if sglang_config.enable_multithread_load: @@ -2415,6 +2438,25 @@ class InferenceEngineConfig: "help": "Whether to output verbose tracing messages for each generation request." }, ) + deterministic_sampling: bool = field( + default=False, + metadata={ + "help": "Use stable request seeds for internal OpenAI-proxy/data-proxy " + "sessions, canonical group ordering, and task-ID ordering of completed " + "rollout results. Concurrent SGLang generation also requires " + "sglang.enable_deterministic_inference. End-to-end determinism is only " + "supported with max_head_offpolicyness=0." + }, + ) + serialize_group_samples: bool = field( + default=False, + metadata={ + "help": "Run RolloutControllerV2 samples within each group sequentially " + "instead of concurrently. This provides stable within-group member " + "submission order at the cost of rollout throughput; it does not " + "serialize requests across groups." + }, + ) check_trajectory_format: bool = field( default=False, metadata={ @@ -2559,6 +2601,13 @@ def __post_init__(self): ) if not self.admin_api_key or not self.admin_api_key.strip(): raise ValueError("admin_api_key must not be empty or whitespace-only") + if self.deterministic_sampling and self.max_head_offpolicyness > 0: + logger.warning( + "deterministic_sampling=True with max_head_offpolicyness=%d does " + "not guarantee deterministic task-to-weight-version mapping; " + "set max_head_offpolicyness=0 for end-to-end determinism.", + self.max_head_offpolicyness, + ) if ( self._version == "v2" and self.agent is not None diff --git a/areal/engine/sglang_remote.py b/areal/engine/sglang_remote.py index cd54bc8a27..c37f6ebeaa 100644 --- a/areal/engine/sglang_remote.py +++ b/areal/engine/sglang_remote.py @@ -76,6 +76,8 @@ def build_generation_request( } if stop: sample_params["stop"] = stop + if gconfig.seed is not None: + sample_params["sampling_seed"] = gconfig.seed payload = { "input_ids": req.input_ids.copy(), diff --git a/areal/engine/vllm_remote.py b/areal/engine/vllm_remote.py index 2d76930a9d..1380c35a76 100644 --- a/areal/engine/vllm_remote.py +++ b/areal/engine/vllm_remote.py @@ -79,6 +79,8 @@ def build_generation_request( } if gconfig.stop: payload["stop"] = gconfig.stop + if gconfig.seed is not None: + payload["seed"] = gconfig.seed if with_lora: lora_name = gconfig.lora_name diff --git a/areal/experimental/openai/client.py b/areal/experimental/openai/client.py index c2edfb71c9..88137c8569 100644 --- a/areal/experimental/openai/client.py +++ b/areal/experimental/openai/client.py @@ -646,6 +646,7 @@ async def create( max_total_tokens: int | None | NotGiven = NOT_GIVEN, metadata: Metadata | None | NotGiven = NOT_GIVEN, n: int | None | NotGiven = NOT_GIVEN, + seed: int | None | NotGiven = NOT_GIVEN, stop: str | None | list[str] | None | NotGiven = NOT_GIVEN, store: bool | None | NotGiven = NOT_GIVEN, temperature: float | None | NotGiven = NOT_GIVEN, @@ -669,6 +670,7 @@ async def create( max_total_tokens: int | None | NotGiven = NOT_GIVEN, metadata: Metadata | None | NotGiven = NOT_GIVEN, n: int | None | NotGiven = NOT_GIVEN, + seed: int | None | NotGiven = NOT_GIVEN, stop: str | None | list[str] | None | NotGiven = NOT_GIVEN, store: bool | None | NotGiven = NOT_GIVEN, temperature: float | None | NotGiven = NOT_GIVEN, @@ -691,6 +693,7 @@ async def create( max_total_tokens: int | None | NotGiven = NOT_GIVEN, metadata: Metadata | None | NotGiven = NOT_GIVEN, n: int | None | NotGiven = NOT_GIVEN, + seed: int | None | NotGiven = NOT_GIVEN, stop: str | None | list[str] | None | NotGiven = NOT_GIVEN, store: bool | None | NotGiven = NOT_GIVEN, temperature: float | None | NotGiven = NOT_GIVEN, @@ -883,6 +886,7 @@ async def create( greedy=temp == 0, frequency_penalty=frequency_penalty, lora_name=self.lora_name, + seed=None if is_omitted(seed) else seed, stop_token_ids=list( set([self.tokenizer.eos_token_id, self.tokenizer.pad_token_id]) ), @@ -1136,6 +1140,7 @@ async def create( instructions: str | None | NotGiven = NOT_GIVEN, max_output_tokens: int | None | NotGiven = NOT_GIVEN, metadata: Metadata | None | NotGiven = NOT_GIVEN, + seed: int | None | NotGiven = NOT_GIVEN, tool_choice: response_create_params.ToolChoice | NotGiven = NOT_GIVEN, tools: Iterable[ToolParam] | NotGiven = NOT_GIVEN, temperature: float | None | NotGiven = NOT_GIVEN, @@ -1291,6 +1296,7 @@ async def create( greedy=temp == 0, frequency_penalty=frequency_penalty, lora_name=self.lora_name, + seed=None if is_omitted(seed) else seed, stop_token_ids=list( set([self.tokenizer.eos_token_id, self.tokenizer.pad_token_id]) ), diff --git a/areal/experimental/openai/proxy/proxy_rollout_server.py b/areal/experimental/openai/proxy/proxy_rollout_server.py index b42ecb6522..8acfeb3241 100644 --- a/areal/experimental/openai/proxy/proxy_rollout_server.py +++ b/areal/experimental/openai/proxy/proxy_rollout_server.py @@ -75,6 +75,10 @@ _warn_lock = threading.Lock() +def _deterministic_sampling_seed(session_id: str, request_index: int) -> int: + return seeding.derive_deterministic_seed(session_id, request_index) + + def _warn_once(msg: str) -> None: """Log a warning message, optionally only once if AREAL_PROXY_WARN_ONCE=1.""" if not _warn_once_enabled: @@ -126,6 +130,9 @@ def _warn_once(msg: str) -> None: _allocated_ports: set[int] = set() _port_alloc_lock = asyncio.Lock() +# Deterministic sampling (set from InferenceEngineConfig at setup time). +_deterministic_sampling: bool = False + # Server config (needed for name_resolve registration) _experiment_name: str | None = None _trial_name: str | None = None @@ -270,8 +277,9 @@ async def alloc_ports(raw_request: Request): def _setup_openai_client(): global _openai_client, _session_timeout_seconds, _admin_api_key - global _message_preprocessors, _prefix_matcher + global _message_preprocessors, _prefix_matcher, _deterministic_sampling config = _engine.config + _deterministic_sampling = bool(getattr(config, "deterministic_sampling", False)) tokenizer = load_hf_tokenizer(config.tokenizer_path) agent_cfg = config.agent _openai_client = ArealOpenAI( @@ -488,6 +496,7 @@ def start_session(request: StartSessionRequest) -> StartSessionResponse: _session_cache[session_id] = SessionData( session_id=session_id, prefix_matcher=_prefix_matcher, + sampling_seed_identity=task_id, ) _api_key_to_session[session_api_key] = session_id _session_to_api_key[session_id] = session_api_key @@ -574,8 +583,11 @@ async def _call_client_create( status_code=410, detail=f"Session {session_id} already ended or expired" ) session_data = _session_cache[session_id] + session_data.update_last_access() - session_data.update_last_access() + request_index = ( + session_data.next_sampling_request_index() if _deterministic_sampling else None + ) sig = inspect.signature(create_fn) areal_client_ignored_args = ["model"] + (extra_ignored_args or []) @@ -621,6 +633,22 @@ def _is_default_value(k: str, v: Any) -> bool: kwargs["top_p"] = 1.0 _warn_once("top_p not set in request, defaulting to 1.0") + if ( + _deterministic_sampling + and kwargs.get("seed") is None + and "seed" in areal_client_allowed_args + ): + assert request_index is not None + # The logical identity excludes the physical session collision suffix. + # Reserve request indices at ingress so concurrent requests remain + # distinct without holding a lock during inference. + # TODO(agent): Strict mapping of concurrent sibling requests to seeds + # requires a stable caller-provided request identity. Group samples use + # separate sessions, so their sample_idx-based identities are stable. + kwargs["seed"] = _deterministic_sampling_seed( + session_data.sampling_seed_identity, request_index + ) + # Strip stream from request body to prevent it from bypassing the explicit # `stream` parameter. Without this, a request with {"stream": true} would # leak through kwargs and cause the client to return an AsyncGenerator even diff --git a/areal/experimental/openai/proxy/server.py b/areal/experimental/openai/proxy/server.py index 9c7f6fe4f0..5bca619ebf 100644 --- a/areal/experimental/openai/proxy/server.py +++ b/areal/experimental/openai/proxy/server.py @@ -67,8 +67,14 @@ class ExportTrajectoriesResponse(BaseModel): class SessionData: """Data associated with a single RL session.""" - def __init__(self, session_id: str, prefix_matcher=None): + def __init__( + self, + session_id: str, + prefix_matcher=None, + sampling_seed_identity: str | None = None, + ): self.session_id = session_id + self.sampling_seed_identity = sampling_seed_identity or session_id self._completed = False self._completions = InteractionCache( @@ -80,6 +86,14 @@ def __init__(self, session_id: str, prefix_matcher=None): self._last_access_time = time.time() self._end_time = None self._lock = threading.Lock() + self._next_sampling_request_index = 0 + + def next_sampling_request_index(self) -> int: + """Reserve a unique request index without serializing request execution.""" + with self._lock: + request_index = self._next_sampling_request_index + self._next_sampling_request_index += 1 + return request_index def update_last_access(self): """Update the last access time for this session.""" diff --git a/areal/experimental/openai/proxy/workflow.py b/areal/experimental/openai/proxy/workflow.py index 862b43887b..887480c8fc 100644 --- a/areal/experimental/openai/proxy/workflow.py +++ b/areal/experimental/openai/proxy/workflow.py @@ -166,7 +166,15 @@ async def _grant_capacity(self, session: aiohttp.ClientSession) -> None: async def arun_episode( self, engine: TRolloutEngine, data: dict[str, Any] ) -> dict[str, InteractionWithTokenLogpReward] | None: - task_id = workflow_context.get().task_id + context = workflow_context.get() + task_id = context.task_id + # Qualify the proxy session with the group sample index so each group + # member owns a distinct, run-stable session namespace. + proxy_task_id = ( + f"{task_id}:{context.sample_idx}" + if context.sample_idx is not None + else str(task_id) + ) http_session = await workflow_context.get_aiohttp_session() @@ -190,7 +198,7 @@ async def arun_episode( proxy_client = OpenAIProxyClient( session=http_session, base_url=self.proxy_addr, - task_id=str(task_id), + task_id=proxy_task_id, admin_api_key=self._admin_api_key, ) proxy_client.session_id = session_info.session_id @@ -220,7 +228,7 @@ async def arun_episode( proxy_client = OpenAIProxyClient( session=http_session, base_url=self.proxy_addr, - task_id=str(task_id), + task_id=proxy_task_id, admin_api_key=self._admin_api_key, ) async with proxy_client: diff --git a/areal/infra/controller/rollout_controller.py b/areal/infra/controller/rollout_controller.py index 492393cee9..2a3a6df72f 100644 --- a/areal/infra/controller/rollout_controller.py +++ b/areal/infra/controller/rollout_controller.py @@ -222,6 +222,7 @@ def initialize( task_factory=self._create_submit_callback, staleness_manager=self._staleness_manager, enable_tracing=self.config.enable_rollout_tracing, + deterministic_order=getattr(self.config, "deterministic_sampling", False), ) # Initialize the dispatcher's async task runner self._dispatcher.initialize(logger=logger) @@ -960,6 +961,8 @@ def submit( # `arun_episode` should return None instead. if task_id is None: task_id = self._task_id_generator.next() + else: + self._task_id_generator.reserve_at_least(task_id) task_input = _RemoteRolloutTaskInput( data=data, workflow=workflow_str, diff --git a/areal/infra/remote_inf_engine.py b/areal/infra/remote_inf_engine.py index e6d5ae8143..ec75739d29 100644 --- a/areal/infra/remote_inf_engine.py +++ b/areal/infra/remote_inf_engine.py @@ -89,9 +89,32 @@ async def arun_episode( ) -> dict[str, Any] | None: from areal.experimental.openai import InteractionWithTokenLogpReward - results = await asyncio.gather( - *[self.workflow.arun_episode(engine, data) for _ in range(self.group_size)] + async def run_sample(sample_idx: int) -> tuple[int, Any]: + from areal.infra import workflow_context + from areal.infra.workflow_context import WorkflowContext + + parent = workflow_context.get() + workflow_context.set( + WorkflowContext( + is_eval=parent.is_eval, + task_id=parent.task_id, + sample_idx=sample_idx, + ) + ) + result = await self.workflow.arun_episode(engine, data) + return sample_idx, result + + indexed_results = await asyncio.gather( + *[run_sample(sample_idx) for sample_idx in range(self.group_size)] ) + indexed_results.sort(key=lambda item: item[0]) + sample_indices = [sample_idx for sample_idx, _ in indexed_results] + if sample_indices != list(range(self.group_size)): + raise RuntimeError( + "Grouped rollout returned invalid sample indices: " + f"expected {list(range(self.group_size))}, got {sample_indices}" + ) + results = [result for _, result in indexed_results] valid_results = [r for r in results if r is not None] @@ -1240,9 +1263,6 @@ def submit( raise ValueError( "workflow must be specified for submit (unless mode='online')" ) - if callback_addr: - self.workflow_executor.dispatcher.register_callback(task_id, callback_addr) - # Resolve workflow to a RolloutWorkflow instance resolved_workflow = self._resolve_workflow( workflow, @@ -1260,6 +1280,7 @@ def submit( should_accept_fn=resolved_should_accept_fn, task_id=task_id, is_eval=is_eval, + callback_addr=callback_addr, ) def wait( diff --git a/areal/infra/workflow_context.py b/areal/infra/workflow_context.py index beb9513807..c98d4e8618 100644 --- a/areal/infra/workflow_context.py +++ b/areal/infra/workflow_context.py @@ -31,10 +31,15 @@ class WorkflowContext: Whether the workflow is running in evaluation mode. task_id : int | None The task ID assigned by the workflow executor. + sample_idx : int | None + Index of this sample within its rollout group, when the workflow runs + under a grouped workflow. Gives group members a stable identity that + does not depend on completion order. """ is_eval: bool = False task_id: int | None = None + sample_idx: int | None = None _current_context: ContextVar[WorkflowContext] = ContextVar( diff --git a/areal/infra/workflow_executor.py b/areal/infra/workflow_executor.py index b2870bf8df..7a55c1ed9a 100644 --- a/areal/infra/workflow_executor.py +++ b/areal/infra/workflow_executor.py @@ -260,6 +260,27 @@ class WithTaskID(Protocol): TResult = TypeVar("TResult") +def _select_results( + drained: list, + count: int, + deterministic: bool, +) -> tuple[list, list]: + """Order drained results, then split them into (selected, pending). + + Normally results are taken oldest-first and the returned batch is + shuffled to avoid systematic ordering bias. Under deterministic sampling, + completed results are ordered by task ID and are not shuffled. + """ + if deterministic: + drained.sort(key=lambda x: x.task_id) + else: + drained.sort(key=lambda x: x.create_time) + selected, pending = drained[:count], drained[count:] + if not deterministic: + random.shuffle(selected) + return selected, pending + + class BatchTaskDispatcher(Generic[TInput, TResult]): """Generic dispatcher for asynchronous task execution with staleness control. @@ -279,6 +300,7 @@ def __init__( task_factory: Callable[[TInput], Callable[[], Awaitable[TResult | None]]], staleness_manager: StalenessManager, enable_tracing: bool = False, + deterministic_order: bool = False, ): self.runner = AsyncTaskRunner( max_queue_size=max_queue_size, @@ -287,12 +309,13 @@ def __init__( self.task_factory = task_factory self.staleness_manager = staleness_manager self.enable_tracing = enable_tracing + self.deterministic_order = deterministic_order self.logger: Logger # Unbounded deques for producer/consumer pattern self._pending_inputs: deque[TInput] = deque() self._pending_results: dict[int, TimedResult[TResult]] = {} - self._active_task_ids: set[int] = set() + self._active_task_ids: dict[int, None] = {} # Condition variables for coordination self._input_lock = threading.Lock() @@ -335,11 +358,19 @@ def _has_runner_capacity(self) -> bool: def register_callback(self, task_id: int, callback_addr: str): """Register a callback address for a task.""" - self._task_callbacks[task_id] = callback_addr + with self._result_cv: + if task_id in self._task_callbacks: + raise ValueError(f"Callback for task {task_id} is already registered") + self._task_callbacks[task_id] = callback_addr - def cancel_callback(self, task_id: int): + def cancel_callback(self, task_id: int, callback_addr: str | None = None): """Remove a registered callback for a task (e.g., on timeout).""" - self._task_callbacks.pop(task_id, None) + with self._result_cv: + if ( + callback_addr is None + or self._task_callbacks.get(task_id) == callback_addr + ): + self._task_callbacks.pop(task_id, None) def _send_callback(self, addr: str, task_id: int, result: TResult): """Send task result to callback address (fire-and-forget).""" @@ -535,14 +566,28 @@ def submit_task_input(self, task_input: TInput) -> None: Task input to be processed. """ self._check_thread_exception() - with self._input_cv: - self._pending_inputs.append(task_input) - self.staleness_manager.on_rollout_enqueued() - if self.enable_tracing: - self.logger.info(f"Enqueue rollout. {self._rollout_stats()}") - self._input_cv.notify() with self._result_cv: - self._active_task_ids.add(task_input.task_id) + if task_input.task_id in self._active_task_ids: + raise ValueError(f"Task id {task_input.task_id} is already active") + self._active_task_ids[task_input.task_id] = None + self._result_cv.notify_all() + try: + with self._input_cv: + self._pending_inputs.append(task_input) + try: + self.staleness_manager.on_rollout_enqueued() + except Exception: + removed = self._pending_inputs.pop() + assert removed is task_input + raise + self._input_cv.notify() + except Exception: + with self._result_cv: + self._active_task_ids.pop(task_input.task_id, None) + self._result_cv.notify_all() + raise + if self.enable_tracing: + self.logger.info(f"Enqueue rollout. {self._rollout_stats()}") def wait_results( self, count: int, timeout: float | None = None, raise_timeout: bool = True @@ -562,6 +607,7 @@ def wait_results( ------- list[TResult | None] List of task results, None for rejected tasks. + """ if count <= 0: raise ValueError(f"count must be positive, got {count}") @@ -573,6 +619,8 @@ def wait_results( with self._result_cv: while len(self._pending_results) < count: self._check_thread_exception() + if self._shutdown_event.is_set(): + raise RuntimeError("Task dispatcher is shutting down") elapsed = time.perf_counter() - start_time remaining = timeout - elapsed @@ -588,18 +636,17 @@ def wait_results( drained: list[TimedResult[TResult]] = list(self._pending_results.values()) self._pending_results.clear() - - drained.sort(key=lambda x: x.create_time) - selected, pending = drained[:count], drained[count:] - with self._result_cv: + selected, pending = _select_results( + drained, + count, + self.deterministic_order, + ) if pending: for result in pending: self._pending_results[result.task_id] = result - self._result_cv.notify_all() for r in selected: - self._active_task_ids.discard(r.task_id) - - random.shuffle(selected) + self._active_task_ids.pop(r.task_id, None) + self._result_cv.notify_all() return [r.data for r in selected] @@ -617,6 +664,10 @@ def wait_for_task( while task_id not in self._pending_results: self._check_thread_exception() + if self._shutdown_event.is_set(): + raise RuntimeError("Task dispatcher is shutting down") + if task_id not in self._active_task_ids: + raise RuntimeError(f"Task {task_id} was consumed by another waiter") elapsed = time.perf_counter() - start_time remaining = timeout - elapsed @@ -628,7 +679,7 @@ def wait_for_task( self._result_cv.wait(timeout=remaining) found_result = self._pending_results.pop(task_id) - self._active_task_ids.remove(task_id) + del self._active_task_ids[task_id] self._result_cv.notify_all() return found_result.data @@ -700,8 +751,9 @@ def active_submit_and_wait( "Input generator exhausted before batch completion. " "Use cycle_dataloader() or provide an infinite generator." ) from None + remaining = batch_size - (total_attempts if dynamic_bs else accepted_cnt) try: - arrived = self.wait_results(count=batch_size - accepted_cnt, timeout=1) + arrived = self.wait_results(count=remaining, timeout=1) except TimeoutError: arrived = [] @@ -743,6 +795,11 @@ def next(self): self._task_cnt += 1 return task_id + def reserve_at_least(self, task_id: int) -> None: + """Keep future automatic IDs above an explicitly supplied task ID.""" + with self._lock: + self._task_cnt = max(self._task_cnt, task_id + 1) + class WorkflowExecutor: """Executor for asynchronous workflow-based rollout generation. @@ -1066,6 +1123,7 @@ def initialize(self, logger=None, train_data_parallel_size: int | None = None): task_factory=self._create_workflow_task, staleness_manager=self._staleness_manager, enable_tracing=self.config.enable_rollout_tracing, + deterministic_order=getattr(self.config, "deterministic_sampling", False), ) # Initialize the dispatcher's async task runner @@ -1255,6 +1313,7 @@ def submit( should_accept_fn: Callable[[dict[str, Any]], bool] = None, task_id: int | None = None, is_eval: bool = False, + callback_addr: str | None = None, ) -> int: """Submit a rollout request to the workflow executor. @@ -1265,6 +1324,8 @@ def submit( """ if task_id is None: task_id = self._task_id_generator.next() + else: + self._task_id_generator.reserve_at_least(task_id) perf_tracer.register_task(task_id) task_input = _RolloutTaskInput( data=data, @@ -1274,8 +1335,14 @@ def submit( is_eval=is_eval, ) - # Delegate to dispatcher - self.dispatcher.submit_task_input(task_input) + if callback_addr is not None: + self.dispatcher.register_callback(task_id, callback_addr) + try: + self.dispatcher.submit_task_input(task_input) + except Exception: + if callback_addr is not None: + self.dispatcher.cancel_callback(task_id, callback_addr) + raise return task_id def wait( diff --git a/areal/utils/seeding.py b/areal/utils/seeding.py index d77862c5da..702b1b6acc 100644 --- a/areal/utils/seeding.py +++ b/areal/utils/seeding.py @@ -19,6 +19,14 @@ def _seed_from_key(key: str) -> int: return int(hashlib.sha256(key.encode()).hexdigest(), 16) & 0xFFFFFFFF +def derive_deterministic_seed(identity: str, request_index: int) -> int: + """Derive a stable non-negative sampling seed from a logical request identity.""" + if request_index < 0: + raise ValueError(f"request_index must be non-negative, got {request_index}") + digest = hashlib.sha256(f"{identity}:{request_index}".encode()).digest() + return int.from_bytes(digest[:4], byteorder="big", signed=False) & 0x7FFFFFFF + + def set_random_seed(base_seed: int, key: str) -> None: global _SEED, _BASE_SEED _BASE_SEED = base_seed diff --git a/areal/v2/inference_service/controller/controller.py b/areal/v2/inference_service/controller/controller.py index 6388bb7c14..3efaedc57a 100644 --- a/areal/v2/inference_service/controller/controller.py +++ b/areal/v2/inference_service/controller/controller.py @@ -477,6 +477,8 @@ async def _async_initialize( "--engine-max-tokens", str(agent_cfg.engine_max_tokens), ] + if cfg.deterministic_sampling: + data_proxy_base_cmd.append("--deterministic-sampling") async def _fork_data_proxy(group_idx: int) -> tuple[str, int, str]: if self.external_mode: @@ -1592,6 +1594,7 @@ def _wrap_agent(self, agent: Any, group_size: int = 1): discount=turn_discount, export_style=export_style, group_size=group_size, + serialize_group_samples=self.config.serialize_group_samples, ) def _resolve_workflow( diff --git a/areal/v2/inference_service/controller/workflow.py b/areal/v2/inference_service/controller/workflow.py index dd51256ba7..eea86a7aff 100644 --- a/areal/v2/inference_service/controller/workflow.py +++ b/areal/v2/inference_service/controller/workflow.py @@ -53,6 +53,7 @@ def __init__( export_style: str = "individual", timeout: float | None = None, group_size: int = 1, + serialize_group_samples: bool = False, ): self.controller = controller self.agent = agent @@ -62,6 +63,7 @@ def __init__( self.export_style = export_style self.timeout = timeout self.group_size = group_size + self.serialize_group_samples = serialize_group_samples @async_http_retry async def _start_session( @@ -144,12 +146,38 @@ async def _run_offline( group_id, sessions = await self._start_session( http_session, str(task_id), group_size=self.group_size ) + execution_mode = "serial" if self.serialize_group_samples else "concurrent" + version = self.controller.get_version() + logger.info( + "V2 rollout group dispatch: task_id=%s group_id=%s version=%s " + "group_size=%d mode=%s sessions=%s", + task_id, + group_id, + version, + len(sessions), + execution_mode, + [session_id for session_id, _ in sessions], + ) assert self.agent is not None http_client = await workflow_context.get_httpx_client() - async def _run_one(session_id: str, session_api_key: str) -> float | None: + async def _run_one( + member_index: int, + session_id: str, + session_api_key: str, + ) -> float | None: """Run one agent session. Returns reward on success, ``None`` on failure.""" + logger.debug( + "V2 rollout member start: task_id=%s group_id=%s member=%d " + "session_id=%s version=%s mode=%s", + task_id, + group_id, + member_index, + session_id, + version, + execution_mode, + ) try: rewards = await self.agent.run( data, @@ -194,10 +222,33 @@ async def _run_one(session_id: str, session_api_key: str) -> float | None: group_id, ) return None + finally: + logger.debug( + "V2 rollout member finish: task_id=%s group_id=%s member=%d " + "session_id=%s version=%s mode=%s", + task_id, + group_id, + member_index, + session_id, + version, + execution_mode, + ) - results = await asyncio.gather( - *[_run_one(sid, api_key) for sid, api_key in sessions] - ) + if self.serialize_group_samples: + results = [] + for member_index, (session_id, session_api_key) in enumerate(sessions): + results.append( + await _run_one(member_index, session_id, session_api_key) + ) + else: + results = await asyncio.gather( + *[ + _run_one(member_index, session_id, session_api_key) + for member_index, (session_id, session_api_key) in enumerate( + sessions + ) + ] + ) session_ids = [sid for sid, _ in sessions] diff --git a/areal/v2/inference_service/data_proxy/__main__.py b/areal/v2/inference_service/data_proxy/__main__.py index e0458bebf9..d706fc1ad0 100644 --- a/areal/v2/inference_service/data_proxy/__main__.py +++ b/areal/v2/inference_service/data_proxy/__main__.py @@ -57,6 +57,11 @@ def main(): "--callback-server-addr", default="", ) + parser.add_argument( + "--deterministic-sampling", + action="store_true", + help="Derive stable per-session request seeds.", + ) parser.add_argument( "--tool-call-parser", default="qwen", @@ -97,6 +102,7 @@ def main(): set_reward_finish_timeout=args.set_reward_finish_timeout, admin_api_key=args.admin_api_key, callback_server_addr=args.callback_server_addr, + deterministic_sampling=args.deterministic_sampling, serving_addr=format_hostport(serving_host, args.port), tool_call_parser=args.tool_call_parser, reasoning_parser=args.reasoning_parser, diff --git a/areal/v2/inference_service/data_proxy/app.py b/areal/v2/inference_service/data_proxy/app.py index 37dbf617bd..e7c3644f94 100644 --- a/areal/v2/inference_service/data_proxy/app.py +++ b/areal/v2/inference_service/data_proxy/app.py @@ -30,6 +30,7 @@ from areal.infra.utils.http import create_httpx_client from areal.utils import logging from areal.utils.data import concat_padded_tensors +from areal.utils.seeding import derive_deterministic_seed from areal.v2.inference_service.data_proxy.config import DataProxyConfig from areal.v2.inference_service.data_proxy.pause import PauseState from areal.v2.inference_service.data_proxy.session import ( @@ -472,7 +473,11 @@ async def start_session( for i in range(group_size): try: session_id, session_api_key = store.start_session( - body.task_id, body.api_key if i == 0 else None + body.task_id, + body.api_key if i == 0 else None, + sampling_seed_identity=( + f"{body.task_id}:{i}" if group_size > 1 else body.task_id + ), ) except ValueError as e: raise HTTPException(status_code=409, detail=str(e)) @@ -661,6 +666,20 @@ async def _stream_and_cache(): if "top_p" not in kwargs: kwargs["top_p"] = 1.0 + deterministic_sampling = ( + session is not None and app.state.config.deterministic_sampling + ) + request_index = ( + session.next_sampling_request_index() if deterministic_sampling else None + ) + if deterministic_sampling and kwargs.get("seed") is None: + assert session is not None + assert request_index is not None + kwargs["seed"] = derive_deterministic_seed( + session.sampling_seed_identity, + request_index, + ) + create_fn: Any = areal_client.chat.completions.create try: diff --git a/areal/v2/inference_service/data_proxy/config.py b/areal/v2/inference_service/data_proxy/config.py index aabc00f27a..55ef5c4a3e 100644 --- a/areal/v2/inference_service/data_proxy/config.py +++ b/areal/v2/inference_service/data_proxy/config.py @@ -17,6 +17,7 @@ class DataProxyConfig: resubmit_wait: float = 0.5 # seconds between is_paused polls admin_api_key: str = "areal-admin-key" # admin key for authentication callback_server_addr: str = "" + deterministic_sampling: bool = False # Resolved serving address (host:port) used as node_addr for RTensor shards. # Set at startup by __main__.py after the host is resolved. serving_addr: str = "" diff --git a/areal/v2/inference_service/data_proxy/session.py b/areal/v2/inference_service/data_proxy/session.py index e6cbe6c2d2..40e146a125 100644 --- a/areal/v2/inference_service/data_proxy/session.py +++ b/areal/v2/inference_service/data_proxy/session.py @@ -133,17 +133,27 @@ def __init__( self, session_id: str, set_reward_finish_timeout: float = 0.0, + sampling_seed_identity: str | None = None, ): self.session_id = session_id + self.sampling_seed_identity = sampling_seed_identity or session_id self._set_reward_finish_timeout = set_reward_finish_timeout self._last_access_time = time.time() self._lock = threading.Lock() self._active_completions = InteractionCache() self._ready_trajectories: OrderedDict[int, ReadyTrajectory] = OrderedDict() self._next_trajectory_id = 0 + self._next_sampling_request_index = 0 self._last_set_reward_time: float | None = None self._last_reward_interaction_id: str | None = None + def next_sampling_request_index(self) -> int: + """Reserve a request index before inference without serializing generation.""" + with self._lock: + request_index = self._next_sampling_request_index + self._next_sampling_request_index += 1 + return request_index + def update_last_access(self) -> None: with self._lock: self._last_access_time = time.time() @@ -388,7 +398,10 @@ def admin_api_key(self) -> str: return self._admin_api_key def start_session( - self, task_id: str, api_key: str | None = None + self, + task_id: str, + api_key: str | None = None, + sampling_seed_identity: str | None = None, ) -> tuple[str, str]: """Start a new session, returning (session_id, session_api_key). @@ -425,6 +438,7 @@ def start_session( self._sessions[session_id] = SessionData( session_id=session_id, set_reward_finish_timeout=self._set_reward_finish_timeout, + sampling_seed_identity=sampling_seed_identity, ) self._api_key_to_session[session_api_key] = session_id self._session_to_api_key[session_id] = session_api_key diff --git a/areal/v2/inference_service/sglang/bridge.py b/areal/v2/inference_service/sglang/bridge.py index 67b835f732..b8ef729099 100644 --- a/areal/v2/inference_service/sglang/bridge.py +++ b/areal/v2/inference_service/sglang/bridge.py @@ -59,6 +59,8 @@ def build_generation_request( } if gconfig.stop: sampling_params["stop"] = gconfig.stop + if gconfig.seed is not None: + sampling_params["sampling_seed"] = gconfig.seed payload: dict[str, Any] = { "input_ids": list(req.input_ids), diff --git a/areal/v2/inference_service/vllm/bridge.py b/areal/v2/inference_service/vllm/bridge.py index bc81c4fded..d2e5d93b90 100644 --- a/areal/v2/inference_service/vllm/bridge.py +++ b/areal/v2/inference_service/vllm/bridge.py @@ -53,6 +53,8 @@ def build_generation_request( "use_beam_search": gconfig.use_beam_search, "stream": False, } + if gconfig.seed is not None: + payload["seed"] = gconfig.seed if with_lora: lora_name = gconfig.lora_name diff --git a/docs/en/cli_reference.md b/docs/en/cli_reference.md index 7eee1e5781..28aa891a18 100644 --- a/docs/en/cli_reference.md +++ b/docs/en/cli_reference.md @@ -534,6 +534,7 @@ Controls text generation behavior for rollout. | `skip_special_tokens` | boolean | `True` | Skip special tokens when decoding/displaying outputs. | | `stop` | list of string \| None | `None` | One or multiple stop words. Generation will stop if one of these words is sampled. | | `frequency_penalty` | float | `0.0` | Penalizes tokens based on their frequency in generation so far. Must be between -2 and 2 where negative numbers encourage repetition. | +| `seed` | integer \| None | `None` | Per-request sampling seed sent to the inference backend. | | `lora_name` | string | `"default_lora"` | Lora name to be used for this generation. | | `use_beam_search` | boolean | `False` | Enable beam search in the vLLM engine. When enabled, sampling parameters like temperature, top-p, and top-k are auto ignored. | | `reward_normalization` | boolean | `False` | If True, apply per-prompt reward normalization across the n_samples rollouts of the same prompt inside GroupedRolloutWorkflow. Only affects InteractionWithTokenLogpReward workflows such as SWE agent workflows. Not supported by RolloutControllerV2 yet. | @@ -545,38 +546,40 @@ Controls text generation behavior for rollout. Configuration for inference servers, including offpolicyness control. -| Parameter | Type | Default | Description | -| ------------------------- | --------------------------------------------------- | ------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `experiment_name` | string \| None | `None` | - | -| `trial_name` | string \| None | `None` | - | -| `fileroot` | string \| None | `None` | Root directory for logs and trajectory dumps. | -| `max_concurrent_rollouts` | integer \| None | `None` | Maximum number of concurrent rollouts to the inference engine. Defaults to consumer_batch_size. | -| `queue_size` | integer \| None | `None` | Input/Output queue size for async rollout. | -| `consumer_batch_size` | integer | `1` | Batch size for consuming rollouts from the queue. | -| `max_head_offpolicyness` | integer | `0` | Maximum off-policyness for the head. If the current version is more than this many versions behind, the request will not be accepted. | -| `enable_rollout_tracing` | boolean | `False` | Whether to output verbose tracing messages for each generation request. | -| `check_trajectory_format` | boolean | `False` | Whether to check the format of produced trajectories of a customized workflow. Useful when debugging the workflow in isolation. Should be False during RL training. | -| `tokenizer_path` | string | `""` | Path to tokenizer for trajectory text decoding. | -| `dump_to_file` | boolean | `False` | Whether to dump the trajectories to files under fileroot. | -| `setup_timeout` | float | `300.0` | Timeout in seconds of connecting to remote servers or launching local servers. | -| `workers_ready_timeout` | float | `30.0` | Timeout (seconds) for initialize() to wait for guards to be ready. | -| `request_timeout` | float | `3600` | Timeout for HTTP requests. | -| `request_retries` | integer | `3` | Number of retries for failed requests. | -| `pause_grace_period` | float | `0.0` | The grace period after calling /pause_generation. Wait until all requests have been dropped. | -| `scheduling_spec` | `tuple` | **Required** | inference engine schedule specs. Can accept 1 or 2 SchedulingSpec: if 1 spec provided, it's used for both worker and engine, engine is embedded in the worker; if 2 specs provided, first one is for worker, second one is for engine. Currently only used by the RolloutController. | -| `backend` | string | **Required** | Backend and parallelism strategy. Must include an explicit backend prefix, e.g. 'sglang:d4', 'vllm:d2t4'. Required. | -| `scheduling_strategy` | [`SchedulingStrategy`](section-scheduling-strategy) | **Required** | The scheduling strategy of this InferenceEngine, either separation or colocation. Currently only used by the RolloutController. | -| `use_lora` | boolean | `False` | Whether to use LoRA. Should be same as actors LORA option. | -| `lora_name` | string | `""` | LoRA adapter name the rollout backend serves. Generation requests select the adapter by this name (plus the weight version). Usually left empty and auto-filled from gconfig.lora_name by PPOConfig.__post_init__ so load and request sides stay in sync. | -| `agent` | [`AgentConfig`](section-agent) | **Required** | Agent workflow configuration used by inference-service rollouts. | -| `return_routed_experts` | boolean | `False` | Return routed expert indices for MoE models. Effective only when using SGLang engine with MoE models. | -| `_version` | string | `"v1"` | Rollout controller implementation version. Use 'v1' for legacy RolloutController, 'v2' for RolloutControllerV2. **Choices:** `v1`, `v2` | -| `model` | string | `"default"` | Model name exposed through the inference-service gateway. | -| `routing_strategy` | string | `"round_robin"` | Routing strategy for the inference-service router. | -| `poll_interval` | float | `5.0` | Health-poll interval in seconds for the inference-service router. | -| `admin_api_key` | string | `"areal-admin-key"` | Admin API key used by the inference-service gateway, router, and data proxies. | -| `api_url` | string \| None | `None` | External OpenAI-compatible base URL for inference-service external model mode. | -| `provider_api_key` | string \| None | `None` | API key for the external OpenAI-compatible provider. | +| Parameter | Type | Default | Description | +| ------------------------- | --------------------------------------------------- | ------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `experiment_name` | string \| None | `None` | - | +| `trial_name` | string \| None | `None` | - | +| `fileroot` | string \| None | `None` | Root directory for logs and trajectory dumps. | +| `max_concurrent_rollouts` | integer \| None | `None` | Maximum number of concurrent rollouts to the inference engine. Defaults to consumer_batch_size. | +| `queue_size` | integer \| None | `None` | Input/Output queue size for async rollout. | +| `consumer_batch_size` | integer | `1` | Batch size for consuming rollouts from the queue. | +| `max_head_offpolicyness` | integer | `0` | Maximum off-policyness for the head. If the current version is more than this many versions behind, the request will not be accepted. | +| `enable_rollout_tracing` | boolean | `False` | Whether to output verbose tracing messages for each generation request. | +| `deterministic_sampling` | boolean | `False` | Use stable request seeds for internal OpenAI-proxy/data-proxy sessions, canonical group ordering, and task-ID ordering of completed rollout results. Concurrent SGLang generation also requires sglang.enable_deterministic_inference. End-to-end determinism is only supported with max_head_offpolicyness=0. | +| `serialize_group_samples` | boolean | `False` | Run RolloutControllerV2 samples within each group sequentially instead of concurrently. This provides stable within-group member submission order at the cost of rollout throughput; it does not serialize requests across groups. | +| `check_trajectory_format` | boolean | `False` | Whether to check the format of produced trajectories of a customized workflow. Useful when debugging the workflow in isolation. Should be False during RL training. | +| `tokenizer_path` | string | `""` | Path to tokenizer for trajectory text decoding. | +| `dump_to_file` | boolean | `False` | Whether to dump the trajectories to files under fileroot. | +| `setup_timeout` | float | `300.0` | Timeout in seconds of connecting to remote servers or launching local servers. | +| `workers_ready_timeout` | float | `30.0` | Timeout (seconds) for initialize() to wait for guards to be ready. | +| `request_timeout` | float | `3600` | Timeout for HTTP requests. | +| `request_retries` | integer | `3` | Number of retries for failed requests. | +| `pause_grace_period` | float | `0.0` | The grace period after calling /pause_generation. Wait until all requests have been dropped. | +| `scheduling_spec` | `tuple` | **Required** | inference engine schedule specs. Can accept 1 or 2 SchedulingSpec: if 1 spec provided, it's used for both worker and engine, engine is embedded in the worker; if 2 specs provided, first one is for worker, second one is for engine. Currently only used by the RolloutController. | +| `backend` | string | **Required** | Backend and parallelism strategy. Must include an explicit backend prefix, e.g. 'sglang:d4', 'vllm:d2t4'. Required. | +| `scheduling_strategy` | [`SchedulingStrategy`](section-scheduling-strategy) | **Required** | The scheduling strategy of this InferenceEngine, either separation or colocation. Currently only used by the RolloutController. | +| `use_lora` | boolean | `False` | Whether to use LoRA. Should be same as actors LORA option. | +| `lora_name` | string | `""` | LoRA adapter name the rollout backend serves. Generation requests select the adapter by this name (plus the weight version). Usually left empty and auto-filled from gconfig.lora_name by PPOConfig.__post_init__ so load and request sides stay in sync. | +| `agent` | [`AgentConfig`](section-agent) | **Required** | Agent workflow configuration used by inference-service rollouts. | +| `return_routed_experts` | boolean | `False` | Return routed expert indices for MoE models. Effective only when using SGLang engine with MoE models. | +| `_version` | string | `"v1"` | Rollout controller implementation version. Use 'v1' for legacy RolloutController, 'v2' for RolloutControllerV2. **Choices:** `v1`, `v2` | +| `model` | string | `"default"` | Model name exposed through the inference-service gateway. | +| `routing_strategy` | string | `"round_robin"` | Routing strategy for the inference-service router. | +| `poll_interval` | float | `5.0` | Health-poll interval in seconds for the inference-service router. | +| `admin_api_key` | string | `"areal-admin-key"` | Admin API key used by the inference-service gateway, router, and data proxies. | +| `api_url` | string \| None | `None` | External OpenAI-compatible base URL for inference-service external model mode. | +| `provider_api_key` | string \| None | `None` | API key for the external OpenAI-compatible provider. | (section-sg-lang)= @@ -615,6 +618,7 @@ https://github.com/sgl-project/sglang for detailed documentation. | `enable_memory_saver` | boolean | `False` | - | | `allow_auto_truncate` | boolean | `False` | - | | `attention_backend` | string \| None | `"fa3"` | - | +| `enable_deterministic_inference` | boolean | `False` | - | | `enable_multimodal` | boolean | `False` | - | | `sampling_backend` | string \| None | `None` | - | | `context_length` | integer \| None | `32768` | - | diff --git a/docs/zh/cli_reference.md b/docs/zh/cli_reference.md index 3a878d0837..75de17d84a 100644 --- a/docs/zh/cli_reference.md +++ b/docs/zh/cli_reference.md @@ -532,6 +532,7 @@ Controls text generation behavior for rollout. | `skip_special_tokens` | boolean | `True` | Skip special tokens when decoding/displaying outputs. | | `stop` | list of string \| None | `None` | One or multiple stop words. Generation will stop if one of these words is sampled. | | `frequency_penalty` | float | `0.0` | Penalizes tokens based on their frequency in generation so far. Must be between -2 and 2 where negative numbers encourage repetition. | +| `seed` | integer \| None | `None` | Per-request sampling seed sent to the inference backend. | | `lora_name` | string | `"default_lora"` | Lora name to be used for this generation. | | `use_beam_search` | boolean | `False` | Enable beam search in the vLLM engine. When enabled, sampling parameters like temperature, top-p, and top-k are auto ignored. | | `reward_normalization` | boolean | `False` | If True, apply per-prompt reward normalization across the n_samples rollouts of the same prompt inside GroupedRolloutWorkflow. Only affects InteractionWithTokenLogpReward workflows such as SWE agent workflows. Not supported by RolloutControllerV2 yet. | @@ -543,38 +544,40 @@ Controls text generation behavior for rollout. Configuration for inference servers, including offpolicyness control. -| Parameter | Type | Default | Description | -| ------------------------- | --------------------------------------------------- | ------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `experiment_name` | string \| None | `None` | - | -| `trial_name` | string \| None | `None` | - | -| `fileroot` | string \| None | `None` | Root directory for logs and trajectory dumps. | -| `max_concurrent_rollouts` | integer \| None | `None` | Maximum number of concurrent rollouts to the inference engine. Defaults to consumer_batch_size. | -| `queue_size` | integer \| None | `None` | Input/Output queue size for async rollout. | -| `consumer_batch_size` | integer | `1` | Batch size for consuming rollouts from the queue. | -| `max_head_offpolicyness` | integer | `0` | Maximum off-policyness for the head. If the current version is more than this many versions behind, the request will not be accepted. | -| `enable_rollout_tracing` | boolean | `False` | Whether to output verbose tracing messages for each generation request. | -| `check_trajectory_format` | boolean | `False` | Whether to check the format of produced trajectories of a customized workflow. Useful when debugging the workflow in isolation. Should be False during RL training. | -| `tokenizer_path` | string | `""` | Path to tokenizer for trajectory text decoding. | -| `dump_to_file` | boolean | `False` | Whether to dump the trajectories to files under fileroot. | -| `setup_timeout` | float | `300.0` | Timeout in seconds of connecting to remote servers or launching local servers. | -| `workers_ready_timeout` | float | `30.0` | Timeout (seconds) for initialize() to wait for guards to be ready. | -| `request_timeout` | float | `3600` | Timeout for HTTP requests. | -| `request_retries` | integer | `3` | Number of retries for failed requests. | -| `pause_grace_period` | float | `0.0` | The grace period after calling /pause_generation. Wait until all requests have been dropped. | -| `scheduling_spec` | `tuple` | **Required** | inference engine schedule specs. Can accept 1 or 2 SchedulingSpec: if 1 spec provided, it's used for both worker and engine, engine is embedded in the worker; if 2 specs provided, first one is for worker, second one is for engine. Currently only used by the RolloutController. | -| `backend` | string | **Required** | Backend and parallelism strategy. Must include an explicit backend prefix, e.g. 'sglang:d4', 'vllm:d2t4'. Required. | -| `scheduling_strategy` | [`SchedulingStrategy`](section-scheduling-strategy) | **Required** | The scheduling strategy of this InferenceEngine, either separation or colocation. Currently only used by the RolloutController. | -| `use_lora` | boolean | `False` | Whether to use LoRA. Should be same as actors LORA option. | -| `lora_name` | string | `""` | LoRA adapter name the rollout backend serves. Generation requests select the adapter by this name (plus the weight version). Usually left empty and auto-filled from gconfig.lora_name by PPOConfig.__post_init__ so load and request sides stay in sync. | -| `agent` | [`AgentConfig`](section-agent) | **Required** | Agent workflow configuration used by inference-service rollouts. | -| `return_routed_experts` | boolean | `False` | Return routed expert indices for MoE models. Effective only when using SGLang engine with MoE models. | -| `_version` | string | `"v1"` | Rollout controller implementation version. Use 'v1' for legacy RolloutController, 'v2' for RolloutControllerV2. **Choices:** `v1`, `v2` | -| `model` | string | `"default"` | Model name exposed through the inference-service gateway. | -| `routing_strategy` | string | `"round_robin"` | Routing strategy for the inference-service router. | -| `poll_interval` | float | `5.0` | Health-poll interval in seconds for the inference-service router. | -| `admin_api_key` | string | `"areal-admin-key"` | Admin API key used by the inference-service gateway, router, and data proxies. | -| `api_url` | string \| None | `None` | External OpenAI-compatible base URL for inference-service external model mode. | -| `provider_api_key` | string \| None | `None` | API key for the external OpenAI-compatible provider. | +| Parameter | Type | Default | Description | +| ------------------------- | --------------------------------------------------- | ------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `experiment_name` | string \| None | `None` | - | +| `trial_name` | string \| None | `None` | - | +| `fileroot` | string \| None | `None` | Root directory for logs and trajectory dumps. | +| `max_concurrent_rollouts` | integer \| None | `None` | Maximum number of concurrent rollouts to the inference engine. Defaults to consumer_batch_size. | +| `queue_size` | integer \| None | `None` | Input/Output queue size for async rollout. | +| `consumer_batch_size` | integer | `1` | Batch size for consuming rollouts from the queue. | +| `max_head_offpolicyness` | integer | `0` | Maximum off-policyness for the head. If the current version is more than this many versions behind, the request will not be accepted. | +| `enable_rollout_tracing` | boolean | `False` | Whether to output verbose tracing messages for each generation request. | +| `deterministic_sampling` | boolean | `False` | Use stable request seeds for internal OpenAI-proxy/data-proxy sessions, canonical group ordering, and task-ID ordering of completed rollout results. Concurrent SGLang generation also requires sglang.enable_deterministic_inference. End-to-end determinism is only supported with max_head_offpolicyness=0. | +| `serialize_group_samples` | boolean | `False` | Run RolloutControllerV2 samples within each group sequentially instead of concurrently. This provides stable within-group member submission order at the cost of rollout throughput; it does not serialize requests across groups. | +| `check_trajectory_format` | boolean | `False` | Whether to check the format of produced trajectories of a customized workflow. Useful when debugging the workflow in isolation. Should be False during RL training. | +| `tokenizer_path` | string | `""` | Path to tokenizer for trajectory text decoding. | +| `dump_to_file` | boolean | `False` | Whether to dump the trajectories to files under fileroot. | +| `setup_timeout` | float | `300.0` | Timeout in seconds of connecting to remote servers or launching local servers. | +| `workers_ready_timeout` | float | `30.0` | Timeout (seconds) for initialize() to wait for guards to be ready. | +| `request_timeout` | float | `3600` | Timeout for HTTP requests. | +| `request_retries` | integer | `3` | Number of retries for failed requests. | +| `pause_grace_period` | float | `0.0` | The grace period after calling /pause_generation. Wait until all requests have been dropped. | +| `scheduling_spec` | `tuple` | **Required** | inference engine schedule specs. Can accept 1 or 2 SchedulingSpec: if 1 spec provided, it's used for both worker and engine, engine is embedded in the worker; if 2 specs provided, first one is for worker, second one is for engine. Currently only used by the RolloutController. | +| `backend` | string | **Required** | Backend and parallelism strategy. Must include an explicit backend prefix, e.g. 'sglang:d4', 'vllm:d2t4'. Required. | +| `scheduling_strategy` | [`SchedulingStrategy`](section-scheduling-strategy) | **Required** | The scheduling strategy of this InferenceEngine, either separation or colocation. Currently only used by the RolloutController. | +| `use_lora` | boolean | `False` | Whether to use LoRA. Should be same as actors LORA option. | +| `lora_name` | string | `""` | LoRA adapter name the rollout backend serves. Generation requests select the adapter by this name (plus the weight version). Usually left empty and auto-filled from gconfig.lora_name by PPOConfig.__post_init__ so load and request sides stay in sync. | +| `agent` | [`AgentConfig`](section-agent) | **Required** | Agent workflow configuration used by inference-service rollouts. | +| `return_routed_experts` | boolean | `False` | Return routed expert indices for MoE models. Effective only when using SGLang engine with MoE models. | +| `_version` | string | `"v1"` | Rollout controller implementation version. Use 'v1' for legacy RolloutController, 'v2' for RolloutControllerV2. **Choices:** `v1`, `v2` | +| `model` | string | `"default"` | Model name exposed through the inference-service gateway. | +| `routing_strategy` | string | `"round_robin"` | Routing strategy for the inference-service router. | +| `poll_interval` | float | `5.0` | Health-poll interval in seconds for the inference-service router. | +| `admin_api_key` | string | `"areal-admin-key"` | Admin API key used by the inference-service gateway, router, and data proxies. | +| `api_url` | string \| None | `None` | External OpenAI-compatible base URL for inference-service external model mode. | +| `provider_api_key` | string \| None | `None` | API key for the external OpenAI-compatible provider. | (section-sg-lang)= @@ -613,6 +616,7 @@ https://github.com/sgl-project/sglang for detailed documentation. | `enable_memory_saver` | boolean | `False` | - | | `allow_auto_truncate` | boolean | `False` | - | | `attention_backend` | string \| None | `"fa3"` | - | +| `enable_deterministic_inference` | boolean | `False` | - | | `enable_multimodal` | boolean | `False` | - | | `sampling_backend` | string \| None | `None` | - | | `context_length` | integer \| None | `32768` | - | diff --git a/tests/test_deterministic_sampling.py b/tests/test_deterministic_sampling.py new file mode 100644 index 0000000000..11d9fc863a --- /dev/null +++ b/tests/test_deterministic_sampling.py @@ -0,0 +1,540 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +import threading +from collections import deque +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from types import SimpleNamespace +from unittest.mock import MagicMock, Mock + +import pytest +import torch + +from areal.api import ModelRequest, ModelResponse +from areal.api.cli_args import GenerationHyperparameters, SGLangConfig +from areal.engine.sglang_remote import SGLangBackend +from areal.experimental.openai.client import ArealOpenAI +from areal.experimental.openai.proxy import proxy_rollout_server +from areal.experimental.openai.proxy.proxy_rollout_server import ( + _deterministic_sampling_seed, +) +from areal.experimental.openai.proxy.server import SessionData +from areal.infra import workflow_context +from areal.infra import workflow_executor as workflow_executor_module +from areal.infra.remote_inf_engine import GroupedRolloutWorkflow +from areal.infra.workflow_executor import ( + BatchTaskDispatcher, + TaskIdGenerator, + WorkflowExecutor, + _select_results, +) +from areal.v2.inference_service.data_proxy.session import SessionData as V2SessionData +from areal.v2.inference_service.sglang.bridge import SGLangBridgeBackend + + +def test_sampling_seed_is_stable_across_calls(): + assert _deterministic_sampling_seed("17:3", 0) == _deterministic_sampling_seed( + "17:3", 0 + ) + + +def test_sampling_seed_differs_per_request_and_per_sample(): + assert _deterministic_sampling_seed("17:3", 0) != _deterministic_sampling_seed( + "17:3", 1 + ) + assert _deterministic_sampling_seed("17:3", 0) != _deterministic_sampling_seed( + "17:4", 0 + ) + + +def test_sampling_seed_identity_ignores_physical_session_suffix(): + sessions = [ + SessionData("17:3-0", sampling_seed_identity="17:3"), + SessionData("17:3-1", sampling_seed_identity="17:3"), + ] + + seeds = [ + _deterministic_sampling_seed( + session.sampling_seed_identity, + session.next_sampling_request_index(), + ) + for session in sessions + ] + + assert seeds[0] == seeds[1] + + +def test_sampling_request_indices_are_unique_under_concurrency(): + session = SessionData("17:3-0", sampling_seed_identity="17:3") + + with ThreadPoolExecutor(max_workers=8) as executor: + indices = list( + executor.map(lambda _: session.next_sampling_request_index(), range(32)) + ) + + assert sorted(indices) == list(range(32)) + + +def test_v2_sampling_request_indices_are_unique_under_concurrency(): + session = V2SessionData("17:3-0", sampling_seed_identity="17:3") + + with ThreadPoolExecutor(max_workers=8) as executor: + indices = list( + executor.map(lambda _: session.next_sampling_request_index(), range(32)) + ) + + assert sorted(indices) == list(range(32)) + + +@pytest.mark.asyncio +async def test_proxy_allocates_unique_seeds_before_concurrent_generation(monkeypatch): + session = SessionData("17:3-0", sampling_seed_identity="17:3") + monkeypatch.setattr(proxy_rollout_server, "_openai_client", object()) + monkeypatch.setattr(proxy_rollout_server, "_deterministic_sampling", True) + monkeypatch.setitem( + proxy_rollout_server._session_cache, session.session_id, session + ) + + async def create_fn(*, areal_cache, seed, temperature, top_p): + await asyncio.sleep(0) + return seed + + seeds = await asyncio.gather( + *[ + proxy_rollout_server._call_client_create( + create_fn, + {"temperature": 1.0, "top_p": 1.0}, + session.session_id, + ) + for _ in range(8) + ] + ) + + expected = { + _deterministic_sampling_seed(session.sampling_seed_identity, i) + for i in range(8) + } + assert set(seeds) == expected + + +@pytest.mark.asyncio +async def test_proxy_explicit_seed_still_consumes_request_index(monkeypatch): + session = SessionData("17:3-0", sampling_seed_identity="17:3") + monkeypatch.setattr(proxy_rollout_server, "_openai_client", object()) + monkeypatch.setattr(proxy_rollout_server, "_deterministic_sampling", True) + monkeypatch.setitem( + proxy_rollout_server._session_cache, session.session_id, session + ) + + async def create_fn(*, areal_cache, seed, temperature, top_p): + return seed + + explicit_seed = await proxy_rollout_server._call_client_create( + create_fn, + {"seed": 123, "temperature": 1.0, "top_p": 1.0}, + session.session_id, + ) + derived_seed = await proxy_rollout_server._call_client_create( + create_fn, + {"temperature": 1.0, "top_p": 1.0}, + session.session_id, + ) + + assert explicit_seed == 123 + assert derived_seed == _deterministic_sampling_seed( + session.sampling_seed_identity, 1 + ) + + +@pytest.mark.asyncio +async def test_grouped_rollout_is_concurrent_and_sample_ordered(): + class _Workflow: + active = 0 + max_active = 0 + + async def arun_episode(self, engine, data): + sample_idx = workflow_context.get().sample_idx + self.active += 1 + self.max_active = max(self.max_active, self.active) + await asyncio.sleep(0.01 * (3 - sample_idx)) + self.active -= 1 + return {"sample_idx": torch.tensor([[sample_idx]])} + + workflow = _Workflow() + grouped = GroupedRolloutWorkflow( + workflow=workflow, + group_size=3, + logger=Mock(), + ) + engine = SimpleNamespace(config=SimpleNamespace(deterministic_sampling=True)) + + result = await grouped.arun_episode(engine, {}) + + assert workflow.max_active == 3 + assert result is not None + assert result["sample_idx"].tolist() == [[0], [1], [2]] + + +def test_sglang_request_forwards_sampling_seed_when_set(): + req = ModelRequest( + input_ids=[1, 2, 3], + gconfig=GenerationHyperparameters(seed=12345), + ) + + request = SGLangBackend().build_generation_request(req, with_lora=False, version=0) + + assert request.payload["sampling_params"]["sampling_seed"] == 12345 + + +def test_sglang_request_omits_sampling_seed_by_default(): + req = ModelRequest(input_ids=[1, 2, 3], gconfig=GenerationHyperparameters()) + + request = SGLangBackend().build_generation_request(req, with_lora=False, version=0) + + assert "sampling_seed" not in request.payload["sampling_params"] + + +def test_sglang_v2_request_forwards_sampling_seed_when_set(): + req = ModelRequest( + input_ids=[1, 2, 3], + gconfig=GenerationHyperparameters(seed=12345), + ) + + request = SGLangBridgeBackend().build_generation_request( + req, with_lora=False, version=0 + ) + + assert request.payload["sampling_params"]["sampling_seed"] == 12345 + + +def test_sglang_v2_request_omits_sampling_seed_by_default(): + req = ModelRequest(input_ids=[1, 2, 3], gconfig=GenerationHyperparameters()) + + request = SGLangBridgeBackend().build_generation_request( + req, with_lora=False, version=0 + ) + + assert "sampling_seed" not in request.payload["sampling_params"] + + +@pytest.mark.asyncio +async def test_areal_openai_forwards_seed_into_model_request(monkeypatch): + monkeypatch.setattr( + "areal.utils.hf_utils.pkg_version.is_version_greater_or_equal", + lambda *_: False, + ) + tokenizer = MagicMock() + tokenizer.apply_chat_template.return_value = [10, 11] + tokenizer.decode.return_value = "ok" + tokenizer.eos_token_id = 2 + tokenizer.pad_token_id = 0 + + class CapturingEngine: + async def agenerate(self, req): + self.request = req + return ModelResponse( + input_tokens=req.input_ids, + output_tokens=[3], + output_logprobs=[-0.1], + output_versions=[0], + stop_reason="length", + tokenizer=tokenizer, + ) + + engine = CapturingEngine() + client = ArealOpenAI(engine=engine, tokenizer=tokenizer, api_key="test") + try: + await client.chat.completions.create( + messages=[{"role": "user", "content": "hi"}], + max_completion_tokens=4, + seed=12345, + ) + finally: + await client.close() + + assert engine.request.gconfig.seed == 12345 + + +def test_sglang_server_args_enable_deterministic_inference(monkeypatch): + monkeypatch.setattr( + "areal.api.cli_args.pkg_version.is_version_greater_or_equal", + lambda *_: True, + ) + args = SGLangConfig.build_args( + SGLangConfig( + model_path="test-model", + enable_deterministic_inference=True, + ), + tp_size=1, + base_gpu_id=0, + ) + + assert args["enable_deterministic_inference"] is True + + +@pytest.mark.parametrize("attention_backend", ["flashinfer", "fa3", "triton", None]) +def test_sglang_deterministic_inference_supported_backend_does_not_warn( + monkeypatch, attention_backend +): + monkeypatch.setattr( + "areal.api.cli_args.pkg_version.is_version_greater_or_equal", + lambda *_: True, + ) + mock_logger = Mock() + monkeypatch.setattr("areal.api.cli_args.logger", mock_logger) + + SGLangConfig.build_args( + SGLangConfig( + model_path="test-model", + attention_backend=attention_backend, + enable_deterministic_inference=True, + ), + tp_size=1, + base_gpu_id=0, + ) + + mock_logger.warning.assert_not_called() + + +def test_sglang_deterministic_inference_unsupported_backend_warns(monkeypatch): + monkeypatch.setattr( + "areal.api.cli_args.pkg_version.is_version_greater_or_equal", + lambda *_: True, + ) + mock_logger = Mock() + monkeypatch.setattr("areal.api.cli_args.logger", mock_logger) + + SGLangConfig.build_args( + SGLangConfig( + model_path="test-model", + attention_backend="torch_native", + enable_deterministic_inference=True, + ), + tp_size=1, + base_gpu_id=0, + ) + + mock_logger.warning.assert_called_once() + assert "torch_native" in mock_logger.warning.call_args.args + + +@dataclass +class _FakeTimedResult: + task_id: int + create_time: float + data: object | None = None + + +def test_select_results_is_task_ordered_without_shuffle_when_deterministic( + monkeypatch, +): + mock_shuffle = Mock() + monkeypatch.setattr(workflow_executor_module.random, "shuffle", mock_shuffle) + # Arrival order (create_time) deliberately disagrees with task id order. + drained = [ + _FakeTimedResult(task_id=2, create_time=1.0), + _FakeTimedResult(task_id=0, create_time=2.0), + _FakeTimedResult(task_id=1, create_time=3.0), + ] + + selected, pending = _select_results(drained, count=2, deterministic=True) + + assert [r.task_id for r in selected] == [0, 1] + assert [r.task_id for r in pending] == [2] + mock_shuffle.assert_not_called() + + +def test_select_results_is_arrival_ordered_and_shuffled_by_default(monkeypatch): + mock_shuffle = Mock() + monkeypatch.setattr(workflow_executor_module.random, "shuffle", mock_shuffle) + drained = [ + _FakeTimedResult(task_id=2, create_time=1.0), + _FakeTimedResult(task_id=0, create_time=2.0), + _FakeTimedResult(task_id=1, create_time=3.0), + ] + + selected, pending = _select_results(drained, count=2, deterministic=False) + + # Oldest-first selection is preserved; the selected order itself is + # shuffled, so only membership is asserted here. + assert {r.task_id for r in selected} == {2, 0} + assert [r.task_id for r in pending] == [1] + mock_shuffle.assert_called_once_with(selected) + + +def test_wait_results_selects_completed_tasks_when_deterministic(): + dispatcher = object.__new__(BatchTaskDispatcher) + dispatcher.deterministic_order = True + dispatcher._result_cv = threading.Condition() + dispatcher._active_task_ids = {0: None, 1: None, 2: None} + dispatcher._shutdown_event = threading.Event() + dispatcher._pending_results = { + 1: _FakeTimedResult(task_id=1, create_time=1.0, data="one"), + 2: _FakeTimedResult(task_id=2, create_time=2.0, data="two"), + } + dispatcher._check_thread_exception = lambda: None + + results = dispatcher.wait_results(2, timeout=0) + + assert results == ["one", "two"] + assert dispatcher._pending_results == {} + assert dispatcher._active_task_ids == {0: None} + + +def test_wait_results_fails_fast_when_dispatcher_is_shutting_down(): + dispatcher = object.__new__(BatchTaskDispatcher) + dispatcher.deterministic_order = True + dispatcher._result_cv = threading.Condition() + dispatcher._active_task_ids = {0: None} + dispatcher._pending_results = {} + dispatcher._shutdown_event = threading.Event() + dispatcher._shutdown_event.set() + dispatcher._check_thread_exception = lambda: None + + with pytest.raises(RuntimeError, match="shutting down"): + dispatcher.wait_results(1, timeout=1) + + +def test_callback_registration_rejects_duplicate_task_id(): + dispatcher = object.__new__(BatchTaskDispatcher) + dispatcher._result_cv = threading.Condition() + dispatcher._task_callbacks = {} + + dispatcher.register_callback(7, "http://first") + + with pytest.raises(ValueError, match="already registered"): + dispatcher.register_callback(7, "http://second") + assert dispatcher._task_callbacks == {7: "http://first"} + + +def test_submit_task_input_rolls_back_when_enqueue_hook_fails(): + dispatcher = object.__new__(BatchTaskDispatcher) + dispatcher._check_thread_exception = lambda: None + dispatcher._result_cv = threading.Condition() + dispatcher._input_cv = threading.Condition() + dispatcher._active_task_ids = {} + dispatcher._pending_inputs = deque() + dispatcher.staleness_manager = SimpleNamespace( + on_rollout_enqueued=Mock(side_effect=RuntimeError("hook failed")) + ) + dispatcher.enable_tracing = False + task_input = SimpleNamespace(task_id=7) + + with pytest.raises(RuntimeError, match="hook failed"): + dispatcher.submit_task_input(task_input) + + assert dispatcher._active_task_ids == {} + assert list(dispatcher._pending_inputs) == [] + + +def test_wait_for_task_removes_result_before_deterministic_batch_selection(): + dispatcher = object.__new__(BatchTaskDispatcher) + dispatcher.deterministic_order = True + dispatcher._result_cv = threading.Condition() + dispatcher._active_task_ids = {0: None, 1: None, 2: None} + dispatcher._shutdown_event = threading.Event() + dispatcher._pending_results = { + task_id: _FakeTimedResult(task_id, float(task_id), str(task_id)) + for task_id in range(3) + } + dispatcher._check_thread_exception = lambda: None + + assert dispatcher.wait_for_task(1, timeout=0) == "1" + assert dispatcher.wait_results(2, timeout=0) == ["0", "2"] + + +def test_workflow_executor_binds_callback_to_allocated_task_id(monkeypatch): + executor = object.__new__(WorkflowExecutor) + executor._task_id_generator = TaskIdGenerator() + executor._dispatcher = Mock() + monkeypatch.setattr( + workflow_executor_module.perf_tracer, "register_task", lambda _: None + ) + + task_id = executor.submit({}, workflow=Mock(), callback_addr="http://callback") + + assert task_id == 0 + executor.dispatcher.register_callback.assert_called_once_with(0, "http://callback") + submitted = executor.dispatcher.submit_task_input.call_args.args[0] + assert submitted.task_id == 0 + + +def test_workflow_executor_cancels_own_callback_when_submit_fails(monkeypatch): + executor = object.__new__(WorkflowExecutor) + executor._task_id_generator = TaskIdGenerator() + executor._dispatcher = Mock() + executor.dispatcher.submit_task_input.side_effect = ValueError("duplicate") + monkeypatch.setattr( + workflow_executor_module.perf_tracer, "register_task", lambda _: None + ) + + with pytest.raises(ValueError, match="duplicate"): + executor.submit({}, workflow=Mock(), callback_addr="http://callback") + + executor.dispatcher.cancel_callback.assert_called_once_with(0, "http://callback") + + +def test_dynamic_batch_counts_rejections_as_attempts(): + dispatcher = object.__new__(BatchTaskDispatcher) + dispatcher._input_cv = threading.Condition() + dispatcher._pending_inputs = [] + dispatcher.staleness_manager = SimpleNamespace(get_pending_limit=lambda: 0) + dispatcher.runner = SimpleNamespace( + max_queue_size=4, get_input_queue_size=lambda: 4 + ) + dispatcher.enable_tracing = False + dispatcher.wait_results = Mock(side_effect=[[None], ["accepted"]]) + + results = dispatcher.active_submit_and_wait(iter(()), batch_size=2, dynamic_bs=True) + + assert results == ["accepted"] + assert [ + call.kwargs["count"] for call in dispatcher.wait_results.call_args_list + ] == [ + 2, + 1, + ] + + +def test_fixed_batch_replaces_rejected_attempts_in_order(): + dispatcher = object.__new__(BatchTaskDispatcher) + dispatcher._input_cv = threading.Condition() + dispatcher._pending_inputs = [] + dispatcher.staleness_manager = SimpleNamespace(get_pending_limit=lambda: 0) + dispatcher.runner = SimpleNamespace( + max_queue_size=4, get_input_queue_size=lambda: 4 + ) + dispatcher.enable_tracing = False + dispatcher.wait_results = Mock(side_effect=[[None, "first"], ["replacement"]]) + + results = dispatcher.active_submit_and_wait(iter(()), batch_size=2) + + assert results == ["first", "replacement"] + assert [ + call.kwargs["count"] for call in dispatcher.wait_results.call_args_list + ] == [ + 2, + 1, + ] + + +def test_task_id_generator_advances_past_explicit_id(): + generator = TaskIdGenerator() + + generator.reserve_at_least(7) + + assert generator.next() == 8 + + +def test_responses_and_completions_both_accept_seed(): + import inspect + + from areal.experimental.openai.client import ( + AsyncCompletionsWithReward, + AsyncResponsesWithReward, + ) + + for cls in (AsyncCompletionsWithReward, AsyncResponsesWithReward): + params = inspect.signature(cls.create).parameters + assert "seed" in params, f"{cls.__name__}.create is missing a seed parameter" diff --git a/tests/test_vllm_generation_request.py b/tests/test_vllm_generation_request.py index 18a42d0ee8..616789c0a7 100644 --- a/tests/test_vllm_generation_request.py +++ b/tests/test_vllm_generation_request.py @@ -18,3 +18,31 @@ def test_vllm_forwards_frequency_penalty_and_stop(): assert payload["frequency_penalty"] == 0.5 assert payload["stop"] == ["STOP"] + + +def test_vllm_forwards_explicit_seed(): + """The V1 vLLM backend forwards an explicitly configured sampling seed.""" + req = ModelRequest( + input_ids=[11, 12], + gconfig=GenerationHyperparameters(max_new_tokens=8, seed=12345), + ) + + payload = ( + VLLMBackend().build_generation_request(req, with_lora=False, version=0).payload + ) + + assert payload["seed"] == 12345 + + +def test_vllm_omits_seed_when_unset(): + """The V1 vLLM backend leaves seed selection to vLLM when it is unset.""" + req = ModelRequest( + input_ids=[11, 12], + gconfig=GenerationHyperparameters(max_new_tokens=8), + ) + + payload = ( + VLLMBackend().build_generation_request(req, with_lora=False, version=0).payload + ) + + assert "seed" not in payload diff --git a/tests/v2/inference_service/test_controller.py b/tests/v2/inference_service/test_controller.py index 0e3eb24849..ee63a15c29 100644 --- a/tests/v2/inference_service/test_controller.py +++ b/tests/v2/inference_service/test_controller.py @@ -3,7 +3,7 @@ from __future__ import annotations import asyncio -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, call, patch import httpx import pytest @@ -76,6 +76,37 @@ def test_dump_to_file_defaults_to_false(self): cfg = InferenceEngineConfig(backend="sglang:d1") assert cfg.dump_to_file is False + def test_deterministic_sampling_with_offpolicy_head_warns(self): + with patch("areal.api.cli_args.logger") as mock_logger: + InferenceEngineConfig( + backend="sglang:d1", + deterministic_sampling=True, + max_head_offpolicyness=1, + ) + + mock_logger.warning.assert_called_once() + assert "task-to-weight-version" in mock_logger.warning.call_args.args[0] + + def test_deterministic_sampling_onpolicy_does_not_warn(self): + with patch("areal.api.cli_args.logger") as mock_logger: + InferenceEngineConfig( + backend="sglang:d1", + deterministic_sampling=True, + max_head_offpolicyness=0, + ) + + mock_logger.warning.assert_not_called() + + def test_nondeterministic_sampling_with_offpolicy_head_does_not_warn(self): + with patch("areal.api.cli_args.logger") as mock_logger: + InferenceEngineConfig( + backend="sglang:d1", + deterministic_sampling=False, + max_head_offpolicyness=1, + ) + + mock_logger.warning.assert_not_called() + # ============================================================================= # RolloutControllerV2 — workflow resolution helpers @@ -141,7 +172,11 @@ def test_resolve_should_accept_fn_callable(self): def test_resolve_workflow_with_agent_class(self): """Test _resolve_workflow wraps agent-like classes in InferenceServiceWorkflow.""" - cfg = InferenceEngineConfig(backend="sglang:d1", admin_api_key="test-key") + cfg = InferenceEngineConfig( + backend="sglang:d1", + admin_api_key="test-key", + serialize_group_samples=True, + ) scheduler = MagicMock(n_gpus_per_node=8) controller = RolloutControllerV2(config=cfg, scheduler=scheduler) controller._gateway_addr = "http://test:8080" @@ -157,6 +192,7 @@ async def run(self, data, **kwargs): assert isinstance(resolved, InferenceServiceWorkflow) assert resolved.agent is not None assert hasattr(resolved, "arun_episode") + assert resolved.serialize_group_samples is True def test_resolve_workflow_agent_class_without_gateway_raises(self): controller = RolloutControllerV2( @@ -357,9 +393,10 @@ def test_config_perf_tracer_is_noop(self): controller.config_perf_tracer() controller.save_perf_tracer() + @pytest.mark.parametrize("deterministic_sampling", [False, True]) @pytest.mark.asyncio - async def test_async_initialize_passes_callback_and_reward_timeout_to_data_proxy( - self, + async def test_async_initialize_passes_config_to_data_proxy( + self, deterministic_sampling ): from areal.api.cli_args import SchedulingSpec from areal.api.io_struct import LocalInfServerInfo @@ -375,6 +412,7 @@ async def test_async_initialize_passes_callback_and_reward_timeout_to_data_proxy backend="sglang:d1", tokenizer_path="mock-tokenizer", request_timeout=15.0, + deterministic_sampling=deterministic_sampling, agent=AgentConfig( agent_cls_path="tests.experimental.openai.utils.SimpleAgent", set_reward_finish_timeout=7.5, @@ -418,6 +456,7 @@ async def test_async_initialize_passes_callback_and_reward_timeout_to_data_proxy assert "7.5" in data_proxy_cmd assert "--callback-server-addr" in data_proxy_cmd assert "http://127.0.0.1:19000" in data_proxy_cmd + assert ("--deterministic-sampling" in data_proxy_cmd) is deterministic_sampling class TestOnlineCallbackFlow: @@ -533,6 +572,80 @@ async def test_cancelled_waiter_buffers_completed_online_result(self): class TestInferenceServiceWorkflow: + async def _run_offline_group( + self, + *, + serialize_group_samples: bool, + failing_member: int | None = None, + ): + active = 0 + max_active = 0 + start_order: list[int] = [] + all_started = asyncio.Event() + + class MockAgent: + async def run(self, data, **kwargs): + del data + nonlocal active, max_active + member_index = int(kwargs["api_key"].rsplit("-", 1)[1]) + start_order.append(member_index) + active += 1 + max_active = max(max_active, active) + try: + if serialize_group_samples: + await asyncio.sleep(0) + else: + if active == 4: + all_started.set() + await asyncio.wait_for(all_started.wait(), timeout=1.0) + if member_index == failing_member: + raise RuntimeError(f"member {member_index} failed") + return float(member_index) + finally: + active -= 1 + + controller = MagicMock() + controller.get_version.return_value = 3 + workflow = InferenceServiceWorkflow( + controller=controller, + agent=MockAgent(), + gateway_addr="http://test:8080", + admin_api_key="test-key", + group_size=4, + serialize_group_samples=serialize_group_samples, + ) + sessions = [(f"task-42-{i}", f"session-key-{i}") for i in range(4)] + workflow._start_session = AsyncMock(return_value=("grp-test-42", sessions)) + workflow._set_last_reward = AsyncMock(return_value=None) + workflow._export_interactions = AsyncMock( + return_value={"chatcmpl-1": MagicMock(reward=1.0)} + ) + + tracker = MagicMock() + with ( + patch( + "areal.v2.inference_service.controller.workflow.workflow_context" + ) as mock_wf_ctx, + patch( + "areal.v2.inference_service.controller.workflow.stats_tracker" + ) as mock_st, + ): + mock_http_session = AsyncMock() + mock_wf_ctx.get_aiohttp_session = AsyncMock(return_value=mock_http_session) + mock_wf_ctx.get.return_value = MagicMock(task_id=42) + mock_wf_ctx.get_httpx_client = AsyncMock(return_value=MagicMock()) + mock_wf_ctx.stat_scope.return_value = "rollout" + mock_st.get.return_value = tracker + + result = await workflow.arun_episode(engine=MagicMock(), data={}) + + workflow._export_interactions.assert_awaited_once_with( + mock_http_session, + [session_id for session_id, _ in sessions], + group_id="grp-test-42", + ) + return result, max_active, start_order, tracker, workflow + @pytest.mark.skip(reason="pending /export_trajectories traj schema migration") @pytest.mark.asyncio async def test_online_mode_waits_on_controller(self): @@ -638,6 +751,57 @@ async def run(self, data, **kwargs): mock_http_session, ["sess-1"], group_id="grp-test-1" ) + @pytest.mark.asyncio + async def test_offline_group_is_concurrent_by_default(self): + result, max_active, start_order, tracker, _ = await self._run_offline_group( + serialize_group_samples=False, + ) + + assert result is not None + assert max_active == 4 + assert sorted(start_order) == [0, 1, 2, 3] + assert tracker.scalar.call_args_list == [ + call(reward=0.0), + call(reward=1.0), + call(reward=2.0), + call(reward=3.0), + ] + + @pytest.mark.asyncio + async def test_offline_group_serial_flag_preserves_within_group_order(self): + result, max_active, start_order, tracker, _ = await self._run_offline_group( + serialize_group_samples=True, + ) + + assert result is not None + assert max_active == 1 + assert start_order == [0, 1, 2, 3] + assert tracker.scalar.call_args_list == [ + call(reward=0.0), + call(reward=1.0), + call(reward=2.0), + call(reward=3.0), + ] + + @pytest.mark.asyncio + async def test_offline_group_serial_flag_exports_after_failure(self): + ( + result, + max_active, + start_order, + tracker, + workflow, + ) = await self._run_offline_group( + serialize_group_samples=True, + failing_member=1, + ) + + assert result is None + assert max_active == 1 + assert start_order == [0, 1, 2, 3] + assert workflow._set_last_reward.await_count == 4 + assert tracker.scalar.call_count == 0 + # ============================================================================= # Multi-node inference configuration diff --git a/tests/v2/inference_service/test_data_proxy_chat.py b/tests/v2/inference_service/test_data_proxy_chat.py index 8c241374fb..d179482c88 100644 --- a/tests/v2/inference_service/test_data_proxy_chat.py +++ b/tests/v2/inference_service/test_data_proxy_chat.py @@ -9,6 +9,7 @@ import pytest import pytest_asyncio +from areal.utils.seeding import derive_deterministic_seed from areal.v2.inference_service.data_proxy.app import ( _flush_ready_trajectories, create_app, @@ -438,6 +439,81 @@ async def test_chat_completions_passes_sampling_params(client, mock_areal_client assert kw["temperature"] == 0.5 assert kw["top_p"] == 0.9 assert kw["max_tokens"] == 100 + assert "seed" not in kw + + +@pytest.mark.asyncio +async def test_chat_completions_deterministic_seed_distinguishes_group_sessions( + client, config, mock_areal_client +): + config.deterministic_sampling = True + resp = await client.post( + "/rl/start_session", + json={"task_id": "seed-test", "group_size": 2}, + headers=admin_headers(), + ) + sessions = resp.json()["sessions"] + + for session in sessions: + resp = await client.post( + "/chat/completions", + json={ + "model": "sglang", + "messages": [{"role": "user", "content": "hi"}], + }, + headers=session_headers(session["session_api_key"]), + ) + assert resp.status_code == 200 + + seeds = [ + call.kwargs["seed"] + for call in mock_areal_client.chat.completions.create.call_args_list + ] + assert seeds == [ + derive_deterministic_seed("seed-test:0", 0), + derive_deterministic_seed("seed-test:1", 0), + ] + assert seeds[0] != seeds[1] + + +@pytest.mark.asyncio +async def test_chat_completions_preserves_explicit_seed( + client, config, mock_areal_client +): + config.deterministic_sampling = True + resp = await client.post( + "/rl/start_session", + json={"task_id": "explicit-seed"}, + headers=admin_headers(), + ) + api_key = resp.json()["sessions"][0]["session_api_key"] + + resp = await client.post( + "/chat/completions", + json={ + "model": "sglang", + "messages": [{"role": "user", "content": "hi"}], + "seed": 12345, + }, + headers=session_headers(api_key), + ) + + assert resp.status_code == 200 + assert mock_areal_client.chat.completions.create.call_args.kwargs["seed"] == 12345 + + resp = await client.post( + "/chat/completions", + json={ + "model": "sglang", + "messages": [{"role": "user", "content": "again"}], + }, + headers=session_headers(api_key), + ) + + assert resp.status_code == 200 + assert mock_areal_client.chat.completions.create.call_args.kwargs[ + "seed" + ] == derive_deterministic_seed("explicit-seed", 1) # ============================================================================= diff --git a/tests/v2/inference_service/test_image_input.py b/tests/v2/inference_service/test_image_input.py index ce0821fc0a..eda87f9935 100644 --- a/tests/v2/inference_service/test_image_input.py +++ b/tests/v2/inference_service/test_image_input.py @@ -357,6 +357,7 @@ def test_vision_messages_get_real_data_uris(self, red_pixel_b64): n_samples=1, max_new_tokens=10, max_tokens=32768, + seed=12345, ) vision_msgs = [ [ @@ -383,6 +384,7 @@ def test_vision_messages_get_real_data_uris(self, red_pixel_b64): http_req = backend.build_generation_request(req, with_lora=False, version=0) assert http_req.endpoint == "/v1/chat/completions" + assert http_req.payload["seed"] == 12345 msg_content = http_req.payload["messages"][0]["content"] image_part = msg_content[1] assert image_part["type"] == "image_url" diff --git a/tests/v2/inference_service/test_inf_bridge.py b/tests/v2/inference_service/test_inf_bridge.py index aad70d546b..4b8b57f2a2 100644 --- a/tests/v2/inference_service/test_inf_bridge.py +++ b/tests/v2/inference_service/test_inf_bridge.py @@ -67,6 +67,7 @@ def _make_request( n_samples: int = 1, greedy: bool = False, temperature: float = 1.0, + seed: int | None = None, metadata: dict[str, Any] | None = None, lora_name: str | None = None, ) -> ModelRequest: @@ -79,6 +80,7 @@ def _make_request( max_tokens=max_tokens, greedy=greedy, temperature=temperature, + seed=seed, ) if lora_name is not None: gconfig.lora_name = lora_name @@ -473,6 +475,19 @@ def test_vllm_build_generation_request_for_text(self): assert http_req.payload["max_tokens"] == 7 assert http_req.payload["stream"] is False + @pytest.mark.parametrize("seed", [None, 12345]) + def test_vllm_build_generation_request_forwards_seed_when_set(self, seed): + """vLLM bridge forwards explicit seeds and omits unspecified seeds.""" + backend = VLLMBridgeBackend() + req = _make_request(input_ids=[11, 12], max_new_tokens=7, seed=seed) + + http_req = backend.build_generation_request(req, with_lora=False, version=0) + + if seed is None: + assert "seed" not in http_req.payload + else: + assert http_req.payload["seed"] == seed + def test_vllm_parse_generation_response_for_chat_format(self): """vLLM bridge parses chat logprobs content format.""" backend = VLLMBridgeBackend() diff --git a/tests/v2/inference_service/test_ipv6_entrypoints.py b/tests/v2/inference_service/test_ipv6_entrypoints.py index 2b7d9b37dd..180d7714ed 100644 --- a/tests/v2/inference_service/test_ipv6_entrypoints.py +++ b/tests/v2/inference_service/test_ipv6_entrypoints.py @@ -20,6 +20,7 @@ def test_data_proxy_main_formats_ipv6_serving_addr(): set_reward_finish_timeout=0.0, admin_api_key="admin-key", callback_server_addr="http://[::1]:19000", + deterministic_sampling=False, tool_call_parser="qwen", reasoning_parser="qwen3", engine_max_tokens=None, @@ -41,6 +42,7 @@ def test_data_proxy_main_formats_ipv6_serving_addr(): config = mock_create_app.call_args.args[0] assert config.serving_addr == "[::1]:8082" + assert config.deterministic_sampling is False mock_run.assert_called_once()