diff --git a/docs/docs.json b/docs/docs.json index 606a095e99..a2ad1f9f56 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -169,6 +169,7 @@ "user-guide/monitoring", "user-guide/customization", "user-guide/rollout-endpoints", + "user-guide/verifiers", "user-guide/fully-async", "user-guide/agentic-chat-template", "user-guide/cli-reference" diff --git a/docs/user-guide/cli-reference.md b/docs/user-guide/cli-reference.md index c6b13c36ee..65ef1d7af0 100644 --- a/docs/user-guide/cli-reference.md +++ b/docs/user-guide/cli-reference.md @@ -191,6 +191,15 @@ Sections mirror the launch-script argument groups. | `--rollout-stop` | str+ | – | Stop strings. | | `--rollout-stop-token-ids` | int+ | – | Stop token IDs. | +### Rollout: Verifiers + +| Flag | Type | Default | Notes | +|---|---|---|---| +| `--verifiers-config` | path | – | Use the built-in Verifiers rollout with this EnvConfig TOML file. | + +See [Verifiers](/user-guide/verifiers) for installation, configuration, and supported +environment behavior. + ### Eval | Flag | Type | Default | Notes | diff --git a/docs/user-guide/index.md b/docs/user-guide/index.md index 033b0c306e..9c7a6b5be0 100644 --- a/docs/user-guide/index.md +++ b/docs/user-guide/index.md @@ -11,6 +11,7 @@ description: Concepts, launch script walkthrough, customization hooks, and a com | [Monitoring & Logging](/user-guide/monitoring) | wandb, structured logs, per-source breakdowns, profiling, router metrics. | | [Customization](/user-guide/customization) | The 21 `--*-path` plug-points for custom Python — rollout, reward, filters, loss, hooks. | | [Rollout Endpoints](/user-guide/rollout-endpoints) | The `/generate` endpoint and the OpenAI chat endpoint for agentic sessions. | +| [Verifiers](/user-guide/verifiers) | Train on Verifiers tasksets and harnesses. | | [Fully Async Rollout](/user-guide/fully-async) | Queue-backed rollout production, tuning knobs, and when to use `train_async.py`. | | [Agentic Chat Templates](/user-guide/agentic-chat-template) | Turning on and verifying TITO so multi-turn agentic rollout stays append-only. | | [CLI Reference](/user-guide/cli-reference) | Every flag Miles accepts, grouped by subsystem. | diff --git a/docs/user-guide/verifiers.md b/docs/user-guide/verifiers.md new file mode 100644 index 0000000000..4972a36e85 --- /dev/null +++ b/docs/user-guide/verifiers.md @@ -0,0 +1,111 @@ +--- +title: Verifiers +description: Train on Verifiers environments with Miles. +--- + +Miles can train on a Verifiers environment in place of a prompt dataset. The +integration requires Python 3.11 or newer and Verifiers 0.2.0. Verifiers 0.2.1 +requires OpenAI 2.9 or newer, while SGLang 0.5.15 pins OpenAI 2.6.1. + +## Install + +Install the optional dependencies with Miles and install the Prime CLI: + +```bash +pip install -e '.[verifiers]' +uv tool install prime +``` + +The recommended workspace keeps local environment packages under `./environments`: + +```text +workspace/ + environments/ + my-environment/ +``` + +From the workspace root, install a local environment by name. For an environment from +the Environments Hub, authenticate and use its `user/environment` ID: + +```bash +# Local: ./environments/my-environment +prime env install my-environment + +# Environments Hub +prime login +prime env install user/my-environment +``` + +## Configure + +Create a Verifiers `EnvConfig` TOML file. A minimal config selects a taskset: + +```toml +[taskset] +id = "gsm8k-v1" +``` + +The config may also define the harness, runtime, judges, retries, and environment +limits supported by Verifiers. Verifiers applies per-rollout and group rewards before +the completed traces are returned to Miles. + +The integration implements Verifiers' V1 environment contract. Legacy V0 environment +configs are rejected during startup. + +## Run + +Add one option to a normal Miles training command: + +```bash +--verifiers-config /path/to/verifiers.toml +``` + +This uses the configured taskset instead of Miles prompt data. Environment behavior +comes from the Verifiers config, while Miles continues to own the model, sampling, +batching, concurrency, reward hooks, and optimizer settings. The Renderers library +formats environment messages with Miles' model and tokenizer settings. + +The standard Miles rollout options keep their existing meaning: + +| Miles option | Verifiers behavior | +|---|---| +| `--rollout-batch-size` | Number of task groups returned by each training rollout | +| `--n-samples-per-prompt` | Rollouts per training task | +| `--n-samples-per-eval-prompt` | Rollouts per evaluation task | +| `--rollout-shuffle` / `--rollout-seed` | Finite taskset order and sampling seeds | +| `--rollout-*` / `--eval-*` sampling options | Sampling and context limits | +| `--apply-chat-template-kwargs` | Typed template options passed to renderers | +| `--sglang-server-concurrency` | Physical engine capacity | +| Miles reward and filtering options | Applied after Verifiers scoring using the standard Miles hooks | + +Evaluation covers every task in the taskset. Training cycles the taskset and advances +from the current Miles rollout ID when a run resumes. + +## Environment Support + +The adapter supports V1 environments that use the Chat Completions dialect with +text-only Renderers inputs. Tools require a model-specific renderer; use a registered +model identity in `--hf-checkpoint` or the existing `--sglang-tokenizer-path` option. +User simulators, multi-turn episodes, environment runtimes, per-rollout rewards, and +group rewards run through the standard Verifiers environment lifecycle. + +Verifiers group rewards apply during both training and evaluation. Miles +`--group-rm` hooks remain training-only, matching the standard Miles rollout path. + +## Limitations + +`--partial-rollout` is not supported. A Verifiers episode owns live harness and +environment state and has no contract for resuming a partially executed episode. Miles +rejects this combination during argument validation. + +`--chat-template-path` is also rejected because Renderers owns message formatting for +Verifiers environments. Use the checkpoint's native template and +`--apply-chat-template-kwargs` instead. + +Streaming model requests, Responses and Anthropic dialects, multimodal inputs, OPD, +routing replay, and indexer replay are not supported by the transport. Miles +rejects the corresponding CLI options when they can be detected during startup. + +Traces with multiple graph branches, including compaction, are rejected. Miles does +not currently preserve a trace's rollout-group boundary when it flattens multiple +training samples, which would make group-relative advantages incorrect. diff --git a/miles/rollout/verifiers_rollout.py b/miles/rollout/verifiers_rollout.py new file mode 100644 index 0000000000..17301b9305 --- /dev/null +++ b/miles/rollout/verifiers_rollout.py @@ -0,0 +1,820 @@ +from __future__ import annotations + +import asyncio +import logging +import random +import sys +import uuid +from argparse import Namespace +from collections import OrderedDict +from collections.abc import Iterable +from importlib import metadata as importlib_metadata +from pathlib import Path +from types import SimpleNamespace +from typing import Any + +import httpx +from packaging.version import InvalidVersion, Version + +from miles.rollout.base_types import ( + RolloutFnConstructorInput, + RolloutFnEvalInput, + RolloutFnEvalOutput, + RolloutFnInput, + RolloutFnOutput, + RolloutFnTrainInput, + RolloutFnTrainOutput, +) +from miles.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter +from miles.rollout.generate_utils.prefill_logprobs import recompute_samples_rollout_logprobs_via_prefill +from miles.utils.lora import LORA_ADAPTER_NAME, is_lora_enabled +from miles.utils.types import Sample + +logger = logging.getLogger(__name__) + +_MIN_VERIFIERS_VERSION = Version("0.2.0") +_MAX_VERIFIERS_VERSION = Version("0.2.1") +_MIN_RENDERERS_VERSION = Version("0.1.8") +_UNSUPPORTED_ERROR_PREFIX = "Miles' Verifiers adapter does not support" + + +def _load_config_data(path: str) -> dict[str, Any]: + config_path = Path(path) + if config_path.suffix.lower() != ".toml": + raise ValueError("--verifiers-config must point to a Verifiers TOML config.") + if sys.version_info < (3, 11): + raise _optional_dependency_error() + + import tomllib + + data = tomllib.loads(config_path.read_text()) + if not isinstance(data, dict): + raise ValueError(f"{path} must contain a mapping at the root.") + return data + + +def _optional_dependency_error() -> RuntimeError: + return RuntimeError( + "Verifiers rollouts require Python 3.11+ and the optional dependencies. " + "Install Miles with `pip install -e '.[verifiers]'`." + ) + + +def _installed_version(package: str) -> str: + try: + return importlib_metadata.version(package) + except importlib_metadata.PackageNotFoundError as error: + raise _optional_dependency_error() from error + + +def _check_version(package: str, raw_version: str, minimum: Version, maximum: Version | None = None) -> None: + try: + installed = Version(raw_version) + except InvalidVersion as error: + raise RuntimeError(f"Could not parse installed {package} version {raw_version!r}.") from error + if installed < minimum: + raise RuntimeError(f"Verifiers rollouts require {package}>={minimum}; found {installed}.") + if maximum is not None and installed >= maximum: + raise RuntimeError( + f"Verifiers rollouts require {package}>={minimum},<{maximum}; found {installed}. " + "Verifiers 0.2.1 requires OpenAI>=2.9, while SGLang 0.5.15 pins OpenAI==2.6.1." + ) + + +def _import_verifiers(): + if sys.version_info < (3, 11): + raise _optional_dependency_error() + _check_version( + "verifiers", + _installed_version("verifiers"), + _MIN_VERIFIERS_VERSION, + _MAX_VERIFIERS_VERSION, + ) + _check_version("renderers", _installed_version("renderers"), _MIN_RENDERERS_VERSION) + try: + from verifiers.v1 import EnvConfig, Environment, ModelContext, SamplingConfig + from verifiers.v1.clients.train import TrainClient + from verifiers.v1.decorators import discover_decorated + from verifiers.v1.errors import OverlongPromptError, ProviderError + except ImportError as error: + raise _optional_dependency_error() from error + return SimpleNamespace( + EnvConfig=EnvConfig, + Environment=Environment, + discover_decorated=discover_decorated, + ModelContext=ModelContext, + OverlongPromptError=OverlongPromptError, + ProviderError=ProviderError, + SamplingConfig=SamplingConfig, + TrainClient=TrainClient, + ) + + +def _renderer_identity(checkpoint: str) -> str | None: + from renderers.base import MODEL_RENDERER_MAP + + if checkpoint in MODEL_RENDERER_MAP: + return checkpoint + + path = Path(checkpoint) + candidates = [] + for part in path.parts: + if part.startswith("models--"): + candidates.append(part.removeprefix("models--").replace("--", "/")) + candidates.extend(model_id for model_id in MODEL_RENDERER_MAP if model_id.rsplit("/", 1)[-1] == path.name) + matches = sorted(set(candidates) & MODEL_RENDERER_MAP.keys()) + if not matches: + return None + renderer_names = {MODEL_RENDERER_MAP[model_id] for model_id in matches} + return matches[0] if len(renderer_names) == 1 else None + + +def _train_client( + runtime, + args: Namespace, + model: str, + pool_size: int, + *, + router_args: Namespace | None = None, +): + tokenizer_source = getattr(args, "sglang_tokenizer_path", None) or model + identity = _renderer_identity(model) or _renderer_identity(tokenizer_source) + + # TrainClient uses one path for both tokenizer loading and renderer lookup. + # Miles checkpoints are commonly local snapshots, so keep local tokenizer + # files while restoring the canonical identity used by Renderers' registry. + class TrainClient(runtime.TrainClient): + @staticmethod + def _unsupported_request(kind: str): + return runtime.ProviderError( + f"{_UNSUPPORTED_ERROR_PREFIX} {kind}.", + status_code=400, + ) + + async def get_response(self, *args, **kwargs): + try: + return await super().get_response(*args, **kwargs) + except NotImplementedError as error: + raise runtime.ProviderError( + f"{_UNSUPPORTED_ERROR_PREFIX} this request: {error}", + status_code=400, + ) from error + except ValueError as error: + if "does not support tools" not in str(error): + raise + raise runtime.ProviderError( + f"{_UNSUPPORTED_ERROR_PREFIX} tools with this renderer: {error} " + "Use a Renderers-registered model identity in " + "--hf-checkpoint or --sglang-tokenizer-path.", + status_code=400, + ) from error + + async def relay(self, *args, **kwargs): + raise self._unsupported_request("streaming requests") + + async def relay_aux(self, *args, **kwargs): + raise self._unsupported_request("auxiliary dialect routes") + + def _renderer_pool(self, requested_model, *, chat_template_kwargs=None): + if identity is None: + return super()._renderer_pool( + requested_model, + chat_template_kwargs=chat_template_kwargs, + ) + if self._pool is None: + from renderers import RendererPool, create_renderer + from renderers.base import load_tokenizer + + source = self.renderer_model_name or requested_model + + def factory(): + tokenizer = load_tokenizer(source) + tokenizer.name_or_path = identity + return create_renderer( + tokenizer, + self.config, + chat_template_kwargs=chat_template_kwargs, + ) + + self._pool = RendererPool(factory, size=self.pool_size) + return self._pool + + return TrainClient( + MilesSGLangTransport(args, router_args=router_args), + pool_size=pool_size, + renderer_model_name=tokenizer_source, + ) + + +def _generate_url(args: Namespace, endpoint: str = "/generate") -> str: + routers = getattr(args, "sglang_model_routers", None) + if routers and "default" in routers: + ip, port = routers["default"] + else: + ip, port = args.sglang_router_ip, args.sglang_router_port + return f"http://{ip}:{port}{endpoint}" + + +async def _sglang_worker_urls(args: Namespace) -> list[str]: + from miles.utils.http_utils import get + + router_url = _generate_url(args).removesuffix("/generate") + if not getattr(args, "use_miles_router", False): + try: + response = await get(f"{router_url}/workers") + return [worker["url"] for worker in response["workers"]] + except Exception: + logger.debug("SGLang /workers lookup failed; trying Miles /list_workers.", exc_info=True) + response = await get(f"{router_url}/list_workers") + return list(response["urls"]) + + +def _finish_reason(output: dict[str, Any]) -> str: + finish_reason = (output.get("meta_info") or {}).get("finish_reason") + if isinstance(finish_reason, dict): + finish_reason = finish_reason.get("type") + if finish_reason == "abort": + raise RuntimeError("SGLang aborted the Verifiers generation request.") + return finish_reason if finish_reason in {"stop", "length", "content_filter"} else "stop" + + +class MilesSGLangTransport: + """Translate Renderers' vLLM generate wire format to Miles' SGLang endpoint.""" + + def __init__(self, args: Namespace, *, router_args: Namespace | None = None): + self.args = args + self.router_args = router_args or args + self._seen_sessions: OrderedDict[str, None] = OrderedDict() + self._session_cache_size = 10_000 + + @property + def base_url(self) -> str: + # The refactored RolloutManager constructs rollout functions before it + # starts SGLang and fills in the router address. + return f"{_generate_url(self.router_args, '').rstrip('/')}/v1" + + async def get(self, _path: str, **_kwargs) -> dict[str, list[Any]]: + # Renderers can discover a vLLM context cap here. Miles owns separate + # prompt and response limits, so the transport enforces them at POST time. + return {"data": []} + + def _sampling_params(self, raw: dict[str, Any], prompt_len: int) -> dict[str, Any]: + values = dict(raw) + values.pop("logprobs", None) + for source, target in ( + ("max_tokens", "max_new_tokens"), + ("min_tokens", "min_new_tokens"), + ("seed", "sampling_seed"), + ): + if source in values: + values[target] = values.pop(source) + + if self.args.rollout_stop is not None: + values.setdefault("stop", self.args.rollout_stop) + if self.args.rollout_stop_token_ids is not None: + renderer_stops = list(values.get("stop_token_ids") or []) + values["stop_token_ids"] = list(dict.fromkeys([*renderer_stops, *self.args.rollout_stop_token_ids])) + values["skip_special_tokens"] = self.args.rollout_skip_special_tokens + values["no_stop_trim"] = True + values["spaces_between_special_tokens"] = False + values["n"] = 1 + + context_limit = self.args.rollout_max_context_len + response_limit = self.args.rollout_max_response_len + requested = int(values.get("max_new_tokens", response_limit)) + if context_limit is not None: + requested = min(requested, context_limit - prompt_len) + values["max_new_tokens"] = min(requested, response_limit) + if values["max_new_tokens"] <= 0: + runtime = _import_verifiers() + raise runtime.OverlongPromptError( + f"prompt has {prompt_len} tokens, rollout_max_context_len={context_limit}" + ) + return values + + async def post(self, endpoint: str, *, body: dict[str, Any], options=None, **_kwargs) -> httpx.Response: + if body.get("features") is not None: + raise NotImplementedError("Miles Verifiers rollouts do not yet support multimodal renderer features.") + + prompt_ids = list(body["token_ids"]) + headers = dict((options or {}).get("headers") or {}) + session_id = headers.get("X-Session-ID") + if session_id: + if session_id in self._seen_sessions: + self._seen_sessions.move_to_end(session_id) + else: + max_prompt_len = getattr(self.args, "rollout_max_prompt_len", None) + if max_prompt_len is not None and len(prompt_ids) > max_prompt_len: + runtime = _import_verifiers() + raise runtime.OverlongPromptError( + f"initial prompt has {len(prompt_ids)} tokens, rollout_max_prompt_len={max_prompt_len}" + ) + self._seen_sessions[session_id] = None + if len(self._seen_sessions) > self._session_cache_size: + self._seen_sessions.popitem(last=False) + + payload: dict[str, Any] = { + "input_ids": prompt_ids, + "sampling_params": self._sampling_params(body.get("sampling_params") or {}, len(prompt_ids)), + "return_logprob": True, + } + if is_lora_enabled(self.args): + payload["lora_path"] = LORA_ADAPTER_NAME + if body.get("priority") is not None: + payload["priority"] = body["priority"] + if body.get("cache_salt") is not None: + payload["extra_key"] = body["cache_salt"] + + request_headers = None + if getattr(self.args, "sglang_router_policy", None) in ("consistent_hashing", "manual") and session_id: + request_headers = {"X-SMG-Routing-Key": session_id} + + from miles.utils.http_utils import post + + output = await post(_generate_url(self.router_args), payload, headers=request_headers) + meta_info = dict(output.get("meta_info") or {}) + token_logprobs = list(meta_info.get("output_token_logprobs") or []) + completion_ids = [int(item[1]) for item in token_logprobs] + completion_logprobs = [float(item[0]) for item in token_logprobs] + expected = int(meta_info.get("completion_tokens", len(completion_ids))) + if len(completion_ids) != expected: + raise RuntimeError( + "SGLang generate response has mismatched completion token metadata: " + f"{len(completion_ids)} != {expected}" + ) + + response_body = { + "request_id": output.get("request_id") or f"vf-{uuid.uuid4().hex}", + "choices": [ + { + "token_ids": completion_ids, + "logprobs": {"content": [{"logprob": value} for value in completion_logprobs]}, + "finish_reason": _finish_reason(output), + } + ], + } + request = httpx.Request("POST", endpoint) + return httpx.Response(200, json=response_body, request=request) + + async def close(self) -> None: + return None + + +def _sample_status(trace) -> Sample.Status: + if trace.has_error: + return Sample.Status.FAILED + if trace.is_truncated: + return Sample.Status.TRUNCATED + return Sample.Status.COMPLETED + + +def _serialize_prompt(prompt): + if isinstance(prompt, list): + return [ + message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message + for message in prompt + ] + return prompt or "" + + +def _validate_group_reward_sample_counts(args: Namespace, tasks, discover_decorated) -> None: + if not any(discover_decorated(task, "group_reward") for task in tasks): + return + if getattr(args, "num_rollout", None) != 0 and args.n_samples_per_prompt < 2: + raise ValueError("Verifiers tasks with @group_reward require --n-samples-per-prompt >= 2.") + if getattr(args, "eval_interval", None) is not None and args.n_samples_per_eval_prompt < 2: + raise ValueError("Verifiers tasks with @group_reward require --n-samples-per-eval-prompt >= 2.") + + +def _raise_for_unsupported_trace_errors(traces) -> None: + for trace in traces: + if trace.error is not None and trace.error.message.startswith(_UNSUPPORTED_ERROR_PREFIX): + raise RuntimeError(trace.error.message) + + +def _branch_to_sample(args: Namespace, trace, branch, *, group_index: int, index: int) -> Sample: + tokens = list(branch.token_ids) + sampled_mask = list(branch.sampled_mask) + logprobs = list(branch.logprobs) + if len(tokens) != len(sampled_mask) or len(tokens) != len(logprobs): + raise ValueError( + f"Trace {trace.id} token metadata mismatch: " + f"tokens={len(tokens)}, mask={len(sampled_mask)}, logprobs={len(logprobs)}" + ) + first_sampled = sampled_mask.index(True) if True in sampled_mask else len(tokens) + response_length = len(tokens) - first_sampled + reward = trace.reward if args.reward_key is None else {**trace.rewards, "reward": trace.reward} + task_data = trace.task.data + label = getattr(task_data, "label", None) + if label is None: + label = getattr(task_data, "answer", None) + + metadata = { + "verifiers": { + "branch_index": branch.index, + "task_index": getattr(task_data, "idx", None), + "rewards": dict(trace.rewards), + "metrics": dict(trace.metrics), + "stop_condition": trace.stop_condition, + } + } + if trace.error is not None: + metadata["verifiers"]["error"] = trace.error.model_dump(mode="json", exclude_none=True) + + sample = Sample( + group_index=group_index, + index=index, + prompt=_serialize_prompt(getattr(task_data, "prompt", "")), + tokens=tokens, + response=trace.last_reply, + response_length=response_length, + label=label, + reward=reward, + loss_mask=[int(value) for value in sampled_mask[first_sampled:]], + rollout_log_probs=logprobs[first_sampled:], + status=_sample_status(trace), + metadata=metadata, + routing_key=trace.id, + ) + sample.validate() + return sample + + +def trace_to_samples(args: Namespace, trace, *, group_index: int, index_start: int) -> list[Sample]: + if not trace.branches: + error = trace.error.model_dump(mode="json", exclude_none=True) if trace.error is not None else None + logger.warning( + "Verifiers trace %s has no graph branches; omitting it from training (error=%s).", + trace.id, + error, + ) + return [] + if len(trace.branches) != 1: + raise NotImplementedError( + "Miles cannot yet preserve Verifiers trace groups when a rollout produces " + f"multiple graph branches (trace {trace.id} produced {len(trace.branches)})." + ) + return [ + _branch_to_sample( + args, + trace, + trace.branches[0], + group_index=group_index, + index=index_start, + ) + ] + + +def trace_to_sample(args: Namespace, trace, *, group_index: int, index: int) -> Sample: + samples = trace_to_samples(args, trace, group_index=group_index, index_start=index) + if len(samples) != 1: + raise ValueError(f"Verifiers trace {trace.id} produced {len(samples)} branches, expected one.") + return samples[0] + + +def _trace_metrics(traces) -> dict[str, float]: + if not traces: + return {} + return { + "verifiers/reward_mean": sum(trace.reward for trace in traces) / len(traces), + "verifiers/error_rate": sum(trace.has_error for trace in traces) / len(traces), + "verifiers/truncated_rate": sum(trace.is_truncated for trace in traces) / len(traces), + "verifiers/num_turns_mean": sum(trace.num_turns for trace in traces) / len(traces), + } + + +def _trace_eval_reward(trace, reward_key: str | None): + if reward_key is None: + return trace.reward + rewards = {**trace.rewards, "reward": trace.reward} + if trace.has_error: + return rewards.get(reward_key) + return rewards[reward_key] + + +def _flatten_samples(values: Iterable[Any]) -> list[Sample]: + flattened = [] + for value in values: + if isinstance(value, list): + flattened.extend(_flatten_samples(value)) + else: + flattened.append(value) + return flattened + + +def _make_eval_args(args: Namespace) -> Namespace: + eval_args = Namespace(**vars(args)) + for eval_name, rollout_name in ( + ("eval_temperature", "rollout_temperature"), + ("eval_top_p", "rollout_top_p"), + ("eval_top_k", "rollout_top_k"), + ("eval_max_response_len", "rollout_max_response_len"), + ("eval_max_context_len", "rollout_max_context_len"), + ): + if (value := getattr(args, eval_name, None)) is not None: + setattr(eval_args, rollout_name, value) + eval_args.rollout_max_prompt_len = args.eval_max_prompt_len + eval_args.rollout_min_new_tokens = args.eval_min_new_tokens + eval_args.reward_key = args.eval_reward_key or args.reward_key + return eval_args + + +class VerifiersRolloutFn: + def __init__(self, input: RolloutFnConstructorInput): + runtime = _import_verifiers() + self.args = input.args + self.data_source = input.data_source + self.config = runtime.EnvConfig.model_validate(_load_config_data(self.args.verifiers_config)) + if self.config.is_legacy: + raise ValueError("Miles' Verifiers integration supports V1 environment configs only.") + if self.config.harness.id == "codex": + raise ValueError( + "Miles' Verifiers adapter does not support the Codex harness because it uses the Responses dialect." + ) + + self.env = runtime.Environment(self.config) + self.model = self.args.hf_checkpoint + self.sampling = self._sampling_config(runtime.SamplingConfig, self.args) + self.eval_args = _make_eval_args(self.args) + self.eval_sampling = self._sampling_config(runtime.SamplingConfig, self.eval_args) + + engine_count = self.args.rollout_num_gpus // self.args.rollout_num_gpus_per_engine + self.max_concurrent = self.args.sglang_server_concurrency * engine_count + pool_size = max(1, min(self.max_concurrent, 16)) + self.client = _train_client(runtime, self.args, self.model, pool_size) + self.eval_client = _train_client( + runtime, + self.eval_args, + self.model, + pool_size, + router_args=self.args, + ) + self.ctx = runtime.ModelContext(client=self.client, model=self.model, sampling=self.sampling) + self.eval_ctx = runtime.ModelContext(client=self.eval_client, model=self.model, sampling=self.eval_sampling) + + from miles.utils.misc import load_function + + self.dynamic_filter = load_function(self.args.dynamic_sampling_filter_path) + self._tasks = list(self.env.taskset.load()) + if self.args.rollout_shuffle: + random.Random(self.args.rollout_seed).shuffle(self._tasks) + if not self._tasks: + raise ValueError("Verifiers taskset selected zero tasks.") + _validate_group_reward_sample_counts(self.args, self._tasks, runtime.discover_decorated) + self._next_train_task_idx = None + self._next_group_index = 0 + self._next_sample_index = 0 + + @staticmethod + def _sampling_config(SamplingConfig, args: Namespace): + data: dict[str, Any] = { + "temperature": args.rollout_temperature, + "top_p": args.rollout_top_p, + "max_tokens": args.rollout_max_response_len, + } + if args.rollout_top_k is not None: + data["top_k"] = args.rollout_top_k + if (min_tokens := getattr(args, "rollout_min_new_tokens", None)) is not None: + data["min_tokens"] = min_tokens + if args.apply_chat_template_kwargs: + data["extra_body"] = {"chat_template_kwargs": args.apply_chat_template_kwargs} + return SamplingConfig.model_validate(data) + + def _task(self, index: int): + return self._tasks[index % len(self._tasks)] + + async def __call__(self, input: RolloutFnInput) -> RolloutFnOutput: + return await (self._call_eval(input) if input.evaluation else self._call_train(input)) + + async def _run_task_group(self, task, n: int, semaphore: asyncio.Semaphore, seed_base: int, ctx=None): + runtime = _import_verifiers() + ctx = ctx or self.ctx + episode = self.env.episode(task, ctx, n=n) + if getattr(self.args, "sglang_enable_deterministic_inference", False): + for offset, rollout in enumerate(episode.rollouts): + sampling = ctx.sampling.model_copy(update={"sampling_seed": seed_base + offset}) + rollout.ctx = runtime.ModelContext(client=ctx.client, model=self.model, sampling=sampling) + return await episode.run(semaphore) + + def _convert_group(self, traces, *, group_index: int, preserve_empty: bool = False): + group = [] + complete = True + for trace in traces: + converted = trace_to_samples( + self.args, + trace, + group_index=group_index, + index_start=self._next_sample_index, + ) + self._next_sample_index += len(converted) + if not converted: + complete = False + if preserve_empty: + group.append(None) + continue + group.append(converted[0]) + return group if complete or preserve_empty else [] + + async def _apply_miles_rewards(self, group) -> None: + from miles.rollout.rm_hub import async_rm, batched_async_rm + + samples = _flatten_samples(group) + if self.args.group_rm: + rewards = await batched_async_rm(self.args, samples) + elif self.args.custom_rm_path is not None or self.args.rm_type: + rewards = await asyncio.gather(*(async_rm(self.args, sample) for sample in samples)) + else: + return + if rewards is None or len(rewards) != len(samples): + raise ValueError("Miles reward model returned an unexpected number of rewards.") + for sample, reward in zip(samples, rewards, strict=True): + sample.reward = reward + + async def _postprocess_train_samples(self, data, all_data) -> None: + from miles.utils.misc import load_function + + if function := load_function(self.args.rollout_sample_filter_path): + function(self.args, data) + if function := load_function(self.args.rollout_all_samples_process_path): + function(self.args, all_data, self.data_source) + await recompute_samples_rollout_logprobs_via_prefill( + self.args, + _flatten_samples(data), + url=_generate_url(self.args), + sampling_params={ + "temperature": self.args.rollout_temperature, + "top_p": self.args.rollout_top_p, + "top_k": self.args.rollout_top_k, + "max_new_tokens": self.args.rollout_max_response_len, + }, + ) + + async def _cancel_pending(self, futures: Iterable[asyncio.Task]) -> None: + pending = [future for future in futures if not future.done()] + for future in pending: + future.cancel() + if pending: + try: + from miles.utils.http_utils import post + + urls = await _sglang_worker_urls(self.args) + await asyncio.gather( + *(post(f"{url}/abort_request", {"abort_all": True}) for url in urls), + return_exceptions=True, + ) + except Exception: + logger.exception("Failed to abort pending Verifiers requests.") + await asyncio.gather(*futures, return_exceptions=True) + + async def _call_train(self, input: RolloutFnTrainInput) -> RolloutFnTrainOutput: + from miles.utils import dumper_utils + + await dumper_utils.configure_sglang(self.args) + target = self.args.rollout_batch_size + if self._next_train_task_idx is None: + self._next_train_task_idx = input.rollout_id * target + + groups = [] + all_groups = [] + all_traces = [] + metrics = MetricGatherer() + semaphore = asyncio.Semaphore(self.max_concurrent) + pending: set[asyncio.Task] = set() + async with self.env.serving(): + try: + while len(groups) < target: + while len(groups) + len(pending) < target: + for _ in range(self.args.over_sampling_batch_size): + task_index = self._next_train_task_idx + self._next_train_task_idx += 1 + seed = self.args.rollout_seed + task_index * self.args.n_samples_per_prompt + pending.add( + asyncio.create_task( + self._run_task_group( + self._task(task_index), + self.args.n_samples_per_prompt, + semaphore, + seed, + ) + ) + ) + done, pending = await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED) + for future in done: + try: + traces = future.result() + except Exception: + logger.exception("Verifiers episode failed; resampling.") + metrics.on_dynamic_filter_drop(reason="episode_error") + continue + all_traces.extend(traces) + _raise_for_unsupported_trace_errors(traces) + if any(trace.has_error for trace in traces): + metrics.on_dynamic_filter_drop(reason="trace_error") + continue + group = self._convert_group(traces, group_index=self._next_group_index) + self._next_group_index += 1 + if len(group) != self.args.n_samples_per_prompt: + metrics.on_dynamic_filter_drop(reason="empty_trace") + continue + await self._apply_miles_rewards(group) + all_groups.append(group) + result = call_dynamic_filter( + self.dynamic_filter, + self.args, + _flatten_samples(group), + ) + if not result.keep: + metrics.on_dynamic_filter_drop(reason=result.reason) + continue + if len(groups) < target: + groups.append(group) + finally: + await self._cancel_pending(pending) + + groups.sort(key=lambda group: _flatten_samples(group)[0].index) + all_groups.sort(key=lambda group: _flatten_samples(group)[0].index) + await self._postprocess_train_samples(groups, all_groups) + output_metrics = metrics.collect() + output_metrics.update(_trace_metrics(all_traces)) + return RolloutFnTrainOutput(samples=groups, metrics=output_metrics) + + async def _call_eval(self, input: RolloutFnEvalInput) -> RolloutFnEvalOutput: + assert not self.args.group_rm, "Group RM is not supported for eval rollout" + + from miles.utils import dumper_utils + + await dumper_utils.configure_sglang(self.args) + semaphore = asyncio.Semaphore(self.max_concurrent) + async with self.env.serving(): + futures = [ + asyncio.create_task( + self._run_task_group( + task, + self.args.n_samples_per_eval_prompt, + semaphore, + self.args.rollout_seed + index * self.args.n_samples_per_eval_prompt, + self.eval_ctx, + ) + ) + for index, task in enumerate(self._tasks) + ] + try: + trace_groups = await asyncio.gather(*futures) + finally: + await self._cancel_pending(futures) + + samples = [] + rewards = [] + truncated = [] + all_traces = [] + reward_key = self.args.eval_reward_key or self.args.reward_key + use_miles_rewards = bool(self.args.custom_rm_path is not None or self.args.rm_type) + for group_index, traces in enumerate(trace_groups): + all_traces.extend(traces) + _raise_for_unsupported_trace_errors(traces) + group = self._convert_group(traces, group_index=group_index, preserve_empty=True) + trainable = [value for value in group if value is not None] + await self._apply_miles_rewards(trainable) + samples.extend(_flatten_samples(trainable)) + for trace, value in zip(traces, group, strict=True): + if use_miles_rewards and value is not None: + sample_rewards = [sample.get_reward_value(self.eval_args) for sample in _flatten_samples([value])] + reward = sum(sample_rewards) / len(sample_rewards) + else: + reward = trace.reward if use_miles_rewards else _trace_eval_reward(trace, reward_key) + rewards.append(reward) + truncated.append(trace.is_truncated) + return RolloutFnEvalOutput( + data={ + self.config.env_id + or "verifiers": { + "rewards": rewards, + "truncated": truncated, + "samples": samples, + } + }, + metrics=_trace_metrics(all_traces), + ) + + +_LEGACY_INSTANCES: dict[tuple[int, int, bool], VerifiersRolloutFn] = {} + + +def generate_rollout( + args: Namespace, + rollout_id: int, + data_source: Any, + evaluation: bool = False, +) -> RolloutFnTrainOutput | RolloutFnEvalOutput: + """Legacy Miles entrypoint backed by one persistent rollout adapter.""" + from miles.utils.async_utils import run + + # The refactored rollout manager constructs separate train and eval adapters. + # Preserve that lifecycle under the legacy function interface as well. + key = (id(args), id(data_source), evaluation) + adapter = _LEGACY_INSTANCES.get(key) + if adapter is None: + adapter = VerifiersRolloutFn(RolloutFnConstructorInput(args=args, data_source=data_source)) + _LEGACY_INSTANCES[key] = adapter + input = RolloutFnEvalInput(rollout_id) if evaluation else RolloutFnTrainInput(rollout_id) + return run(adapter(input)) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 6138f65cd5..2b79a52e23 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -40,6 +40,12 @@ def reset_arg(parser, name, **kwargs): _FT_CHOICES = ["rollout", "train"] +VERIFIERS_ROLLOUT_FUNCTION_PATH = ( + "miles.rollout.verifiers_rollout.VerifiersRolloutFn" + if enable_experimental_rollout_refactor() + else "miles.rollout.verifiers_rollout.generate_rollout" +) + def get_miles_extra_args_provider(add_custom_arguments=None): def add_miles_arguments(parser): @@ -348,6 +354,15 @@ def add_rollout_arguments(parser): "and `truncated`." ), ) + parser.add_argument( + "--verifiers-config", + type=str, + default=None, + help=( + "Path to a Verifiers EnvConfig TOML file. Setting this uses the built-in Verifiers " + "rollout and disables Miles prompt-data loading." + ), + ) parser.add_argument( "--rollout-temperature", type=float, @@ -2293,6 +2308,35 @@ def miles_validate_args(args): args.ft_components = _resolve_ft_components(args) args.eval_datasets = _resolve_eval_datasets(args) + if args.verifiers_config is not None: + args.rollout_function_path = VERIFIERS_ROLLOUT_FUNCTION_PATH + args.rollout_global_dataset = False + + if args.partial_rollout: + raise ValueError( + "--partial-rollout is not supported for Verifiers because an episode " + "cannot be resumed from partially executed environment state." + ) + if args.multimodal_keys is not None: + raise ValueError( + "--multimodal-keys is not supported with --verifiers-config. " + "The Verifiers transport supports text-only renderer inputs." + ) + unsupported = [ + flag + for enabled, flag in ( + (args.use_opd, "--use-opd"), + (args.use_rollout_routing_replay, "--use-rollout-routing-replay"), + (args.use_rollout_indexer_replay, "--use-rollout-indexer-replay"), + ) + if enabled + ] + if unsupported: + raise ValueError( + f"{', '.join(unsupported)} is not supported with --verifiers-config because " + "the Verifiers SGLang transport does not preserve its additional token metadata." + ) + if args.mini_ft_controller_enable and args.control_server_port == 0: raise ValueError("--mini-ft-controller-enable requires --control-server-port to be set (non-zero)") @@ -2404,6 +2448,12 @@ def miles_validate_args(args): if not os.path.isfile(args.chat_template_path): raise FileNotFoundError(f"--chat-template-path file not found: {args.chat_template_path}") args.sglang_chat_template = args.chat_template_path + if args.verifiers_config is not None: + raise ValueError( + "--chat-template-path is not supported with --verifiers-config because renderers " + "does not accept a custom Jinja template. Use the checkpoint's template and " + "--apply-chat-template-kwargs." + ) if args.kl_coef != 0 or args.use_kl_loss: if not os.path.exists(args.ref_load): @@ -2476,7 +2526,7 @@ def miles_validate_args(args): args.ckpt_step = args.ref_ckpt_step args.start_rollout_id = 0 - if args.eval_interval is not None: + if args.eval_interval is not None and args.verifiers_config is None: assert args.eval_datasets, "Evaluation datasets must be configured when eval_interval is set." if args.save_interval is not None: diff --git a/setup.py b/setup.py index de7281b742..216d3450fa 100644 --- a/setup.py +++ b/setup.py @@ -51,6 +51,12 @@ def get_tag(self): "mlflow": [ "mlflow>=2.0", ], + "verifiers": [ + "verifiers>=0.2.0,<0.2.1", + "renderers>=0.1.8", + # SGLang 0.5.15 pins OpenAI 2.6.1; newer Verifiers and Agents require a newer SDK. + "openai-agents<0.5", + ], }, python_requires=">=3.10", classifiers=[ diff --git a/tests/e2e/long/test_qwen3_0.6B_verifiers.py b/tests/e2e/long/test_qwen3_0.6B_verifiers.py new file mode 100644 index 0000000000..8920dbd7c8 --- /dev/null +++ b/tests/e2e/long/test_qwen3_0.6B_verifiers.py @@ -0,0 +1,138 @@ +import math +import os +import shutil +import sys +from collections import Counter +from pathlib import Path + +import pytest +import torch +from tests.ci.ci_register import register_cuda_ci + +import miles.utils.external_utils.command_utils as U + +register_cuda_ci(est_time=900, suite="stage-c-2-gpu-h200", labels=["long"]) + +MODEL_NAME = "Qwen3-0.6B" +MODEL_TYPE = "qwen3-0.6B" +NUM_GPUS = 2 +MODEL_DIR = Path(os.environ.get("MILES_E2E_MODEL_DIR", "/root/models")) +MEGATRON_PATH = Path(os.environ.get("MILES_E2E_MEGATRON_PATH", "/root/Megatron-LM")) +RUN_DIR = Path(os.environ.get("MILES_E2E_RUN_DIR", "/tmp/miles-verifiers-e2e")) +VERIFIERS_DIR = Path("/tmp/verifiers-v0.2.0") + + +def prepare(): + U.exec_command(f"mkdir -p {MODEL_DIR} {RUN_DIR}") + U.exec_command(f"hf download Qwen/{MODEL_NAME} --local-dir {MODEL_DIR}/{MODEL_NAME}") + U.exec_command(f"{sys.executable} -m pip install -e '{U.repo_base_dir}[verifiers]'") + U.exec_command("uv tool install 'prime==0.6.19'") + if not VERIFIERS_DIR.exists(): + U.exec_command( + f"git clone --depth 1 --branch v0.2.0 " + f"https://github.com/PrimeIntellect-ai/verifiers.git {VERIFIERS_DIR}" + ) + shutil.copytree( + VERIFIERS_DIR / "environments" / "code_golf_v1", + RUN_DIR / "environments" / "code_golf_v1", + dirs_exist_ok=True, + ) + U.exec_command(f"cd {RUN_DIR} && prime --plain env install code-golf-v1") + U.convert_checkpoint( + model_name=MODEL_NAME, + megatron_model_type=MODEL_TYPE, + num_gpus_per_node=NUM_GPUS, + dir_dst=str(MODEL_DIR), + hf_checkpoint=str(MODEL_DIR / MODEL_NAME), + megatron_path=str(MEGATRON_PATH), + ) + + +def execute(): + config_path = RUN_DIR / "code-golf.toml" + config_path.write_text('[taskset]\nid = "code-golf-v1"\n') + dump_dir = RUN_DIR / "dump" + + train_args = " ".join( + [ + f"--hf-checkpoint {MODEL_DIR}/{MODEL_NAME}", + "--sglang-tokenizer-path Qwen/Qwen3-0.6B", + f"--ref-load {MODEL_DIR}/{MODEL_NAME}_torch_dist", + f"--verifiers-config {config_path}", + "--num-rollout 1", + "--rollout-batch-size 3", + "--n-samples-per-prompt 4", + "--over-sampling-batch-size 3", + "--rollout-max-response-len 512", + "--rollout-max-context-len 2048", + "--rollout-temperature 0.8", + "--global-batch-size 12", + "--balance-data", + "--advantage-estimator grpo", + "--entropy-coef 0.0", + "--eps-clip 0.2", + "--eps-clip-high 0.28", + "--optimizer adam", + "--lr 1e-6", + "--lr-decay-style constant", + "--weight-decay 0.1", + "--adam-beta1 0.9", + "--adam-beta2 0.98", + "--no-gradient-accumulation-fusion", + "--rollout-num-gpus-per-engine 1", + "--sglang-mem-fraction-static 0.6", + "--sglang-enable-metrics", + "--tensor-model-parallel-size 1", + "--pipeline-model-parallel-size 1", + "--context-parallel-size 1", + "--use-dynamic-batch-size", + "--max-tokens-per-gpu 4096", + "--actor-num-nodes 1", + f"--actor-num-gpus-per-node {NUM_GPUS}", + "--colocate", + f"--dump-details {dump_dir}", + U.get_default_wandb_args(__file__), + ] + ) + + U.execute_train( + train_args=train_args, + num_gpus_per_node=NUM_GPUS, + megatron_model_type=MODEL_TYPE, + extra_env_vars={"MILES_EXPERIMENTAL_ROLLOUT_REFACTOR": "0"}, + megatron_path=str(MEGATRON_PATH), + ) + verify(dump_dir) + + +def verify(dump_dir: Path): + samples = torch.load(dump_dir / "rollout_data" / "0.pt", weights_only=False)["samples"] + + assert len(samples) == 12 + assert Counter(sample["group_index"] for sample in samples) == {0: 4, 1: 4, 2: 4} + assert {sample["status"] for sample in samples} <= {"completed", "truncated"} + assert any(sample["status"] == "completed" for sample in samples) + assert all(math.isfinite(sample["reward"]) for sample in samples) + assert all( + len(sample["rollout_log_probs"]) == len(sample["loss_mask"]) == sample["response_length"] for sample in samples + ) + assert all(math.isfinite(value) for sample in samples for value in sample["rollout_log_probs"]) + + for sample in samples: + metadata = sample["metadata"]["verifiers"] + assert set(metadata["rewards"]) == {"correct", "fastest", "most_concise"} + assert metadata["metrics"]["passed"] in {0.0, 1.0} + assert math.isfinite(metadata["metrics"]["latency"]) + assert "error" not in metadata + + assert any(sample["metadata"]["verifiers"]["metrics"]["passed"] for sample in samples) + for group_index in range(3): + group = [sample for sample in samples if sample["group_index"] == group_index] + assert sum(sample["metadata"]["verifiers"]["rewards"]["fastest"] for sample in group) == pytest.approx(0.5) + + +if __name__ == "__main__": + prepare() + for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"): + os.environ.pop(proxy_var, None) + execute() diff --git a/tests/fast/rollout/test_verifiers_rollout.py b/tests/fast/rollout/test_verifiers_rollout.py new file mode 100644 index 0000000000..5aca9771f5 --- /dev/null +++ b/tests/fast/rollout/test_verifiers_rollout.py @@ -0,0 +1,639 @@ +import asyncio +import sys +from argparse import Namespace +from contextlib import asynccontextmanager +from types import SimpleNamespace + +import pytest +from packaging.version import Version + +if sys.version_info < (3, 11): + pytest.skip("Verifiers requires Python 3.11+", allow_module_level=True) + +from miles.rollout.verifiers_rollout import ( + MilesSGLangTransport, + VerifiersRolloutFn, + _check_version, + _finish_reason, + _load_config_data, + _make_eval_args, + _raise_for_unsupported_trace_errors, + _renderer_identity, + _trace_eval_reward, + _train_client, + _validate_group_reward_sample_counts, + trace_to_sample, + trace_to_samples, +) +from miles.utils.types import Sample + + +def _args(**overrides) -> Namespace: + values = { + "lora_adapter_path": None, + "lora_rank": 0, + "reward_key": None, + "rollout_max_context_len": 64, + "rollout_max_prompt_len": None, + "rollout_max_response_len": 8, + "rollout_skip_special_tokens": True, + "rollout_stop": None, + "rollout_stop_token_ids": None, + "sglang_model_routers": None, + "sglang_router_ip": "127.0.0.1", + "sglang_router_policy": "round_robin", + "sglang_router_port": 30000, + "sglang_tokenizer_path": None, + } + values.update(overrides) + return Namespace(**values) + + +def _branch(*, index=0, token_ids=None, sampled_mask=None, logprobs=None): + return SimpleNamespace( + index=index, + token_ids=[10, 11, 20, 21, 22] if token_ids is None else token_ids, + sampled_mask=[False, False, True, False, True] if sampled_mask is None else sampled_mask, + logprobs=[0.0, 0.0, -0.1, 0.0, -0.2] if logprobs is None else logprobs, + ) + + +def _trace(**overrides): + values = { + "id": "trace-1", + "branches": [_branch()], + "task": SimpleNamespace(data=SimpleNamespace(prompt="solve this", idx="task-1")), + "rewards": {"score": 1.25, "bonus": 0.75}, + "metrics": {"turns": 2.0}, + "stop_condition": "done", + "error": None, + "has_error": False, + "is_truncated": False, + "reward": 2.0, + "last_reply": "answer", + "num_turns": 2, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def test_config_loader_uses_verifiers_toml_format(tmp_path): + path = tmp_path / "verifiers.toml" + path.write_text('[taskset]\nid = "gsm8k-v1"\n') + + assert _load_config_data(path) == {"taskset": {"id": "gsm8k-v1"}} + + +def test_config_loader_rejects_non_toml_formats(tmp_path): + path = tmp_path / "verifiers.yaml" + path.write_text("taskset: gsm8k-v1\n") + + with pytest.raises(ValueError, match="Verifiers TOML"): + _load_config_data(path) + + +def test_verifiers_0_2_1_is_rejected_with_compatibility_reason(): + with pytest.raises(RuntimeError, match="SGLang 0.5.15 pins OpenAI"): + _check_version("verifiers", "0.2.1", Version("0.2.0"), Version("0.2.1")) + + +def test_transport_does_not_treat_aborted_generation_as_complete(): + with pytest.raises(RuntimeError, match="aborted"): + _finish_reason({"meta_info": {"finish_reason": {"type": "abort"}}}) + + +@pytest.mark.parametrize( + ("checkpoint", "expected"), + [ + ("/models/Qwen3-4B-Instruct-2507", "Qwen/Qwen3-4B-Instruct-2507"), + ( + "/cache/models--Qwen--Qwen3-4B-Instruct-2507/snapshots/revision", + "Qwen/Qwen3-4B-Instruct-2507", + ), + ("/models/private-finetune", None), + ], +) +def test_renderer_identity_is_inferred_from_standard_checkpoint_paths(checkpoint, expected): + pytest.importorskip("renderers", minversion="0.1.8") + + assert _renderer_identity(checkpoint) == expected + + +def test_train_client_uses_local_tokenizer_with_inferred_renderer_identity(monkeypatch): + renderers = pytest.importorskip("renderers", minversion="0.1.8") + checkpoint = "/cache/models--Qwen--Qwen3-4B-Instruct-2507/snapshots/revision" + seen = {} + + class BaseTrainClient: + def __init__(self, openai, pool_size, config=None, renderer_model_name=None): + self.openai = openai + self.pool_size = pool_size + self.config = config + self.renderer_model_name = renderer_model_name + self._pool = None + + class RendererPool: + def __init__(self, factory, size): + seen["size"] = size + seen["renderer"] = factory() + + def load_tokenizer(source): + seen["source"] = source + return SimpleNamespace(name_or_path=source) + + def create_renderer(tokenizer, config, *, chat_template_kwargs=None): + seen["identity"] = tokenizer.name_or_path + seen["config"] = config + seen["kwargs"] = chat_template_kwargs + return "renderer" + + monkeypatch.setattr(renderers, "RendererPool", RendererPool) + monkeypatch.setattr(renderers, "create_renderer", create_renderer) + monkeypatch.setattr("renderers.base.load_tokenizer", load_tokenizer) + runtime = SimpleNamespace(TrainClient=BaseTrainClient) + + args = _args(sglang_tokenizer_path="/models/custom-tokenizer") + client = _train_client(runtime, args, checkpoint, pool_size=3) + pool = client._renderer_pool(checkpoint, chat_template_kwargs={"enable_thinking": False}) + + assert isinstance(pool, RendererPool) + assert seen == { + "size": 3, + "source": "/models/custom-tokenizer", + "identity": "Qwen/Qwen3-4B-Instruct-2507", + "config": None, + "kwargs": {"enable_thinking": False}, + "renderer": "renderer", + } + + +@pytest.mark.asyncio +async def test_train_client_reports_unsupported_tool_renderer_as_configuration_error(): + class ProviderError(Exception): + def __init__(self, message, *, status_code): + super().__init__(message) + self.status_code = status_code + + class BaseTrainClient: + def __init__(self, openai, pool_size, config=None, renderer_model_name=None): + pass + + async def get_response(self, *args, **kwargs): + raise ValueError("RendererPool does not support tools.") + + runtime = SimpleNamespace(ProviderError=ProviderError, TrainClient=BaseTrainClient) + client = _train_client(runtime, _args(), "/models/private-finetune", pool_size=1) + + with pytest.raises(ProviderError, match="--sglang-tokenizer-path") as error: + await client.get_response() + + assert error.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_train_client_reports_unsupported_dialect_as_configuration_error(): + class ProviderError(Exception): + def __init__(self, message, *, status_code): + super().__init__(message) + self.status_code = status_code + + class BaseTrainClient: + def __init__(self, openai, pool_size, config=None, renderer_model_name=None): + pass + + async def get_response(self, *args, **kwargs): + raise NotImplementedError("only the chat-completions dialect is supported") + + runtime = SimpleNamespace(ProviderError=ProviderError, TrainClient=BaseTrainClient) + client = _train_client(runtime, _args(), "/models/test", pool_size=1) + + with pytest.raises(ProviderError, match="does not support this request") as error: + await client.get_response() + + assert error.value.status_code == 400 + + +@pytest.mark.parametrize(("method", "message"), [("relay", "streaming"), ("relay_aux", "auxiliary")]) +@pytest.mark.asyncio +async def test_train_client_reports_unsupported_relay_paths_as_configuration_errors(method, message): + class ProviderError(Exception): + def __init__(self, text, *, status_code): + super().__init__(text) + self.status_code = status_code + + class BaseTrainClient: + def __init__(self, openai, pool_size, config=None, renderer_model_name=None): + pass + + runtime = SimpleNamespace(ProviderError=ProviderError, TrainClient=BaseTrainClient) + client = _train_client(runtime, _args(), "/models/test", pool_size=1) + + with pytest.raises(ProviderError, match=message) as error: + await getattr(client, method)() + + assert error.value.status_code == 400 + + +def test_trace_to_sample_preserves_training_fields_and_verifiers_reward(): + sample = trace_to_sample(_args(), _trace(), group_index=3, index=9) + + assert sample.group_index == 3 + assert sample.index == 9 + assert sample.prompt == "solve this" + assert sample.tokens == [10, 11, 20, 21, 22] + assert sample.response == "answer" + assert sample.response_length == 3 + assert sample.loss_mask == [1, 0, 1] + assert sample.rollout_log_probs == [-0.1, 0.0, -0.2] + assert sample.reward == 2.0 + assert sample.routing_key == "trace-1" + assert sample.status == Sample.Status.COMPLETED + assert sample.metadata["verifiers"]["task_index"] == "task-1" + sample.validate() + + +def test_trace_to_sample_preserves_named_rewards_when_reward_key_is_set(): + args = _args(reward_key="score") + + sample = trace_to_sample(args, _trace(), group_index=0, index=0) + + assert sample.reward == {"score": 1.25, "bonus": 0.75, "reward": 2.0} + assert sample.get_reward_value(args) == 1.25 + + +def test_trace_to_sample_serializes_structured_prompt_messages(): + class Message: + def model_dump(self, **kwargs): + assert kwargs == {"mode": "json", "exclude_none": True} + return {"role": "user", "content": "solve this"} + + trace = _trace(task=SimpleNamespace(data=SimpleNamespace(prompt=[Message()], idx="task-1"))) + + sample = trace_to_sample(_args(), trace, group_index=0, index=0) + + assert sample.prompt == [{"role": "user", "content": "solve this"}] + + +def test_trace_to_sample_marks_error_before_truncation(): + error = SimpleNamespace(model_dump=lambda **_kwargs: {"type": "ProviderError"}) + + sample = trace_to_sample( + _args(), + _trace(error=error, has_error=True, is_truncated=True), + group_index=0, + index=0, + ) + + assert sample.status == Sample.Status.FAILED + assert sample.metadata["verifiers"]["error"] == {"type": "ProviderError"} + + +def test_failed_eval_trace_with_missing_named_reward_returns_none(): + trace = _trace(has_error=True, rewards={}) + + assert _trace_eval_reward(trace, "score") is None + + +def test_successful_eval_trace_requires_configured_named_reward(): + with pytest.raises(KeyError, match="score"): + _trace_eval_reward(_trace(rewards={}), "score") + + +def test_unsupported_trace_error_is_not_resampled_forever(): + error = SimpleNamespace(message="Miles' Verifiers adapter does not support this request: ResponsesDialect") + + with pytest.raises(RuntimeError, match="ResponsesDialect"): + _raise_for_unsupported_trace_errors([_trace(error=error, has_error=True)]) + + +def test_graph_branches_fail_before_miles_can_corrupt_trace_groups(): + trace = _trace(branches=[_branch(index=0), _branch(index=1)]) + + with pytest.raises(NotImplementedError, match="multiple graph branches"): + trace_to_samples(_args(), trace, group_index=4, index_start=10) + + +def test_convert_group_uses_standard_miles_group_shape(): + from miles.ray.rollout.rollout_data_conversion import postprocess_rollout_data + + adapter = object.__new__(VerifiersRolloutFn) + adapter.args = _args() + adapter._next_sample_index = 0 + + group = adapter._convert_group( + [_trace(id="first"), _trace(id="second")], + group_index=2, + ) + + assert len(group) == 2 + assert all(isinstance(sample, Sample) for sample in group) + args = SimpleNamespace( + disable_rollout_trim_samples=True, + global_batch_size=1, + use_dynamic_global_batch_size=False, + ) + flattened, _ = postprocess_rollout_data(args, [group], train_parallel_config={"dp_size": 1}) + assert [sample.routing_key for sample in flattened] == ["first", "second"] + + +def test_standard_dynamic_filter_accepts_converted_branch_group(): + from miles.rollout.filter_hub.dynamic_sampling_filters import check_reward_nonzero_std + + adapter = object.__new__(VerifiersRolloutFn) + adapter.args = _args() + adapter._next_sample_index = 0 + group = adapter._convert_group( + [_trace(id="low", reward=0.0), _trace(id="high", reward=1.0)], + group_index=0, + ) + + assert check_reward_nonzero_std(adapter.args, group).keep + + +@pytest.mark.parametrize( + ("overrides", "option"), + [ + ( + {"eval_interval": None, "n_samples_per_prompt": 1, "n_samples_per_eval_prompt": 1}, + "--n-samples-per-prompt", + ), + ( + {"eval_interval": 1, "n_samples_per_prompt": 2, "n_samples_per_eval_prompt": 1}, + "--n-samples-per-eval-prompt", + ), + ], +) +def test_group_reward_tasks_require_multiple_rollouts(overrides, option): + args = _args(**overrides) + tasks = [object()] + + with pytest.raises(ValueError, match=option): + _validate_group_reward_sample_counts(args, tasks, lambda _task, _kind: [object()]) + + +def test_group_reward_eval_count_is_ignored_when_eval_is_disabled(): + args = _args( + eval_interval=None, + n_samples_per_prompt=2, + n_samples_per_eval_prompt=1, + ) + + _validate_group_reward_sample_counts(args, [object()], lambda _task, _kind: [object()]) + + +def test_group_reward_train_count_is_ignored_for_eval_only_runs(): + args = _args( + num_rollout=0, + eval_interval=1, + n_samples_per_prompt=1, + n_samples_per_eval_prompt=2, + ) + + _validate_group_reward_sample_counts(args, [object()], lambda _task, _kind: [object()]) + + +@pytest.mark.asyncio +async def test_verifiers_episode_owns_group_reward_computation(): + traces = [_trace(id="a", reward=0.0), _trace(id="b", reward=0.0)] + + class Episode: + rollouts = [] + + async def run(self, semaphore): + assert semaphore is not None + traces[0].reward = -1.0 + traces[1].reward = 1.0 + return traces + + class Environment: + def episode(self, task, ctx, n): + assert task == "task" + assert ctx == "ctx" + assert n == 2 + return Episode() + + adapter = object.__new__(VerifiersRolloutFn) + adapter.args = _args(sglang_enable_deterministic_inference=False) + adapter.env = Environment() + adapter.ctx = "ctx" + + result = await adapter._run_task_group("task", 2, asyncio.Semaphore(2), seed_base=0) + + assert [trace.reward for trace in result] == [-1.0, 1.0] + + +def test_sampling_config_preserves_miles_minimum_tokens(): + class SamplingConfig: + @staticmethod + def model_validate(data): + return data + + config = VerifiersRolloutFn._sampling_config( + SamplingConfig, + _args( + apply_chat_template_kwargs={}, + rollout_min_new_tokens=3, + rollout_temperature=0.7, + rollout_top_k=20, + rollout_top_p=0.9, + ), + ) + + assert config["min_tokens"] == 3 + + +def test_eval_args_clear_training_prompt_cap_and_preserve_other_fallbacks(): + args = _args( + eval_max_context_len=128, + eval_max_prompt_len=None, + eval_max_response_len=None, + eval_min_new_tokens=None, + eval_reward_key=None, + eval_temperature=None, + eval_top_k=None, + eval_top_p=None, + reward_key="score", + rollout_max_context_len=64, + rollout_max_prompt_len=32, + rollout_max_response_len=8, + ) + + eval_args = _make_eval_args(args) + + assert eval_args.rollout_max_context_len == 128 + assert eval_args.rollout_max_prompt_len is None + assert eval_args.rollout_max_response_len == 8 + assert eval_args.reward_key == "score" + + +@pytest.mark.asyncio +async def test_transport_translates_renderer_request_to_sglang(monkeypatch): + requests = [] + + async def fake_post(url, payload, headers=None): + requests.append((url, payload, headers)) + return { + "request_id": "request-id", + "meta_info": { + "completion_tokens": 2, + "finish_reason": {"type": "stop"}, + "output_token_logprobs": [[-0.1, 20], [-0.2, 21]], + }, + } + + monkeypatch.setattr("miles.utils.http_utils.post", fake_post) + transport = MilesSGLangTransport(_args(sglang_router_policy="manual")) + + response = await transport.post( + "http://127.0.0.1:30000/inference/v1/generate", + body={ + "model": "test/model", + "token_ids": [10, 11], + "sampling_params": {"temperature": 0.2, "max_tokens": 2, "stop_token_ids": [99], "logprobs": 1}, + }, + options={"headers": {"X-Session-ID": "trace-id"}}, + ) + + assert response.json()["choices"][0]["token_ids"] == [20, 21] + assert requests == [ + ( + "http://127.0.0.1:30000/generate", + { + "input_ids": [10, 11], + "sampling_params": { + "temperature": 0.2, + "stop_token_ids": [99], + "max_new_tokens": 2, + "skip_special_tokens": True, + "no_stop_trim": True, + "spaces_between_special_tokens": False, + "n": 1, + }, + "return_logprob": True, + }, + {"X-SMG-Routing-Key": "trace-id"}, + ) + ] + + +@pytest.mark.asyncio +async def test_transport_rejects_multimodal_features(): + transport = MilesSGLangTransport(_args()) + + with pytest.raises(NotImplementedError, match="multimodal"): + await transport.post( + "http://127.0.0.1:30000/inference/v1/generate", + body={"token_ids": [1], "sampling_params": {}, "features": {}}, + ) + + +@pytest.mark.asyncio +async def test_transport_bounds_seen_sessions(monkeypatch): + async def fake_post(_url, _payload, headers=None): + return { + "meta_info": { + "completion_tokens": 1, + "finish_reason": {"type": "stop"}, + "output_token_logprobs": [[-0.1, 20]], + } + } + + monkeypatch.setattr("miles.utils.http_utils.post", fake_post) + transport = MilesSGLangTransport(_args()) + transport._session_cache_size = 2 + body = {"token_ids": [10], "sampling_params": {}} + + for session_id in ("one", "two", "three"): + await transport.post( + "http://127.0.0.1:30000/inference/v1/generate", + body=body, + options={"headers": {"X-Session-ID": session_id}}, + ) + + assert list(transport._seen_sessions) == ["two", "three"] + + +def test_transport_resolves_router_address_lazily(): + args = _args(sglang_router_ip=None, sglang_router_port=None) + transport = MilesSGLangTransport(args) + + args.sglang_router_ip = "10.0.0.4" + args.sglang_router_port = 3210 + + assert transport.base_url == "http://10.0.0.4:3210/v1" + + +def test_transport_uses_default_model_router(): + args = _args( + sglang_model_routers={ + "default": ("10.0.0.5", 3211), + "ref": ("10.0.0.6", 3212), + } + ) + + assert MilesSGLangTransport(args).base_url == "http://10.0.0.5:3211/v1" + + +def test_eval_transport_keeps_live_router_args(): + rollout_args = _args(sglang_router_ip=None, sglang_router_port=None) + eval_args = _args(sglang_router_ip=None, sglang_router_port=None) + transport = MilesSGLangTransport(eval_args, router_args=rollout_args) + + rollout_args.sglang_model_routers = {"default": ("10.0.0.7", 3213)} + + assert transport.base_url == "http://10.0.0.7:3213/v1" + + +@pytest.mark.asyncio +async def test_eval_rejects_miles_group_reward_model(): + adapter = object.__new__(VerifiersRolloutFn) + adapter.args = _args(group_rm=True) + + with pytest.raises(AssertionError, match="Group RM is not supported for eval rollout"): + await adapter._call_eval(SimpleNamespace(rollout_id=0)) + + +@pytest.mark.asyncio +async def test_eval_extracts_structured_miles_reward(monkeypatch): + import miles.utils as miles_utils + + @asynccontextmanager + async def serving(): + yield + + async def configure_sglang(_args): + return None + + async def run_group(*_args, **_kwargs): + return [_trace(rewards={"bonus": 0.25})] + + async def apply_reward(group): + group[0].reward = {"score": 0.75, "details": "ok"} + + dumper_utils = SimpleNamespace(configure_sglang=configure_sglang) + monkeypatch.setitem(sys.modules, "miles.utils.dumper_utils", dumper_utils) + monkeypatch.setattr(miles_utils, "dumper_utils", dumper_utils, raising=False) + adapter = object.__new__(VerifiersRolloutFn) + adapter.args = _args( + custom_rm_path="tests.fake_reward", + eval_reward_key="score", + group_rm=False, + n_samples_per_eval_prompt=1, + rm_type=None, + rollout_batch_size=1, + rollout_seed=1, + ) + adapter.eval_args = _args(reward_key="score") + adapter.config = SimpleNamespace(env_id="test-env") + adapter.env = SimpleNamespace(serving=serving) + adapter.eval_ctx = object() + adapter.max_concurrent = 1 + adapter.model = "test/model" + adapter._tasks = [object()] + adapter._next_sample_index = 0 + adapter._run_task_group = run_group + adapter._apply_miles_rewards = apply_reward + + output = await adapter._call_eval(SimpleNamespace(rollout_id=0)) + + assert output.data["test-env"]["rewards"] == [0.75] diff --git a/tests/fast/rollout/test_verifiers_runtime.py b/tests/fast/rollout/test_verifiers_runtime.py new file mode 100644 index 0000000000..8f322462d6 --- /dev/null +++ b/tests/fast/rollout/test_verifiers_runtime.py @@ -0,0 +1,123 @@ +import sys +from argparse import Namespace +from types import SimpleNamespace + +import pytest + +if sys.version_info < (3, 11): + pytest.skip("Verifiers requires Python 3.11+", allow_module_level=True) + +pytest.importorskip("verifiers", minversion="0.2.0") +pytest.importorskip("renderers", minversion="0.1.8") + +from verifiers.v1.clients.train import TrainClient +from verifiers.v1.dialects import ChatDialect, ResponsesDialect +from verifiers.v1.env import EnvConfig, Environment +from verifiers.v1.types import SamplingConfig + +from miles.rollout.verifiers_rollout import MilesSGLangTransport + + +def _args(**overrides): + values = { + "lora_adapter_path": None, + "lora_rank": 0, + "rollout_max_context_len": 64, + "rollout_max_prompt_len": None, + "rollout_max_response_len": 8, + "rollout_skip_special_tokens": True, + "rollout_stop": None, + "rollout_stop_token_ids": None, + "sglang_model_routers": None, + "sglang_router_ip": "127.0.0.1", + "sglang_router_policy": "round_robin", + "sglang_router_port": 30000, + "sglang_tokenizer_path": None, + } + values.update(overrides) + return Namespace(**values) + + +def test_minimal_env_config_uses_the_v1_environment_contract(): + config = EnvConfig.model_validate({"taskset": {"id": "harbor"}}) + + environment = Environment(config) + + assert config.is_legacy is False + assert config.env_id == "harbor" + assert type(environment.taskset).__name__ == "HarborTaskset" + + +class _Rendered: + token_ids = [10, 11] + multi_modal_data = None + is_content = [True, True] + + @staticmethod + def message_token_spans(): + return [(0, 2)] + + +class _Renderer: + supports_tools = True + + def render(self, messages, *, tools, add_generation_prompt): + assert messages == [{"role": "user", "content": "question"}] + assert tools is None + assert add_generation_prompt is True + return _Rendered() + + @staticmethod + def get_stop_token_ids(): + return [99] + + @staticmethod + def parse_response(token_ids, *, tools): + assert token_ids == [20, 21] + assert tools is None + return SimpleNamespace(content="answer", reasoning_content=None, tool_calls=[]) + + +@pytest.mark.asyncio +async def test_published_train_client_runs_through_miles_transport(monkeypatch): + async def fake_post(_url, _payload, headers=None): + assert headers is None + return { + "request_id": "request-id", + "meta_info": { + "completion_tokens": 2, + "finish_reason": {"type": "stop"}, + "output_token_logprobs": [[-0.1, 20], [-0.2, 21]], + }, + } + + monkeypatch.setattr("miles.utils.http_utils.post", fake_post) + client = TrainClient(MilesSGLangTransport(_args()), renderer_model_name="test/model") + client._pool = _Renderer() + + response = await client.get_response( + ChatDialect(), + {"messages": [{"role": "user", "content": "question"}]}, + "test/model", + SamplingConfig(temperature=0.2, max_tokens=2), + session_id="trace-id", + ) + + assert response.message.content == "answer" + assert response.tokens.prompt_ids == [10, 11] + assert response.tokens.completion_ids == [20, 21] + assert response.tokens.completion_logprobs == [-0.1, -0.2] + ChatDialect().validate_response(response.raw) + + +@pytest.mark.asyncio +async def test_published_train_client_rejects_non_chat_dialects(): + client = TrainClient(MilesSGLangTransport(_args()), renderer_model_name="test/model") + + with pytest.raises(NotImplementedError, match="chat-completions dialect"): + await client.get_response( + ResponsesDialect(), + {"input": "question"}, + "test/model", + SamplingConfig(max_tokens=2), + ) diff --git a/tests/fast/utils/test_arguments.py b/tests/fast/utils/test_arguments.py index 20a7bfb758..c4f24ab47f 100644 --- a/tests/fast/utils/test_arguments.py +++ b/tests/fast/utils/test_arguments.py @@ -7,6 +7,7 @@ import pytest from miles.utils.arguments import ( + VERIFIERS_ROLLOUT_FUNCTION_PATH, _maybe_apply_dumper_overrides, _resolve_ft_components, get_miles_extra_args_provider, @@ -149,6 +150,74 @@ def test_recompute_logprobs_via_prefill_flag_is_parsed(): assert args.recompute_logprobs_via_prefill is True +def test_verifiers_config_is_parsed(): + parser = argparse.ArgumentParser() + get_miles_extra_args_provider()(parser) + + args = parser.parse_args(["--verifiers-config", "vf.toml"] + REQUIRED_ARGS) + + assert args.verifiers_config == "vf.toml" + assert VERIFIERS_ROLLOUT_FUNCTION_PATH.startswith("miles.rollout.verifiers_rollout.") + + +def _parse_verifiers_args(*extra: str): + parser = argparse.ArgumentParser() + get_miles_extra_args_provider()(parser) + return parser.parse_args( + [ + "--verifiers-config", + "vf.toml", + "--num-rollout", + "1", + *extra, + ] + + REQUIRED_ARGS + ) + + +def test_verifiers_validation_selects_adapter_and_disables_prompt_dataset(): + args = _parse_verifiers_args("--hf-checkpoint", "test/model") + + miles_validate_args(args) + + assert args.rollout_function_path == VERIFIERS_ROLLOUT_FUNCTION_PATH + assert args.rollout_global_dataset is False + + +def test_verifiers_validation_rejects_partial_rollout(): + args = _parse_verifiers_args("--partial-rollout") + + with pytest.raises(ValueError, match="cannot be resumed"): + miles_validate_args(args) + + +@pytest.mark.parametrize( + "flag", + ["--use-opd", "--use-rollout-routing-replay", "--use-rollout-indexer-replay"], +) +def test_verifiers_validation_rejects_unpreserved_token_metadata(flag): + args = _parse_verifiers_args(flag) + + with pytest.raises(ValueError, match="additional token metadata"): + miles_validate_args(args) + + +def test_verifiers_validation_rejects_multimodal_dataset_mapping(): + args = _parse_verifiers_args("--multimodal-keys", '{"image": "image"}') + + with pytest.raises(ValueError, match="text-only renderer inputs"): + miles_validate_args(args) + + +def test_verifiers_validation_rejects_custom_chat_template(tmp_path): + template = tmp_path / "custom.jinja" + template.write_text("{{ messages }}") + args = _parse_verifiers_args("--chat-template-path", str(template)) + + with pytest.raises(ValueError, match="renderers does not accept a custom Jinja template"): + miles_validate_args(args) + + def test_custom_megatron_post_save_hook_path_is_parsed(): parser = argparse.ArgumentParser() get_miles_extra_args_provider()(parser)