Skip to content
Closed
23 changes: 18 additions & 5 deletions docs/user-guide/rollout-endpoints.md
Original file line number Diff line number Diff line change
Expand Up @@ -86,14 +86,14 @@ Helpers:

- `compute_prompt_ids_from_sample` and `compute_request_payload` from
`miles/rollout/generate_utils/generate_endpoint_utils.py` build `/generate` requests.
- For multi-sample outputs, set `--generate-multi-samples` and return a list.
- Returning a `list[Sample]` from a generate function is supported natively; no flag is needed.

### Reference generators

- **`single_turn.py`**: single-turn generation via `/generate`. Text or multimodal prompts.
- **`multi_turn.py`**: multi-turn tool calling via `/generate`. Adds CLI flags
`--generate-max-turns`, `--generate-tool-specs-path`, `--generate-tool-call-parser`,
`--generate-execute-tool-function-path`, `--generate-multi-samples`.
`--generate-execute-tool-function-path`.
- **`benchmarkers.py`**: forces random output sequence length for benchmarking.

---
Expand Down Expand Up @@ -181,6 +181,18 @@ remain a `messages` list. SGLang handles templating server-side.

</Warning>

<Warning>

**Agentic output is a `list[Sample]`.** `agentic_tool_call.generate` always returns a list
(one merged TITO sample per linear run today). Consequences:

- A custom reward model (`--custom-rm-path`) is called in batch form with a
`list[Sample]` argument; it must handle that shape.
- `--group-rm`, `--partial-rollout`, and `--recompute-logprobs-via-prefill` are not
supported in combination with the agentic generator.

</Warning>

### Optional teardown: the `abort` hook

The module named by `--custom-agent-function-path` may expose an optional `abort`
Expand Down Expand Up @@ -217,8 +229,9 @@ is a thin wrapper around the custom agent. It:
1. Creates a session on MilesRouter and builds a session-scoped `base_url`.
2. Calls the custom agent (from `--custom-agent-function-path`) to send one or more
chat requests.
3. Collects session records via `OpenAIEndpointTracer`.
4. Converts records into `Sample` objects via `compute_samples_from_openai_records`.
3. Collects server-assembled `Sample` objects via `OpenAIEndpointTracer.collect_samples`
(the session server converts records into samples, truncates and merges on the
owning instance; records never leave the server).

For broader customization beyond the OpenAI wrapper, see the `/generate` path above.

