Implement PartialRolloutFn abstraction for clean partial rollout (#172)

* more

* more

* more

* more

* more

* fmt

* mv

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* fmt

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* rm

* Revert "rm"

This reverts commit 53184f608f501af4757709cd396ae94516f00dd6.

* cp

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* more

* fmt

* fmt

---------

Co-authored-by: Zilin Zhu <zhuzilinallen@gmail.com>
This commit is contained in:
fzyzcjy
2025-08-13 19:04:41 +08:00
committed by GitHub
co-authored by Zilin Zhu
parent 118b6a02b4
commit 2a95a217c2
6 changed files with 474 additions and 39 deletions
+2 -2
View File
@@ -21,8 +21,8 @@ class RolloutFnCallParams:
@dataclass
class RolloutFnCallOutput:
samples: Optional[list[list[Sample]]]
metrics: Optional[dict[str, Any]] # TODO what is the type
samples: Optional[list[list[Sample]]] = None
metrics: Optional[dict[str, Any]] = None # TODO what is the type
class BaseRolloutFn(Protocol):
@@ -37,14 +37,6 @@ class LegacyAdapterRolloutFn:
)
del self.rollout_id
def __call__(self, params: RolloutFnCallParams) -> RolloutFnCallOutput:
raw_output = self.original_fn(
self.init_params.args,
params.rollout_id,
self.init_params.buffer,
evaluation=self.init_params.evaluation,
)
if self.init_params.evaluation:
return RolloutFnCallOutput(samples=None, metrics=raw_output)
else:
@@ -0,0 +1,91 @@
from functools import partial
from typing import Callable
from slime.rollout.components.base_rollout_fn import RolloutFnInitParams, RolloutFnCallParams, RolloutFnCallOutput
from slime.utils.misc import load_function
from slime.utils.types import Sample
class PartialRolloutFn:
"""A rollout fn to support partial rollout.
The user only needs to provide `generate_one_step` to do arbitrary generation for one step (the
"rollout worker" green box in paper), and this class maintains an aborted_samples_buffer (the
"Replay Buffer" red box in paper) for partial rollout.
For more details, please visit https://arxiv.org/abs/2501.12599
"""
def __init__(
self,
params: RolloutFnInitParams,
generate_one_step: Callable,
):
self.args = params.args
self.data_source = params.data_source
self.generate_one_step = generate_one_step
# a list of sample group.
# each group has n_samples_per_prompt samples, all of them has the same prompt.
self.aborted_samples_buffer: list[list[Sample]] = []
if (p := self.args.buffer_filter_path) is not None:
self.buffer_filter = load_function(p)
else:
self.buffer_filter = _buffer_filter_pop_first
def __call__(self, params: RolloutFnCallParams) -> RolloutFnCallOutput:
output, aborted_samples = self.generate_one_step(
params=params,
get_samples=partial(self._get_samples, rollout_id=params.rollout_id),
)
self._add_samples_to_buffer(aborted_samples)
return output
# TODO simplify
def _get_samples(self, num_samples: int, rollout_id: int) -> list[list[Sample]]:
"""
Return num_samples samples
"""
samples = self._get_samples_from_buffer(num_samples, rollout_id=rollout_id)
num_samples -= len(samples)
if num_samples == 0:
return samples
samples += self._get_samples_from_data_source(num_samples=num_samples)
return samples
def _get_samples_from_buffer(self, num_samples: int, rollout_id: int) -> list[list[Sample]]:
if len(self.aborted_samples_buffer) == 0 or num_samples == 0:
return []
samples = self.buffer_filter(self.args, rollout_id, self.aborted_samples_buffer, num_samples)
return samples
def _get_samples_from_data_source(self, num_samples: int) -> list[list[Sample]]:
return self.data_source.get_samples(num_samples=num_samples)
def _add_samples_to_buffer(self, samples: list[list[Sample]]):
if not samples:
return
# TODO improve code, e.g. separate assertion and addition
assert isinstance(samples, list), f"samples must be a list, got {type(samples)}"
assert isinstance(samples[0], list), f"the elements of samples must be list, got {type(samples[0])}"
for i in range(0, len(samples)):
assert (
len(samples[i]) == self.args.n_samples_per_prompt
), f"the length of the elements of samples must be equal to n_samples_per_prompt, got {len(samples[i])} != {self.args.n_samples_per_prompt}"
group = samples[i] # type: ignore
self.aborted_samples_buffer.append(group)
def _buffer_filter_pop_first(
args, rollout_id, aborted_samples_buffer: list[list[Sample]], num_samples: int
) -> list[list[Sample]]:
num_to_pop = min(len(aborted_samples_buffer), num_samples)
samples = aborted_samples_buffer[:num_to_pop]
del aborted_samples_buffer[:num_to_pop]
return samples
+23 -28
View File
@@ -1,5 +1,6 @@
import asyncio
import copy
from functools import partial
from tqdm import tqdm
from transformers import AutoTokenizer
@@ -10,10 +11,12 @@ from slime.utils.http_utils import get, post
from slime.utils.misc import SingletonMeta, load_function
from slime.utils.types import Sample
from slime.rollout.components.sample_generator import generate_one_sample_vanilla
from .components.base_rollout_fn import RolloutFnCallParams, RolloutFnInitParams, RolloutFnCallOutput
from .components.partial_rollout_fn import PartialRolloutFn
from .rm_hub import async_rm, batched_async_rm
__all__ = ["generate_rollout"]
__all__ = ["create_rollout_fn"]
class GenerateState(metaclass=SingletonMeta):
@@ -154,13 +157,13 @@ async def abort(args, rollout_id: int):
return aborted_samples
async def generate_rollout_async(args, rollout_id: int, data_source) -> list[list[Sample]]:
async def generate_rollout_async(args, rollout_id: int, get_samples):
"""An example to implement the generate_rollout function for an rule based rm rollout generation.
Args:
args: the whole args
rollout_id: int, the id of the rollout, used for deterministic data generation
data_source: the data source to fetch
get_samples: the data source to fetch
Returns:
list[list[Sample]]: a list of samples generated by the rollout, the length of the list is exactly the same as the `rollout_batch_size`
@@ -186,7 +189,7 @@ async def generate_rollout_async(args, rollout_id: int, data_source) -> list[lis
while len(data) < target_data_size:
while state.remaining_batch_size < target_data_size:
# get samples from the buffer and submit the generation requests.
samples = data_source(args.over_sampling_batch_size)
samples = get_samples(args.over_sampling_batch_size)
state.submit_generate_tasks(samples)
# wait for the generation to finish
@@ -229,7 +232,7 @@ async def generate_rollout_async(args, rollout_id: int, data_source) -> list[lis
# reset the global state to prevent effects on the next rollout or eval.
state.reset()
return data, aborted_samples
return RolloutFnCallOutput(samples=data), aborted_samples
EVAL_PROMPT_DATASET = {}
@@ -241,7 +244,7 @@ async def eval_rollout(args, rollout_id):
for i in range(0, len(args.eval_prompt_data), 2):
name, path = args.eval_prompt_data[i : i + 2]
results.update(await eval_rollout_single_dataset(args, rollout_id, name, path))
return results, []
return RolloutFnCallOutput(metrics=results), []
async def eval_rollout_single_dataset(args, rollout_id, name, path):
@@ -326,28 +329,20 @@ async def eval_rollout_single_dataset(args, rollout_id, name, path):
}
# TODO remove this temp function
def generate_rollout(args, rollout_id, data_buffer, evaluation=False):
"""An example to implement the generate_rollout function for an rule based rm rollout generation.
def _generate_one_step(
init_params: RolloutFnInitParams,
params: RolloutFnCallParams,
get_samples,
):
if init_params.evaluation:
return run(eval_rollout(init_params.args, params.rollout_id))
else:
return run(generate_rollout_async(init_params.args, params.rollout_id, get_samples))
Args:
args: the whole args
rollout_id: int, the id of the rollout, used for deterministic data generation
data_buffer: the data buffer to store the generated samples
evaluation: bool, whether the rollout is for evaluation or not
Returns:
list[list[Sample]]: a list of list of samples generated by the rollout
"""
completed_samples, aborted_samples = generate_abortable_samples(
args, rollout_id, data_buffer.get_samples, evaluation=evaluation
def create_rollout_fn(params: RolloutFnInitParams):
assert params.args.rollout_global_dataset
return PartialRolloutFn(
params=params,
generate_one_step=partial(_generate_one_step, init_params=params),
)
data_buffer.add_samples(aborted_samples)
return completed_samples
def generate_abortable_samples(args, rollout_id, data_source, evaluation=False):
assert args.rollout_global_dataset
if evaluation:
return run(eval_rollout(args, rollout_id))
return run(generate_rollout_async(args, rollout_id, data_source))
+357
View File
@@ -0,0 +1,357 @@
"""
This file demonstrates the legacy rollout fn API usage. Prefer to use the new APIs.
"""
import asyncio
import copy
from tqdm import tqdm
from transformers import AutoTokenizer
from slime.utils.async_utils import run
from slime.utils.data import Dataset
from slime.utils.http_utils import get, post
from slime.utils.misc import SingletonMeta, load_function
from slime.utils.types import Sample
from slime.rollout.components.sample_generator import generate_one_sample_vanilla
from .rm_hub import async_rm, batched_async_rm
__all__ = ["generate_rollout"]
class GenerateState(metaclass=SingletonMeta):
"""
The global state for the generation process.
"""
def __init__(self, args):
# persistant state for the generation process
self.args = args
self.tokenizer = AutoTokenizer.from_pretrained(args.hf_checkpoint, trust_remote_code=True)
self.semaphore = asyncio.Semaphore(
args.sglang_server_concurrency * args.rollout_num_gpus // args.rollout_num_gpus_per_engine
)
self.sampling_params = dict(
temperature=args.rollout_temperature,
top_p=args.rollout_top_p,
top_k=args.rollout_top_k,
max_new_tokens=args.rollout_max_response_len,
stop=args.rollout_stop,
stop_token_ids=args.rollout_stop_token_ids,
skip_special_tokens=args.rollout_skip_special_tokens,
no_stop_trim=True,
spaces_between_special_tokens=False,
)
self.reset()
def reset(self):
self.remaining_batch_size = 0
self.pendings = set()
self.aborted = False
def submit_generate_tasks(self, samples: list[list[Sample]]):
for group in samples:
self.pendings.add(
asyncio.create_task(
# submit a group of samples as a single task.
generate_and_rm_group(
self.args,
group,
sampling_params=self.sampling_params.copy(),
evaluation=False,
)
)
)
self.remaining_batch_size += len(samples)
async def generate_and_rm(args, sample: Sample, sampling_params: dict, evaluation=False) -> Sample:
# For samples with existing response, check if they're complete
if sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED:
assert sample.response is not None
if not args.group_rm:
assert sample.reward is not None
return sample
state = GenerateState(args)
# generate
async with state.semaphore:
if state.aborted:
sample.status = Sample.Status.ABORTED
return sample
if args.custom_generate_function_path is not None:
custom_generate_func = load_function(args.custom_generate_function_path)
sample = await custom_generate_func(args, sample, sampling_params)
else:
sample = await generate_one_sample_vanilla(args, state.tokenizer, sample, sampling_params)
if sample.status == Sample.Status.ABORTED:
return sample
# for the rm that need the whole group, we will not do the rm here
if args.group_rm:
return sample
sample.reward = await async_rm(args, sample)
return sample
async def generate_and_rm_group(args, group: list[Sample], sampling_params: dict, evaluation=False) -> list[Sample]:
state = GenerateState(args)
if state.aborted:
return group
group = await asyncio.gather(
*[generate_and_rm(args, sample, sampling_params.copy(), evaluation=evaluation) for sample in group]
)
# for the rm that need the whole group, we will not do the rm here
if not state.aborted and args.group_rm:
rewards = await batched_async_rm(args, group)
for sample, reward in zip(group, rewards):
sample.reward = reward
return group
async def abort(args, rollout_id: int):
aborted_samples = []
state = GenerateState(args)
assert not state.aborted
state.aborted = True
response = await get(
f"http://{args.sglang_router_ip}:{args.sglang_router_port}/list_workers", use_http2=args.use_http2
)
# abort all the requests
for url in response["urls"]:
print(f"Abort request for {url}", flush=True)
await post(f"{url}/abort_request", {"abort_all": True}, use_http2=False)
# make sure all the pending tasks are finished
count = 0
while state.pendings:
done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED)
if not args.partial_rollout:
continue
# for partial rollout, collect the partial samples into the data buffer
for task in done:
group = task.result()
for sample in group:
if sample.response and "start_rollout_id" not in sample.metadata:
sample.metadata["start_rollout_id"] = rollout_id
aborted_samples.append(group)
count += len(group)
if args.partial_rollout:
print(f"Collected {count} partial samples into the data buffer", flush=True)
return aborted_samples
async def generate_rollout_async(args, rollout_id: int, data_source) -> list[list[Sample]]:
"""An example to implement the generate_rollout function for an rule based rm rollout generation.
Args:
args: the whole args
rollout_id: int, the id of the rollout, used for deterministic data generation
data_source: the data source to fetch
Returns:
list[list[Sample]]: a list of samples generated by the rollout, the length of the list is exactly the same as the `rollout_batch_size`
"""
assert args.rollout_global_dataset
state = GenerateState(args)
# instantiate data filters
dynamic_filter = (
load_function(args.dynamic_sampling_filter_path) if args.dynamic_sampling_filter_path is not None else None
)
over_sampling_filter = (
load_function(args.over_sampling_filter_path) if args.over_sampling_filter_path is not None else None
)
# target_data_size is the total number of valid samples to get
target_data_size = args.over_sampling_batch_size if over_sampling_filter is not None else args.rollout_batch_size
data = []
do_print = True
pbar = tqdm(total=target_data_size * args.n_samples_per_prompt, desc="Rollout generation")
while len(data) < target_data_size:
while state.remaining_batch_size < target_data_size:
# get samples from the buffer and submit the generation requests.
samples = data_source(args.over_sampling_batch_size)
state.submit_generate_tasks(samples)
# wait for the generation to finish
done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED)
for task in done:
group: list[Sample] = task.result()
if do_print:
print(
f"First rollout sample: {[group[0].prompt + group[0].response]}, label: {group[0].label}, reward: {group[0].reward}",
flush=True,
)
do_print = False
assert len(group) == args.n_samples_per_prompt
if dynamic_filter is not None and not dynamic_filter(args, group):
state.remaining_batch_size -= 1
continue
# add the samples to the data
# NOTE: here we have not stored all the unused samples back to the data buffer.
if len(data) < target_data_size:
data.append(group)
pbar.update(args.n_samples_per_prompt)
pbar.close()
print(
f"Finish rollout: {[data[-1][0].prompt + data[-1][0].response]}, label: {data[-1][0].label}, reward: {data[-1][0].reward}",
flush=True,
)
# there are still some unfinished requests, abort them
aborted_samples = await abort(args, rollout_id)
if over_sampling_filter is not None:
data = over_sampling_filter(args, data)[: args.rollout_batch_size]
assert len(data) == args.rollout_batch_size, f"Got {len(data)} samples, expected {args.rollout_batch_size}"
data = sorted(data, key=lambda group: group[0].index)
# reset the global state to prevent effects on the next rollout or eval.
state.reset()
return data, aborted_samples
EVAL_PROMPT_DATASET = {}
async def eval_rollout(args, rollout_id):
assert not args.group_rm, "Group RM is not supported for eval rollout"
results = {}
for i in range(0, len(args.eval_prompt_data), 2):
name, path = args.eval_prompt_data[i : i + 2]
results.update(await eval_rollout_single_dataset(args, rollout_id, name, path))
return results, []
async def eval_rollout_single_dataset(args, rollout_id, name, path):
"""An example to implement the eval_rollout function for an rule based rm rollout generation.
Args:
args: the whole args
rollout_id: int, the id of the rollout, used for deterministic data generation
name: str, the name of the dataset
path: str, the path of the dataset
"""
assert not args.group_rm, "Group RM is not supported for eval rollout"
global EVAL_PROMPT_DATASET
if name not in EVAL_PROMPT_DATASET:
tokenizer = AutoTokenizer.from_pretrained(args.hf_checkpoint, trust_remote_code=True)
EVAL_PROMPT_DATASET[name] = Dataset(
path,
tokenizer=tokenizer,
max_length=args.rollout_max_prompt_len,
prompt_key=args.input_key if args.eval_input_key is None else args.eval_input_key,
label_key=args.label_key if args.eval_label_key is None else args.eval_label_key,
metadata_key=args.metadata_key,
tool_key=args.tool_key if args.eval_tool_key is None else args.eval_tool_key,
apply_chat_template=args.apply_chat_template,
)
dataset = EVAL_PROMPT_DATASET[name]
sampling_params = dict(
temperature=args.rollout_temperature if args.eval_temperature is None else args.eval_temperature,
top_p=args.rollout_top_p if args.eval_top_p is None else args.eval_top_p,
top_k=args.rollout_top_k if args.eval_top_k is None else args.eval_top_k,
max_new_tokens=(
args.rollout_max_response_len if args.eval_max_response_len is None else args.eval_max_response_len
),
stop=args.rollout_stop,
stop_token_ids=args.rollout_stop_token_ids,
skip_special_tokens=args.rollout_skip_special_tokens,
no_stop_trim=True,
spaces_between_special_tokens=False,
)
tasks = []
# do multiple samples for eval prompts
sample_index = 0
for i, prompt_sample in enumerate(dataset.samples):
for j in range(args.n_samples_per_eval_prompt):
# use the same prompt for multiple samples
sample = copy.deepcopy(prompt_sample)
sample.index = sample_index
sample_index += 1
tasks.append(
generate_and_rm(
args,
sample,
sampling_params=sampling_params,
evaluation=True,
)
)
data = []
do_print = True
pbar = tqdm(total=len(tasks), desc="Rollout generation", disable=not do_print)
for coro in asyncio.as_completed(tasks):
sample = await coro
if do_print:
print([sample.prompt + sample.response], sample.reward, flush=True)
do_print = False
data.append(sample)
pbar.update(1)
pbar.close()
data.sort(key=lambda sample: sample.index)
reward_key = args.reward_key or args.eval_reward_key
return {
name: {
"rewards": [sample.reward if not reward_key else sample.reward[reward_key] for sample in data],
"truncated": [sample.status == Sample.Status.TRUNCATED for sample in data],
}
}
# TODO remove this temp function
def generate_rollout(args, rollout_id, data_buffer, evaluation=False):
"""An example to implement the generate_rollout function for an rule based rm rollout generation.
Args:
args: the whole args
rollout_id: int, the id of the rollout, used for deterministic data generation
data_buffer: the data buffer to store the generated samples
evaluation: bool, whether the rollout is for evaluation or not
Returns:
list[list[Sample]]: a list of list of samples generated by the rollout
"""
completed_samples, aborted_samples = generate_abortable_samples(
args, rollout_id, data_buffer.get_samples, evaluation=evaluation
)
data_buffer.add_samples(aborted_samples)
return completed_samples
def generate_abortable_samples(args, rollout_id, data_source, evaluation=False):
assert args.rollout_global_dataset
if evaluation:
return run(eval_rollout(args, rollout_id))
return run(generate_rollout_async(args, rollout_id, data_source))
+1 -1
View File
@@ -101,7 +101,7 @@ def get_slime_extra_args_provider(add_custom_arguments=None):
parser.add_argument(
"--rollout-function-path",
type=str,
default="slime.rollout.sglang_rollout.generate_rollout",
default="slime.rollout.sglang_rollout.create_rollout_fn",
help=(
"Path to the rollout generation function."
"You should use this model to create your own custom rollout function, "