mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Merge remote-tracking branch 'upstream/main'
This commit is contained in:
+2
-1
@@ -178,4 +178,5 @@ outputs/
|
||||
tests/
|
||||
local/
|
||||
**/rollout_data/
|
||||
**/buffer_stats/
|
||||
**/buffer_stats/
|
||||
*.out
|
||||
|
||||
@@ -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
@@ -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?**
|
||||
|
||||
|
||||
@@ -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
@@ -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 处理?**
|
||||
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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 {},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user