Expand All @@ -234,7 +247,7 @@ TITO needs two things from every SGLang response:
By default, `build_chat_request_kwargs` sets both flags. The session middleware
forwards raw `messages` to SGLang, which tokenizes the prompt and returns the
response. `_compute_sample_from_openai_record` in
[`openai_endpoint_utils.py`](https://github.com/radixark/miles/blob/main/miles/rollout/generate_utils/openai_endpoint_utils.py)
[`samples.py`](https://github.com/radixark/miles/blob/main/miles/rollout/session/samples.py)
extracts prompt and output ids from the response and concatenates them into
`sample.tokens`. You don't need to provide `input_ids` yourself.

Expand Down
6 changes: 3 additions & 3 deletions examples/experimental/swe-agent-v2/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -289,8 +289,6 @@ merge_samples() ->
logprobs: -------- [real] [0.0] [real]
```

Without TITO, use `--generate-multi-samples` to skip merge and train on per-turn samples instead (current default in `run.sh`).

## Troubleshooting

### Harbor containers can't reach Miles Router
Expand All @@ -312,7 +310,9 @@ The task directory for the given `instance_id` doesn't exist under `HARBOR_TASKS

### `b.tokens must start with a.tokens` assertion error

Multi-turn merge fails due to BPE re-tokenization inconsistency. Use `--generate-multi-samples` (already default in `run.sh`) to skip merge and train on per-turn samples.
Multi-turn merge fails due to BPE re-tokenization inconsistency. The session-server TITO path
(pretokenized `input_ids`) avoids the re-tokenization entirely; check that the run goes through
`--use-session-server` and that the chat template round-trips (see the TITO docs).

### Trace-viewer shows no trajectories

Expand Down
24 changes: 12 additions & 12 deletions miles/rollout/filter_hub/dynamic_sampling_filters.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,25 +6,25 @@
__all__ = ["check_reward_nonzero_std", "check_no_aborted"]


def check_reward_nonzero_std(args, samples: list[Sample], **kwargs):
rewards = [sample.get_reward_value(args) for sample in samples]
keep = torch.tensor(rewards, dtype=torch.float64).std() > 1e-8
return DynamicFilterOutput(
keep=keep,
reason=None if keep else f"zero_std_{round(rewards[0], 1)}",
)


def _flatten_samples(samples):
"""Flatten samples that may contain nested lists (from --generate-multi-samples)."""
def _flatten_samples(samples: list[Sample | list[Sample]]):
"""Flatten a group whose elements are `Sample` or `list[Sample]` (generate-function dependent)."""
for s in samples:
if isinstance(s, list):
yield from s
else:
yield s


def check_no_aborted(args, samples: list[Sample], **kwargs):
def check_reward_nonzero_std(args, samples: list[Sample | list[Sample]], **kwargs):
rewards = [sample.get_reward_value(args) for sample in _flatten_samples(samples)]
keep = torch.tensor(rewards, dtype=torch.float64).std() > 1e-8
return DynamicFilterOutput(
keep=keep,
reason=None if keep else f"zero_std_{round(rewards[0], 1)}",
)


def check_no_aborted(args, samples: list[Sample | list[Sample]], **kwargs):
"""Reject entire group if any sample was aborted (e.g. env timeout, Docker crash)."""
if any(s.status == Sample.Status.ABORTED for s in _flatten_samples(samples)):
return DynamicFilterOutput(keep=False, reason="group_has_aborted")
Expand Down
70 changes: 24 additions & 46 deletions miles/rollout/generate_hub/agentic_tool_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,9 @@
The agent logic is fully encapsulated in a user-provided async function
(--custom-agent-function-path). This generate function only handles:
1. TITO session tracing (OpenAIEndpointTracer)
2. Converting session records to training samples
3. Multi-turn merge
2. Collecting the worker-assembled training samples (the session server
converts records to samples, truncates and merges in the owning worker)
3. Driver-side metadata application (agent_metadata, session_metadata)

Agent function contract:
async def my_agent(
Expand All @@ -32,12 +33,7 @@ async def my_agent(
from sglang.srt.entrypoints.openai.protocol import ChatCompletionRequest

from miles.rollout.base_types import GenerateFnInput, GenerateFnOutput
from miles.rollout.generate_utils.openai_endpoint_utils import (
OpenAIEndpointTracer,
compute_samples_from_openai_records,
truncate_samples_by_total_tokens,
)
from miles.rollout.generate_utils.sample_utils import merge_samples
from miles.rollout.generate_utils.openai_endpoint_utils import OpenAIEndpointTracer
from miles.utils.misc import load_function
from miles.utils.types import Sample

Expand Down Expand Up @@ -84,69 +80,51 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput:
logger.warning(f"{log_prefix} Agent function failed: {e}", exc_info=True)

finally:
logger.debug(f"{log_prefix} Calling collect_records...")
records, session_metadata = await tracer.collect_records()
logger.debug(f"{log_prefix} collect_records done: {len(records)} records")
# The session server assembles the samples on the owning instance; records
# never leave it. Runs even when the agent function failed, like the old
# collect_records; a collect failure (422/5xx/timeout) raises loudly.
logger.debug(f"{log_prefix} Calling collect_samples...")
result = await tracer.collect_samples(input.sample, max_seq_len=max_seq_len)
logger.debug(
f"{log_prefix} collect_samples done: {len(result.samples)} samples, "
f"total_time={time.monotonic()-t_start:.1f}s"
)

if not records:
logger.warning("No model calls recorded for sample")
if not result.samples:
if result.empty_reason == "all_truncated":
logger.warning("All samples truncated (prompt already exceeds max_seq_len)")
else:
logger.warning("No model calls recorded for sample")
sample = deepcopy(input.sample)
sample.status = Sample.Status.ABORTED
return GenerateFnOutput(samples=sample)

logger.debug(f"{log_prefix} Computing samples from {len(records)} records...")
samples = compute_samples_from_openai_records(
input.args,
input.sample,
records,
input.state.tokenizer,
accumulated_token_ids=session_metadata.get("accumulated_token_ids"),
max_trim_tokens=session_metadata.get("max_trim_tokens", 0),
)
return GenerateFnOutput(samples=[sample])

logger.debug(
f"{log_prefix} compute_samples done: {len(samples)} samples, total_time={time.monotonic()-t_start:.1f}s"
)
samples = result.samples
for s in samples:
s.metadata.update(agent_metadata or {})

# If the agent function reports wall-clock time spent outside policy generation
# (env/tool steps), surface it on Sample.non_generation_time so throughput
# accounting subtracts it. Must be equal across all turn-samples: merge_samples
# collapses them with _merge_equal_value, which asserts the values match.
# accounting subtracts it (applied to every returned sample).
ngt = ((agent_metadata or {}).get("agent_metrics") or {}).get("total_tool_time")
if ngt is not None:
for s in samples:
s.non_generation_time = ngt

if max_seq_len is not None:
samples = truncate_samples_by_total_tokens(samples, max_seq_len, input.state.tokenizer)

if not samples:
logger.warning("All samples truncated (prompt already exceeds max_seq_len)")
sample = deepcopy(input.sample)
sample.status = Sample.Status.ABORTED
return GenerateFnOutput(samples=sample)

if not input.args.generate_multi_samples:
samples = merge_samples(samples, input.state.tokenizer)
samples.metadata.update(session_metadata)
else:
samples[-1].metadata.update(session_metadata)
samples[-1].metadata.update(result.session_metadata)
return GenerateFnOutput(samples=samples)


def _add_arguments(parser: argparse.ArgumentParser):
parser.add_argument("--custom-agent-function-path", type=str)
parser.add_argument("--generate-multi-samples", action="store_true", default=False)
parser.add_argument(
"--max-seq-len",
type=int,
default=None,
dest="max_seq_len",
help="Max sequence length in tokens (prompt + completion, including env responses) "
"per session. Truncates samples on the Miles side and is forwarded to the "
"Harbor agent server (as max_seq_len) to abort the trial early.",
"per session. Truncation happens inside the session server during sample assembly; "
"also forwarded to the Harbor agent server (as max_seq_len) to abort the trial early.",
)


Expand Down
13 changes: 1 addition & 12 deletions miles/rollout/generate_hub/multi_turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,8 +36,6 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput:
tool_specs = load_function(args.generate_tool_specs_path)
tool_call_parser = create_tool_call_parser(tool_specs, args.generate_tool_call_parser)

multi_samples = []

# ----------------------- Initial prompts -------------------------

prompt_tokens_ids = compute_prompt_ids_from_sample(input.state, sample, tools=tool_specs)
Expand All @@ -50,19 +48,11 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput:
payload, halt_status = compute_request_payload(args, sample.tokens, input.sampling_params)
if payload is None:
sample.status = halt_status
if args.generate_multi_samples and multi_samples:
multi_samples[-1].status = halt_status
break

if args.generate_multi_samples:
sample = deepcopy(input.sample)

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

if args.generate_multi_samples:
multi_samples.append(deepcopy(sample))

if output["meta_info"]["finish_reason"]["type"] in ("abort", "length"):
break

Expand All @@ -75,15 +65,14 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput:
tool_messages = await execute_tool_calls(tool_calls, execute_tool_function)
update_sample_with_tool_responses(sample, tool_messages, tokenizer=tokenizer)

return GenerateFnOutput(samples=multi_samples if args.generate_multi_samples else sample)
return GenerateFnOutput(samples=sample)


def _add_arguments(parser: argparse.ArgumentParser):
parser.add_argument("--generate-max-turns", type=int, default=16)
parser.add_argument("--generate-tool-specs-path", type=str)
parser.add_argument("--generate-tool-call-parser", type=str)
parser.add_argument("--generate-execute-tool-function-path", type=str)
parser.add_argument("--generate-multi-samples", action="store_true")


generate.add_arguments = _add_arguments
Loading
Loading