mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
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:
@@ -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
|
||||
@@ -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))
|
||||
|
||||
@@ -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))
|
||||
@@ -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, "
|
||||
|
||||
Reference in New Issue
Block a user