Merge remote-tracking branch 'upstream/main'

This commit is contained in:
Yusheng Su
2025-07-01 07:17:31 +00:00
20 changed files with 413 additions and 336 deletions
+2 -1
View File
@@ -178,4 +178,5 @@ outputs/
tests/
local/
**/rollout_data/
**/buffer_stats/
**/buffer_stats/
*.out
+1 -4
View File
@@ -351,10 +351,7 @@ In Rollout Buffer, you need to implement the key function `run_rollout()` for yo
### Filtering and Reward Configuration
- **`--buffer-filter-path`**: Data buffer filter path, supports customization, defaults to using the latest batch of data in Buffer for updates
```bash
--buffer-filter-path slime.rollout.filter_hub.buffer_filters.pop_first
```
- **`--buffer-filter-path`**: Data buffer filter path, supports customization, defaults to using the latest batch of data in Buffer for updates.
- **`--disable-rewards-normalization`**: Disable reward normalization, if rewards are already normalized in Rollout Buffer, please enable this parameter
```bash
+2 -2
View File
@@ -37,9 +37,9 @@
6. **How is the batch size calculated?**
A single rollout uses `rollout_batch_size` prompts. For each prompt, `n_samples_per_prompts` samples are generated. Therefore, one rollout contains a total of `rollout_batch_size * n_samples_per_prompts` data entries.
A single rollout uses `rollout_batch_size` prompts. For each prompt, `n_samples_per_prompt` samples are generated. Therefore, one rollout contains a total of `rollout_batch_size * n_samples_per_prompt` data entries.
You can use `--num-steps-per-rollout` to determine how many steps to run per rollout. This is equivalent to setting the `global_batch_size` to `rollout_batch_size * n_samples_per_prompts // num_steps_per_rollout`.
You can use `--num-steps-per-rollout` to determine how many steps to run per rollout. This is equivalent to setting the `global_batch_size` to `rollout_batch_size * n_samples_per_prompt // num_steps_per_rollout`.
7. **Does slime perform data packing / variable-length (varlen) processing?**
+1 -4
View File
@@ -352,10 +352,7 @@ class CustomTaskLossMaskGenerator(MultiTurnLossMaskGenerator):
#### 过滤与奖励配置
- **`--buffer-filter-path`**:数据缓冲区过滤器路径,支持自定义,默认采用 Buffer 中最新的一批数据进行更新
```bash
--buffer-filter-path slime.rollout.filter_hub.buffer_filters.pop_first
```
- **`--buffer-filter-path`**:数据缓冲区过滤器路径,支持自定义,默认采用 Buffer 中最新的一批数据进行更新。
- **`--disable-rewards-normalization`**:禁用奖励归一化,如果 Rollout Buffer 中已经归一化奖励,请启用此参数
```bash
+2 -2
View File
@@ -37,9 +37,9 @@
1. **batch size 是如何计算的?**
一个 rollout 会用 `rollout_batch_size` 条 prompt,每一条会采 `n_samples_per_prompts` 条,所以一个 rollout 共 `rollout_batch_size * n_samples_per_prompts` 条数据。
一个 rollout 会用 `rollout_batch_size` 条 prompt,每一条会采 `n_samples_per_prompt` 条,所以一个 rollout 共 `rollout_batch_size * n_samples_per_prompt` 条数据。
可以用 `--num-steps-per-rollout` 来决定每一个 rollout 跑多少步。这相当于是把 `global_batch_size` 设置成 `rollout_batch_size * n_samples_per_prompts // num_steps_per_rollout`。
可以用 `--num-steps-per-rollout` 来决定每一个 rollout 跑多少步。这相当于是把 `global_batch_size` 设置成 `rollout_batch_size * n_samples_per_prompt // num_steps_per_rollout`。
1. **slime 是否进行了 data packing / varlen 处理?**
-1
View File
@@ -140,7 +140,6 @@ ray job submit --address="http://127.0.0.1:8265" \
--agent-rollout-buffer-url http://${MASTER_ADDR}:8889 \
--keep-old-actor \
--update-rollout-weights-interval 1 \
--buffer-filter-path slime.rollout.filter_hub.buffer_filters.pop_first \
--disable-rewards-normalization \
--offload-rollout \
--offload-ref \
-2
View File
@@ -33,9 +33,7 @@ ROLLOUT_ARGS=(
--label-key label
--apply-chat-template
--rollout-shuffle
--rm-type deepscaler
--num-rollout 3000
--rollout-batch-size 32
--n-samples-per-prompt 8
+1
View File
@@ -17,6 +17,7 @@ from megatron.core.utils import get_model_config
from megatron.training.global_vars import get_args
from megatron.training.training import get_model
from .checkpoint import load_checkpoint, save_checkpoint
from .data import get_batch, set_local_storage
from .loss import get_log_probs_and_entropy, policy_loss_func
+88 -58
View File
@@ -6,9 +6,9 @@ from typing import Any, Union
import ray
import torch
import wandb
from transformers import AutoTokenizer
import wandb
from slime.utils.data import JsonlDataset
from slime.utils.misc import load_function
from slime.utils.types import Sample
@@ -17,30 +17,11 @@ logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)
def convert_samples_to_train_data(samples: list[Sample]):
"""
Convert inference generated samples to training data.
"""
samples = sorted(samples, key=lambda x: x.index)
# print([item.index for item in samples][:32])
train_data = {
"tokens": [sample.tokens for sample in samples],
"response_lengths": [sample.response_length for sample in samples],
"rewards": [sample.reward for sample in samples],
"truncated": [1 if sample.truncated else 0 for sample in samples],
}
if samples[0].loss_mask:
train_data["loss_masks"] = []
for sample in samples:
assert (
len(sample.loss_mask) == sample.response_length
), f"loss mask length {len(sample.loss_mask)} != response length {sample.response_length}"
train_data["loss_masks"].append(sample.loss_mask)
if samples[0].metadata and "raw_reward" in samples[0].metadata:
train_data["raw_reward"] = [sample.metadata["raw_reward"] for sample in samples]
if samples[0].metadata and "round_number" in samples[0].metadata:
train_data["round_number"] = [sample.metadata["round_number"] for sample in samples]
return train_data
def pop_first(args, rollout_id, buffer: list[list[Sample]], num_samples: int) -> list[list[Sample]]:
num_to_pop = min(len(buffer), num_samples)
samples = buffer[:num_to_pop]
del buffer[:num_to_pop]
return samples
@ray.remote
@@ -48,8 +29,13 @@ class Buffer:
def __init__(self, args):
self.args = args
self.buffer = []
self.buffer_filter = load_function(self.args.buffer_filter_path)
# a list of sample group.
# each group has n_samples_per_prompt samples, all of them has the same prompt.
self.buffer: list[list[Sample]] = []
if self.args.buffer_filter_path is None:
self.buffer_filter = pop_first
else:
self.buffer_filter = load_function(self.args.buffer_filter_path)
self.train_data_pool = {}
self.eval_data_pool = {}
@@ -123,71 +109,80 @@ class Buffer:
wandb.init(**wandb_config, settings=wandb.Settings(mode="shared"))
async def get_samples(self, num_samples) -> list[Sample]:
async def get_samples(self, num_samples: int) -> list[list[Sample]]:
"""
Return num_samples samples
"""
samples = await self._get_samples_from_buffer(num_samples)
num_samples -= len(samples)
assert num_samples % self.args.n_samples_per_prompt == 0
num_prompts = num_samples // self.args.n_samples_per_prompt
if num_samples == 0:
return samples
if self.dataset is not None:
if self.sample_offset + num_prompts <= len(self.dataset):
prompt_samples = self.dataset.samples[self.sample_offset : self.sample_offset + num_prompts]
self.sample_offset += num_prompts
if self.sample_offset + num_samples <= len(self.dataset):
prompt_samples = self.dataset.samples[self.sample_offset : self.sample_offset + num_samples]
self.sample_offset += num_samples
else:
prompt_samples = self.dataset.samples[self.sample_offset :]
num_prompts -= len(prompt_samples)
num_samples -= len(prompt_samples)
self.epoch_id += 1
if self.args.rollout_shuffle:
self.dataset.shuffle(self.epoch_id)
prompt_samples += self.dataset.samples[:num_prompts]
self.sample_offset = num_prompts
prompt_samples += self.dataset.samples[:num_samples]
self.sample_offset = num_samples
for prompt_sample in prompt_samples:
group = []
for _ in range(self.args.n_samples_per_prompt):
sample = copy.deepcopy(prompt_sample)
sample.index = self.sample_index
self.sample_index += 1
samples.append(sample)
group.append(sample)
samples.append(group)
else:
for _ in range(num_samples):
sample = Sample(
index=self.sample_index,
)
self.sample_index += 1
samples.append(sample)
assert len(samples) == num_samples
group = []
for _ in range(self.args.n_samples_per_prompt):
sample = Sample(
index=self.sample_index,
)
self.sample_index += 1
group.append(sample)
samples.append(group)
return samples
async def _get_samples_from_buffer(self, num_samples) -> list[Sample]:
async def _get_samples_from_buffer(self, num_samples: int) -> list[list[Sample]]:
if len(self.buffer) == 0 or num_samples == 0:
return []
samples = self.buffer_filter(self.buffer, num_samples)
samples = self.buffer_filter(self.args, self.rollout_id, self.buffer, num_samples)
return samples
async def add_samples(self, samples: list[Sample]):
# TODO: we can save some partial rollout data here.
async def add_samples(self, samples: list[list[Sample]]):
"""
Add a sample group to buffer.
"""
if not samples:
return
assert len(samples) % self.args.n_samples_per_prompt == 0
self.buffer.extend(samples)
for i in range(0, len(samples), self.args.n_samples_per_prompt):
group = samples[i : i + self.args.n_samples_per_prompt]
self.buffer.append(group)
def generate(self, rollout_id, evaluation=False):
self.rollout_id = rollout_id
if not evaluation and self.args.load_debug_rollout_data:
data = pickle.load(
open(self.args.load_debug_rollout_data.format(rollout_id=rollout_id), "rb"),
)
data = [Sample(**sample) for sample in data]
self.train_data_pool[rollout_id] = convert_samples_to_train_data(data)
return
else:
generate_rollout = self.eval_generate_rollout if evaluation else self.generate_rollout
data = generate_rollout(self.args, rollout_id, self, evaluation=evaluation)
generate_rollout = self.eval_generate_rollout if evaluation else self.generate_rollout
data = generate_rollout(self.args, rollout_id, self, evaluation=evaluation)
self.set_data(rollout_id, data, evaluation=evaluation)
self._set_data(data, evaluation=evaluation)
def get_data(self, rollout_id, evaluation=False):
data_pool = self.train_data_pool if not evaluation else self.eval_data_pool
@@ -196,16 +191,51 @@ class Buffer:
del data_pool[rollout_id]
return data
def set_data(self, rollout_id, data: Union[list[Sample], Any], evaluation=False):
def _convert_samples_to_train_data(self, samples: list[Sample]):
"""
Convert inference generated samples to training data.
"""
samples = sorted(samples, key=lambda x: x.index)
train_data = {
"tokens": [sample.tokens for sample in samples],
"response_lengths": [sample.response_length for sample in samples],
# some reward model, e.g. remote rm, may return multiple rewards,
# we could use key to select the reward.
"rewards": [
sample.reward if not self.args.reward_key else sample.rewards[self.args.reward_key]
for sample in samples
],
"truncated": [1 if sample.status == Sample.Status.TRUNCATED else 0 for sample in samples],
}
if samples[0].loss_mask:
train_data["loss_masks"] = []
for sample in samples:
assert (
len(sample.loss_mask) == sample.response_length
), f"loss mask length {len(sample.loss_mask)} != response length {sample.response_length}"
train_data["loss_masks"].append(sample.loss_mask)
# overwriting the raw reward
if samples[0].metadata and "raw_reward" in samples[0].metadata:
train_data["raw_reward"] = [sample.metadata["raw_reward"] for sample in samples]
# For rollout buffer
if samples[0].metadata and "round_number" in samples[0].metadata:
train_data["round_number"] = [sample.metadata["round_number"] for sample in samples]
return train_data
def _set_data(self, data: Union[list[Sample], Any], evaluation=False):
data_pool = self.eval_data_pool if evaluation else self.train_data_pool
if not evaluation:
if self.args.save_debug_rollout_data:
pickle.dump(
[sample.__dict__ for sample in data],
open(self.args.save_debug_rollout_data.format(rollout_id=rollout_id), "wb"),
open(self.args.save_debug_rollout_data.format(rollout_id=self.rollout_id), "wb"),
)
data = convert_samples_to_train_data(data)
data_pool[rollout_id] = data
data = self._convert_samples_to_train_data(data)
data_pool[self.rollout_id] = data
def update_metadata(self, metadata: dict):
self.metadata.update(metadata)
+17 -2
View File
@@ -1,3 +1,4 @@
import socket
import ray
from ray.util.placement_group import placement_group
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
@@ -13,9 +14,23 @@ class InfoActor:
def sort_key(x):
index, node_ip, gpu_id = x
index, node_identifier, gpu_id = x
# Sort by node IP number and then by GPU ID
node_ip_parts = list(map(int, node_ip.split(".")))
try:
# try to parse it as an IP address.
ip_address = node_identifier
node_ip_parts = list(map(int, ip_address.split(".")))
except ValueError:
# Try to resolve the hostname to an IP address.
try:
ip_address = socket.gethostbyname(node_identifier)
node_ip_parts = list(map(int, ip_address.split(".")))
except (socket.gaierror, TypeError):
# Instead, we convert each character of the original identifier string
# to its ASCII value. This provides a stable and consistent numerical
# representation that allows for sorting.
node_ip_parts = [ord(c) for c in node_identifier]
return (node_ip_parts, gpu_id)
+5 -6
View File
@@ -4,15 +4,14 @@ from typing import Any, Dict, List, Optional
import aiohttp
import requests
import wandb
from transformers import AutoTokenizer
import wandb
from slime.ray.buffer import Buffer
from slime.utils.async_utils import run
from slime.utils.mask_utils import MultiTurnLossMaskGenerator
from slime.utils.types import Sample
__all__ = ["generate_agent_rollout"]
@@ -235,12 +234,12 @@ async def generate_agent_rollout(
print(f"start rollout id: {rollout_id}")
START_ROLLOUT = False
data_number_to_fetch = args.rollout_batch_size * args.n_samples_per_prompt - data_buffer.get_buffer_length()
data_number_to_fetch = (args.rollout_batch_size - data_buffer.get_buffer_length()) * args.n_samples_per_prompt
if data_number_to_fetch <= 0:
print(
f"❕buffer length: {data_buffer.get_buffer_length()}, buffer has enough data, return {args.rollout_batch_size * args.n_samples_per_prompt} samples"
f"❕buffer length: {data_buffer.get_buffer_length()}, buffer has enough data, return {args.rollout_batch_size} prompts"
)
return await data_buffer.get_samples(args.rollout_batch_size * args.n_samples_per_prompt)
return await data_buffer.get_samples(args.rollout_batch_size)
assert (
data_number_to_fetch % args.n_samples_per_prompt == 0
), "data_number_to_fetch must be a multiple of n_samples_per_prompt"
@@ -307,7 +306,7 @@ async def generate_agent_rollout(
final_return_results = []
await data_buffer.add_samples(sample_results)
final_return_results = await data_buffer.get_samples(args.rollout_batch_size * args.n_samples_per_prompt)
final_return_results = await data_buffer.get_samples(args.rollout_batch_size)
return final_return_results
@@ -1,10 +0,0 @@
def pop_first(buffer, num_samples):
samples = []
for _ in range(num_samples):
if buffer:
samples.append(buffer.pop(0))
return samples
def get_newest_samples(buffer, num_samples):
return buffer[-num_samples:]
@@ -1,18 +1,19 @@
import torch
from slime.utils.types import Samples
from slime.utils.types import Sample
__all__ = ["sort_by_reward_std"]
def sort_by_reward_std(args, samples: list[Samples], **kwargs):
def sort_by_reward_std(args, samples: list[Sample], **kwargs):
args.n_samples_per_prompt
samples_with_std = []
for i in range(0, len(samples), args.n_samples_per_prompt):
batch = samples[i : i + args.n_samples_per_prompt]
rewards = [item[3] for item in batch]
rewards = [item.reward for item in batch]
std = torch.tensor(rewards, dtype=torch.float).std()
for j in range(args.n_samples_per_prompt):
samples_with_std.append(batch[i + j], torch.tensor(rewards, std))
for sample in batch:
samples_with_std.append((sample, std))
# python sort is stable, so the order of samples with the same std is preserved
samples_with_std.sort(key=lambda x: x[1], reverse=True)
return [item[0] for item in samples_with_std]
+187 -175
View File
@@ -1,6 +1,5 @@
import asyncio
import copy
from dataclasses import dataclass
from tqdm import tqdm
from transformers import AutoTokenizer
@@ -8,7 +7,7 @@ from transformers import AutoTokenizer
from slime.utils.async_utils import run
from slime.utils.data import JsonlDataset
from slime.utils.http_utils import get, post
from slime.utils.misc import load_function
from slime.utils.misc import SingletonMeta, load_function
from slime.utils.types import Sample
from .rm_hub import async_rm, batched_async_rm
@@ -16,62 +15,112 @@ from .rm_hub import async_rm, batched_async_rm
__all__ = ["generate_rollout"]
@dataclass
class GenerateState:
remaining_batch_size: int = 0
pendings: set = None
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()
TOKENIZER = None
SEMAPHORE = None
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(args, sample: Sample, sampling_params) -> Sample:
global TOKENIZER, SEMAPHORE
if TOKENIZER is None:
TOKENIZER = AutoTokenizer.from_pretrained(args.hf_checkpoint, trust_remote_code=True)
if SEMAPHORE is None:
SEMAPHORE = asyncio.Semaphore(
args.sglang_server_concurrency * args.rollout_num_gpus // args.rollout_num_gpus_per_engine
)
state = GenerateState(args)
url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate"
assert (
sample.status == Sample.Status.PENDING or sample.status == Sample.Status.ABORTED
), f"Sample status is {sample.status}"
# Handle partial rollout samples: continue generation from existing response
input_text = sample.prompt + sample.response
payload = {
"text": sample.prompt,
"text": input_text,
"sampling_params": sampling_params,
}
while True:
try:
async with SEMAPHORE:
output = await post(url, payload, use_http2=args.use_http2)
except Exception as e:
print(f"Error: {e}, retrying...")
await asyncio.sleep(1)
continue
break
prompt_tokens_ids = TOKENIZER(sample.prompt, add_special_tokens=False)["input_ids"]
response_token_ids = TOKENIZER(output["text"], add_special_tokens=False)["input_ids"]
output = await post(url, payload, use_http2=args.use_http2)
sample.response += output["text"]
if output["meta_info"]["finish_reason"]["type"] == "abort":
sample.status = Sample.Status.ABORTED
return sample
prompt_tokens_ids = state.tokenizer(sample.prompt, add_special_tokens=False)["input_ids"]
response_token_ids = state.tokenizer(sample.response, add_special_tokens=False)["input_ids"]
sample.tokens = prompt_tokens_ids + response_token_ids
sample.response_length = len(response_token_ids)
sample.truncated = output["meta_info"]["finish_reason"]["type"] == "length"
sample.response = output["text"]
sample.aborted = output["meta_info"]["finish_reason"]["type"] == "abort"
match output["meta_info"]["finish_reason"]["type"]:
case "length":
sample.status = Sample.Status.TRUNCATED
case "abort":
sample.status = Sample.Status.ABORTED
case "stop":
sample.status = Sample.Status.COMPLETED
return sample
async def generate_and_rm(args, sample: Sample, sampling_params: dict, evaluation=False) -> Sample:
# generate
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(args, sample, sampling_params)
# 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 and sample.reward is not None
return sample
if sample.aborted:
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(args, 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
@@ -88,7 +137,60 @@ async def generate_and_rm(args, sample: Sample, sampling_params: dict, evaluatio
return sample
async def generate_rollout_async(args, rollout_id, data_buffer) -> list[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, 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, data_buffer):
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", {"rid": ""}, 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
await data_buffer.add_samples(group)
count += len(group)
if args.partial_rollout:
print(f"Collected {count} partial samples into the data buffer", flush=True)
async def generate_rollout_async(args, rollout_id: int, data_buffer) -> list[Sample]:
"""An example to implement the generate_rollout function for an rule based rm rollout generation.
Args:
@@ -97,162 +199,71 @@ async def generate_rollout_async(args, rollout_id, data_buffer) -> list[Sample]:
data_buffer: the data buffer to store the generated samples
Returns:
list[Sample]: a list of samples generated by the rollout
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
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,
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
)
# if over_sampling is set, the sampling batch size can be larger than
# the required rollout batch size
sampling_batch_size = (
args.over_sampling_batch_size if args.over_sampling_batch_size is not None else args.rollout_batch_size
)
# get data from the global_dataset
samples = await data_buffer.get_samples(sampling_batch_size * args.n_samples_per_prompt)
# 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
state = GenerateState(
remaining_batch_size=0,
pendings=set(),
)
def submit_generate_tasks(samples):
for sample in samples:
state.pendings.add(
asyncio.create_task(
generate_and_rm(
args,
sample,
sampling_params=sampling_params,
evaluation=False,
)
)
)
state.remaining_batch_size += len(samples) // args.n_samples_per_prompt
# submit the generation requests.
submit_generate_tasks(samples)
do_dynamic_sampling = args.over_sampling_batch_size and args.dynamic_sampling_filter_path is not None
# load multiple time, so the filter should have no side effect, which should be rational?
if do_dynamic_sampling:
assert args.dynamic_sampling_filter_path is not None
dynamic_sampling_filter = load_function(args.dynamic_sampling_filter_path)
elif args.over_sampling_batch_size is not None:
assert args.over_sampling_filter_path is not None
over_sampling_filter = load_function(args.over_sampling_filter_path)
data_group = {}
data = []
do_print = True
# when doing dynamic sampling, we will use the first rollout_batch_size samples.
target_data_size = (
args.rollout_batch_size if do_dynamic_sampling else sampling_batch_size
) * args.n_samples_per_prompt
pbar = tqdm(total=target_data_size, desc="Rollout generation")
pbar = tqdm(total=target_data_size * args.n_samples_per_prompt, desc="Rollout generation")
while len(data) < target_data_size:
done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED)
# Always finish all done tasks. This will make the code of partial rollout cleaner.
# The assumption here is that group_rm is not too slow.
for task in done:
sample = task.result()
while state.remaining_batch_size < target_data_size:
# get samples from the buffer and submit the generation requests.
samples = await data_buffer.get_samples(args.over_sampling_batch_size)
state.submit_generate_tasks(samples)
# add sample to its group
group_index = sample.index // args.n_samples_per_prompt
if group_index not in data_group:
data_group[group_index] = []
data_group[group_index].append(sample)
# 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([sample.prompt + sample.response], flush=True)
print(
f"First rollout sample: {[group[0].prompt + group[0].response]}, reward: {group[0].reward}",
flush=True,
)
do_print = False
if not len(data_group[group_index]) == args.n_samples_per_prompt:
# wait for the data_group for this prompt finishing
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
# For some rm, we need to do the rm for all samples in the group at the same time
if args.group_rm:
# TODO: this will stuck the asyncio loop.
rewards = await batched_async_rm(args, data_group[group_index])
for i, sample in enumerate(data_group[group_index]):
if args.reward_key:
sample.reward = rewards[i][args.reward_key]
else:
sample.reward = rewards[i]
if do_dynamic_sampling:
# the group is ready
if dynamic_sampling_filter(args, data_group[group_index]):
# When having enough samples, don't add to data.
if len(data) == target_data_size:
continue
data.extend(data_group[group_index])
del data_group[group_index]
pbar.update(args.n_samples_per_prompt)
else:
# Delete the invalid samples, don't use them in partial rollout.
del data_group[group_index]
state.remaining_batch_size -= 1
if state.remaining_batch_size < args.rollout_batch_size:
print(
f"Remaining batch size not enough, add {sampling_batch_size} prompts, "
f"sample response: {[sample.prompt + sample.response]}"
)
new_samples = await data_buffer.get_samples(sampling_batch_size * args.n_samples_per_prompt)
submit_generate_tasks(new_samples)
else:
# if not dynamic sampling, we will just add the samples to the data
data.extend(data_group[group_index])
# 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)
del data_group[group_index]
pbar.close()
print(f"Finish rollout: {[data[-1][0].prompt + data[-1][0].response]}, reward: {data[-1][0].reward}", flush=True)
print(f"Got {len(data)} samples, sample response: {[sample.prompt + sample.response]}")
# there are still some unfinished requests, abort them
await abort(args, rollout_id, data_buffer)
if do_dynamic_sampling:
response = await get(
f"http://{args.sglang_router_ip}:{args.sglang_router_port}/list_workers", use_http2=args.use_http2
)
for url in response["urls"]:
# abort all the requests
print(f"Abort request for {url}")
await post(f"{url}/abort_request", {"rid": ""}, use_http2=False)
if over_sampling_filter is not None:
data = over_sampling_filter(args, data)[: args.rollout_batch_size]
while state.pendings:
done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED)
for task in done:
sample = task.result()
group_index = sample.index // args.n_samples_per_prompt
if group_index not in data_group:
data_group[group_index] = []
data_group[group_index].append(sample)
assert len(data) == args.rollout_batch_size, f"Got {len(data)} samples, expected {args.rollout_batch_size}"
for group_index, samples in data_group.items():
assert (
len(samples) == args.n_samples_per_prompt
), f"Got {len(samples)} samples, expected {args.n_samples_per_prompt}"
if args.partial_rollout:
data_buffer.add_samples(samples)
assert len(data) == target_data_size, f"Got {len(data)} samples, expected {target_data_size}"
if not do_dynamic_sampling and args.over_sampling_batch_size is not None:
data = over_sampling_filter(args, data)[: args.rollout_batch_size * args.n_samples_per_prompt]
else:
data.sort(key=lambda sample: sample.index)
# flatten the data for backward compatibility
data = sum(data, [])
# reset the global state to prevent effects on the next rollout or eval.
state.reset()
return data
@@ -333,7 +344,7 @@ async def eval_rollout_single_dataset(args, rollout_id, name, path):
for coro in asyncio.as_completed(tasks):
sample = await coro
if do_print:
print([sample.prompt + sample.response], sample.reward)
print([sample.prompt + sample.response], sample.reward, flush=True)
do_print = False
data.append(sample)
pbar.update(1)
@@ -341,10 +352,11 @@ async def eval_rollout_single_dataset(args, rollout_id, name, path):
data.sort(key=lambda sample: sample.index)
reward_key = args.reward_key or args.eval_reward_key
return {
name: {
"rewards": [sample.reward for sample in data],
"truncated": [sample.truncated for sample in data],
"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],
}
}
+35 -33
View File
@@ -194,22 +194,28 @@ def get_slime_extra_args_provider(add_custom_arguments=None):
),
)
# over sampling
# sampling
parser.add_argument(
"--over-sampling-batch-size",
type=int,
default=None,
help=(
"The batch size for over sampling. "
"There are 2 cases for over sampling: "
"1. If `over_sampling_batch_size` is set, and `dynamic_sampling_filter_path` is set, "
"we will do dynamic sampling as in DAPO, in which, we will first sample `over_sampling_batch_size` of prompts, "
"use the function in `dynamic_sampling_filter_path` to check if the responses of the prompt is valid, "
"e.g. not all correct or all wrong. And if there are not enough remaining prompts, "
"we will take another `over_sampling_batch_size` of prompts. When there are enough valid prompts, we will abort the ongoing sampling."
"2. If `over_sampling_batch_size` is set, and `over_sampling_filter_path` is set, "
"the first `rollout_batch_size` of `over_sampling_batch_size` will be selected as the result of the prompt. "
"The `over_sampling_filter_path` should be able to sort prompts by its responses and rewards."
"This defines the granularity of the sampling batch in the rollout function. "
"When the number of available samples falls below the target, a sampling "
"operation of size over_sampling_batch_size will be triggered."
"Regardless of whether partial rollout is used or filters are applied, "
"the sampling granularity is always determined by this value. "
"If this value is None, rollout_batch_size will be used as the default over_sampling_batch_size."
),
)
parser.add_argument(
"--over-sampling-filter-input-size",
type=int,
default=None,
help=(
"This is the input size for the over sampling filter."
"This value will replace the rollout_batch_size as target batch size "
"(number of complete, valid samples to be generated) when the over sampling filter is applied."
),
)
parser.add_argument(
@@ -217,26 +223,24 @@ def get_slime_extra_args_provider(add_custom_arguments=None):
type=str,
default=None,
help=(
"Path to the over-sampling filter function. "
"It should be able to sort prompts by its responses and rewards"
"When --over-sampling-filter-path is set, the first `rollout_batch_size` of "
"`over_sampling_batch_size` will be selected as the result of the prompt. "
"You could use `slime.rollout.filter_hub.oversampling_sampling_filters.sort_by_reward_std` as an example."
"This parameter is used with the over_sampling_filter_input_size. "
"The over sampling filter is applied only after enough data has been generated."
"You could use `slime.rollout.filter_hub.over_sampling_filters.sort_by_reward_std` as an example."
),
)
# dynamic sampling
parser.add_argument(
"--dynamic-sampling-filter-path",
type=str,
default=None,
help=(
"Path to the dynamic sampling filter function. "
"It should be able to judge whether the result of a prompt should be selected or not. "
"When --dynamic-sampling-filter-path is set, the first `rollout_batch_size` that satisfy the filter "
"will be selected as the result of the prompt and --over-sampling-filter-path will be ignored. "
"This is the filter function for dynamic sampling. "
"It should be able to judge whether the result of a prompt should be selected or not."
"We will do dynamic filter for sampling as in DAPO. e.g. not all correct or all wrong samples."
"You could use `slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std` as an example."
),
)
# partial rollout
parser.add_argument(
"--partial-rollout",
action="store_true",
@@ -247,7 +251,6 @@ def get_slime_extra_args_provider(add_custom_arguments=None):
"This is useful for long responses."
),
)
parser.add_argument(
"--custom-generate-function-path",
type=str,
@@ -261,14 +264,13 @@ def get_slime_extra_args_provider(add_custom_arguments=None):
parser.add_argument(
"--buffer-filter-path",
type=str,
default="slime.rollout.filter_hub.buffer_filters.pop_first",
default=None,
help=(
"Path to the buffer filter function. "
"It should be able to select the samples in the buffer. "
"The function should take a list of samples and return a list of samples."
"The function should take list[list[Sample]] and return list[Sample]."
),
)
# update weight
parser.add_argument(
"--update-weight-buffer-size",
@@ -810,14 +812,6 @@ def parse_args(add_custom_arguments=None):
if args.eps_clip_high is None:
args.eps_clip_high = args.eps_clip
if args.over_sampling_batch_size is not None:
assert (
args.over_sampling_batch_size >= args.rollout_batch_size
), "over_sampling_batch_size must be greater than rollout_batch_size"
assert (
args.dynamic_sampling_filter_path is not None or args.over_sampling_filter_path is not None
), "over_sampling_batch_size must be used with dynamic_sampling_filter_path or over_sampling_filter_path"
if args.eval_reward_key is None:
args.eval_reward_key = args.reward_key
@@ -870,6 +864,14 @@ def parse_args(add_custom_arguments=None):
if args.vocab_size and not args.padded_vocab_size:
args.padded_vocab_size = _vocab_size_with_padding(args.vocab_size, args)
if args.over_sampling_batch_size is None:
args.over_sampling_batch_size = args.rollout_batch_size
assert args.over_sampling_batch_size >= args.rollout_batch_size, (
f"over_sampling_batch_size {args.over_sampling_batch_size} should be greater than or equal to "
f"rollout_batch_size {args.rollout_batch_size}"
)
# placeholders
args.seq_length = 4096
args.max_position_embeddings = args.seq_length
+1 -1
View File
@@ -41,7 +41,7 @@ class JsonlDataset:
Sample(
prompt=prompt,
label=data[label_key] if label_key is not None else None,
metadata=data.get(metadata_key, None),
metadata=data.get(metadata_key) or {},
)
)
+23 -7
View File
@@ -1,6 +1,8 @@
import asyncio
import multiprocessing
import random
import socket
import httpx
@@ -67,16 +69,30 @@ def terminate_process(process: multiprocessing.Process, timeout: float = 1.0) ->
process.join()
async def post(url, payload, use_http2=False):
async def post(url, payload, use_http2=False, max_retries=60):
# never timeout
timeout = httpx.Timeout(None)
async with httpx.AsyncClient(http1=not use_http2, http2=use_http2, timeout=timeout) as client:
response = await client.post(url, json=payload or {})
response.raise_for_status()
max_retries = 60
retry_count = 0
while retry_count < max_retries:
try:
output = response.json()
except:
output = response.text
async with httpx.AsyncClient(http1=not use_http2, http2=use_http2, timeout=timeout) as client:
response = await client.post(url, json=payload or {})
response.raise_for_status()
try:
output = response.json()
except:
output = response.text
except Exception as e:
retry_count += 1
print(f"Error: {e}, retrying... (attempt {retry_count}/{max_retries})")
if retry_count >= max_retries:
print(f"Max retries ({max_retries}) reached, failing...")
raise e
await asyncio.sleep(1)
continue
break
return output
+14
View File
@@ -10,3 +10,17 @@ def load_function(path):
module_path, _, attr = path.rpartition(".")
module = importlib.import_module(module_path)
return getattr(module, attr)
class SingletonMeta(type):
"""
A metaclass for creating singleton classes.
"""
_instances = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
+8 -12
View File
@@ -1,21 +1,16 @@
from time import time
from functools import wraps
from contextlib import contextmanager
from functools import wraps
from time import time
from .misc import SingletonMeta
__all__ = ["Timer", "timer"]
class Timer:
_instance = None
def __new__(cls):
if cls._instance is None:
cls._instance = super(Timer, cls).__new__(cls)
cls._instance.timers = {}
cls._instance.start_time = {}
cls._instance.seq_lens = None
return cls._instance
class Timer(metaclass=SingletonMeta):
def __init__(self):
self.timers = {}
self.start_time = {}
def start(self, name):
assert name not in self.timers, f"Timer {name} already started."
@@ -70,6 +65,7 @@ def timer(name_or_func):
return Timer().context(name)
func = name_or_func
@wraps(func)
def wrapper(*args, **kwargs):
with Timer().context(func.__name__):
+20 -11
View File
@@ -1,5 +1,7 @@
from dataclasses import dataclass
from typing import Optional
from dataclasses import dataclass, field
from enum import Enum
from typing import Optional, Union
import torch
@@ -8,17 +10,24 @@ class Sample:
"""The sample generated"""
index: Optional[int] = None
prompt: Optional[str] = None
# prompt
prompt: str = ""
tokens: list[int] = field(default_factory=list)
# response
response: str = ""
response_length: int = 0
label: Optional[str] = None
response: Optional[str] = None
tokens: Optional[list[int]] = None
response_length: Optional[int] = None
truncated: Optional[bool] = None
reward: Optional[float] = None
reward: Optional[Union[float, dict[str, float]]] = None
loss_mask: Optional[list[int]] = None
metadata: Optional[dict] = None
version: int = 0
aborted: bool = False
class Status(Enum):
PENDING = "pending"
COMPLETED = "completed"
TRUNCATED = "truncated"
ABORTED = "aborted"
status: Status = Status.PENDING
metadata: dict = field(default_factory=dict)
@dataclass