From 8765a0f8fd5415d0c182b1d747bfbd5fdd55a7af Mon Sep 17 00:00:00 2001 From: zyzshishui <492129152@qq.com> Date: Thu, 19 Jun 2025 22:06:03 +0000 Subject: [PATCH 1/6] Squash partial rollout implementation Co-authored-by: Jiajun Li Co-authored-by: Yuzhen Zhou <492129152@qq.com> --- .gitignore | 1 + docs/partial.md | 38 +++ scripts/run-moonlight.sh | 227 +++++++++++++ scripts/run-qwen-1.5B-partial.sh | 162 ++++++++++ scripts/run-qwen-1.5B.sh | 162 ++++++++++ slime/backends/megatron_utils/data.py | 1 + slime/ray/buffer.py | 64 +++- ...lters.py => diversity_sampling_filters.py} | 0 .../filter_hub/over_sampling_filters.py | 11 +- .../filter_hub/partial_rollout_filters.py | 13 + slime/rollout/rm_hub/deepscaler.py | 2 +- slime/rollout/sglang_example.py | 301 +++++++++++++----- slime/utils/arguments.py | 94 ++++-- slime/utils/types.py | 16 +- 14 files changed, 956 insertions(+), 136 deletions(-) create mode 100644 docs/partial.md create mode 100644 scripts/run-moonlight.sh create mode 100644 scripts/run-qwen-1.5B-partial.sh create mode 100644 scripts/run-qwen-1.5B.sh rename slime/rollout/filter_hub/{dynamic_sampling_filters.py => diversity_sampling_filters.py} (100%) create mode 100644 slime/rollout/filter_hub/partial_rollout_filters.py diff --git a/.gitignore b/.gitignore index b1912e65ce..c840800627 100644 --- a/.gitignore +++ b/.gitignore @@ -5,6 +5,7 @@ __pycache__/ # C extensions *.so +*.out # Distribution / packaging .Python diff --git a/docs/partial.md b/docs/partial.md new file mode 100644 index 0000000000..f035b50994 --- /dev/null +++ b/docs/partial.md @@ -0,0 +1,38 @@ +partial rollout 开发过程中,sgl-router 需要手动重新安装 + +参考教程:[sglang/sgl-router/README.md at 971a0dfa32f7521c77c2eeb1180cc9a4fa0100aa · sgl-project/sglang](https://github.com/sgl-project/sglang/blob/971a0dfa32f7521c77c2eeb1180cc9a4fa0100aa/sgl-router/README.md)(用 Option A!!! B 应该是不行的) + +#### Build Rust Project + +```bash +cargo build +``` + +#### Build Python Binding + +##### Build and Install Wheel + +Build the wheel package: + +``` +pip install setuptools-rust wheel build +python -m build +``` + +上面的指令如果报错说 Rust 编译器版本太旧,不能构建 `icu_normalizer v2.0.0`,它需要 Rust 1.82 或更新版本,就手动升级一下 + +```bash +curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh +source $HOME/.cargo/env +# 检查是否成功,应该看到 rustc 1.82.0 或更新版本。 +rustc --version +``` + +然后重新执行上面的命令 + +Install the generated wheel: + +``` +pip install --force-reinstall dist/*.whl +``` + diff --git a/scripts/run-moonlight.sh b/scripts/run-moonlight.sh new file mode 100644 index 0000000000..eca2b62bbb --- /dev/null +++ b/scripts/run-moonlight.sh @@ -0,0 +1,227 @@ +#!/bin/bash + +# 调试时用来重启任务 +pkill -9 sglang +sleep 3 +ray stop --force +pkill -9 ray +pkill -9 python +sleep 3 +pkill -9 ray +pkill -9 python + +set -ex + +export PYTHONBUFFERED=16 + +# env +export MASTER_ADDR=${MASTER_PORT:-"127.0.0.1"} +export MASTER_PORT=${MLP_WORKER_0_PORT:-"12345"} +export no_proxy=localhost,127.0.0.1,0.0.0.0,${MASTER_ADDR} + +export TP_SIZE=4 +export PP_SIZE=1 +export CP_SIZE=1 +export EP_SIZE=4 +export ETP_SIZE=2 + +TARGET_VOCAB=163840 + +EXP_NAME="moonlight" + +MOE_ROUTED_EXPERTS=64 +MOE_ACTIVE_ROUTED_EXPERTS=6 +MOE_SHARED_EXPERTS=2 + +NHIDDEN=2048 +MOE_FFN_HIDDEN=1408 +MOE_SHARED_EXPERT_INTERMEDIATE_SIZE=$(($MOE_FFN_HIDDEN * $MOE_SHARED_EXPERTS)) +MOE_ROUTER_GROUP_TOPK=1 +MOE_ROUTER_NUM_GROUPS=1 +MOE_ROUTER_TOPK_SCALING_FACTOR=2.446 +FFN_HIDDEN=11264 +NLAYERS=27 +NHEADS=16 +FIRST_K_DENSE_REPLACE=1 + +SEQ_LEN=8192 + +# 1) 构造数组 +arr=() +for ((i=0; i list[Sample]: samples = await self._get_samples_from_buffer(num_samples) num_samples -= len(samples) + if num_samples == 0: + return 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): @@ -161,20 +162,65 @@ async def get_samples(self, num_samples) -> list[Sample]: ) self.sample_index += 1 samples.append(sample) - - assert len(samples) == num_samples return samples async def _get_samples_from_buffer(self, num_samples) -> list[Sample]: if len(self.buffer) == 0 or num_samples == 0: return [] - samples = self.buffer_filter(self.buffer, num_samples) + + if self.args.partial_rollout: + partial_ids = [i for i, s in enumerate(self.buffer) if s.status != Sample.Status.PENDING] + new_ids = [i for i, s in enumerate(self.buffer) if s.status == Sample.Status.PENDING] + + # Use configurable ratio of partial samples for diversity + max_partial_samples = min(len(partial_ids), int(num_samples * self.args.partial_rollout_mix_ratio / self.args.n_samples_per_prompt) * self.args.n_samples_per_prompt) + selected_partial_ids = partial_ids[:max_partial_samples] + selected_partial_samples = [self.buffer[i] for i in selected_partial_ids] + + # Get remaining samples using regular filter + needed_new_samples_num = min(len(new_ids), num_samples - len(selected_partial_samples)) + if needed_new_samples_num > 0: + selected_new_ids = new_ids[:needed_new_samples_num] + selected_new_samples = [self.buffer[i] for i in selected_new_ids] + else: + selected_new_ids = [] + selected_new_samples = [] + + samples = selected_partial_samples + selected_new_samples + # Manually remove samples from buffer + to_remove = set(selected_partial_ids + selected_new_ids) + self.buffer = [s for idx, s in enumerate(self.buffer) if idx not in to_remove] + else: + # Original behavior for non-partial rollout + samples = self.buffer_filter(self.buffer, num_samples) return samples async def add_samples(self, samples: list[Sample]): - # TODO: we can save some partial rollout data here. - assert len(samples) % self.args.n_samples_per_prompt == 0 - self.buffer.extend(samples) + if not samples: + return + + partial_samples = [] + new_samples = [] + + assert self.sample_index % self.args.n_samples_per_prompt == 0, f"Samples in a group should be continuous and divisible by n_samples_per_prompt {self.args.n_samples_per_prompt}" + for sample in samples: + sample.index = self.sample_index + self.sample_index += 1 + if (sample.status != Sample.Status.PENDING): + partial_samples.append(sample) + else: + new_samples.append(sample) + + if partial_samples: + # Add partial samples to front of buffer for priority processing + self.buffer[:0] = partial_samples + + if new_samples: + # Add new samples to end of buffer + self.buffer.extend(new_samples) + + assert (len(new_samples) + len(partial_samples)) % self.args.n_samples_per_prompt == 0 + print(f"Buffer size after adding samples: {len(self.buffer)} (adding {len(partial_samples)} partial samples and {len(new_samples)} new samples)", flush=True) def generate(self, rollout_id, evaluation=False): if not evaluation and self.args.load_debug_rollout_data: diff --git a/slime/rollout/filter_hub/dynamic_sampling_filters.py b/slime/rollout/filter_hub/diversity_sampling_filters.py similarity index 100% rename from slime/rollout/filter_hub/dynamic_sampling_filters.py rename to slime/rollout/filter_hub/diversity_sampling_filters.py diff --git a/slime/rollout/filter_hub/over_sampling_filters.py b/slime/rollout/filter_hub/over_sampling_filters.py index f00dfe40b1..2417287149 100644 --- a/slime/rollout/filter_hub/over_sampling_filters.py +++ b/slime/rollout/filter_hub/over_sampling_filters.py @@ -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] diff --git a/slime/rollout/filter_hub/partial_rollout_filters.py b/slime/rollout/filter_hub/partial_rollout_filters.py new file mode 100644 index 0000000000..c7965c3a27 --- /dev/null +++ b/slime/rollout/filter_hub/partial_rollout_filters.py @@ -0,0 +1,13 @@ +from slime.utils.types import Sample + + +__all__ = ["valid_partial_sample"] + + +def valid_partial_sample(args, sample: Sample, **kwargs): + if (sample.response and + len(sample.response.strip()) < args.partial_rollout_min_response_length): + return False + if (sample.response_length and sample.response_length < args.partial_rollout_min_tokens): + return False + return True diff --git a/slime/rollout/rm_hub/deepscaler.py b/slime/rollout/rm_hub/deepscaler.py index 347f839e46..2c56492d3c 100644 --- a/slime/rollout/rm_hub/deepscaler.py +++ b/slime/rollout/rm_hub/deepscaler.py @@ -7,7 +7,7 @@ def get_deepscaler_rule_based_reward(response, label): elif "###Response" in response: model_solution = response.split("###Response")[1] else: - return 0 + model_solution = response model_answer = extract_answer(model_solution) if model_answer is None: diff --git a/slime/rollout/sglang_example.py b/slime/rollout/sglang_example.py index 4f69d150ea..d81e77cce9 100644 --- a/slime/rollout/sglang_example.py +++ b/slime/rollout/sglang_example.py @@ -4,6 +4,7 @@ from tqdm import tqdm from transformers import AutoTokenizer +import wandb from slime.utils.async_utils import run from slime.utils.data import JsonlDataset @@ -18,15 +19,29 @@ @dataclass class GenerateState: - remaining_batch_size: int = 0 - pendings: set = None - + remaining_batch_size: int = 0 # the number of samples that are done/running/pending + pendings: set = None # the set of pending tasks + input_pending_samples: int = 0 + input_aborted_samples: int = 0 + input_completed_samples: int = 0 + input_truncated_samples: int = 0 + output_aborted_samples: int = 0 + output_completed_samples: int = 0 + output_truncated_samples: int = 0 + cached_pending_samples: int = 0 + cached_aborted_samples: int = 0 + cached_completed_samples: int = 0 + cached_truncated_samples: int = 0 + diversity_filter_excluded_samples: int = 0 + over_sampling_filter_excluded_samples: int = 0 + def stats_to_string(self): + return f"Input pending samples: {self.input_pending_samples}, Input aborted samples: {self.input_aborted_samples}, Input completed samples: {self.input_completed_samples}, Input truncated samples: {self.input_truncated_samples}, Output aborted samples: {self.output_aborted_samples}, Output completed samples: {self.output_completed_samples}, Output truncated samples: {self.output_truncated_samples}, Cached pending samples: {self.cached_pending_samples}, Cached aborted samples: {self.cached_aborted_samples}, Cached completed samples: {self.cached_completed_samples}, Cached truncated samples: {self.cached_truncated_samples}, Diversity filter excluded samples: {self.diversity_filter_excluded_samples}, Over sampling filter excluded samples: {self.over_sampling_filter_excluded_samples}" TOKENIZER = None SEMAPHORE = None -async def generate(args, sample: Sample, sampling_params) -> Sample: +async def generate(args, rollout_id: int, sample: Sample, sampling_params) -> Sample: global TOKENIZER, SEMAPHORE if TOKENIZER is None: TOKENIZER = AutoTokenizer.from_pretrained(args.hf_checkpoint, trust_remote_code=True) @@ -37,41 +52,81 @@ async def generate(args, sample: Sample, sampling_params) -> Sample: ) 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}" + is_partial_sample = (sample.status == Sample.Status.ABORTED) + # Handle partial rollout samples: continue generation from existing response + if is_partial_sample: + # Continue generation from partial response + input_text = sample.prompt + sample.response + else: + # Regular generation from prompt + input_text = sample.prompt + payload = { - "text": sample.prompt, + "text": input_text, "sampling_params": sampling_params, } - while True: + + max_retries = 5 + retry_count = 0 + while retry_count < max_retries: try: async with SEMAPHORE: output = await post(url, payload, use_http2=args.use_http2) + # print(f"Prompt: {input_text[:100]}...", flush=True) + # print(f"Output: {output['text'][:200]}...", flush=True) except Exception as e: - print(f"Error: {e}, retrying...") + # TODO(jiajun): what will happen if the worker has been aborted? + 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 prompt_tokens_ids = TOKENIZER(sample.prompt, add_special_tokens=False)["input_ids"] - response_token_ids = TOKENIZER(output["text"], add_special_tokens=False)["input_ids"] + + if is_partial_sample: + # For partial rollout: combine existing response with new generation + sample.response = sample.response + output["text"] + else: + # Regular generation + sample.response = output["text"] + response_token_ids = 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" + if "rollout_id" not in sample.metadata: + sample.metadata["rollout_id"] = rollout_id + 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: +async def generate_and_rm(args, rollout_id: int, 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 and sample.reward is not None + assert sample.metadata.get("rollout_id", -1) != rollout_id + return 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) + sample = await generate(args, rollout_id, sample, sampling_params) - if sample.aborted: + if sample.status == Sample.Status.ABORTED: return sample # for the rm that need the whole group, we will not do the rm here @@ -97,7 +152,7 @@ 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 @@ -113,25 +168,46 @@ async def generate_rollout_async(args, rollout_id, data_buffer) -> list[Sample]: spaces_between_special_tokens=False, ) - # 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) + # sampling_batch_size refers to the number of samples to get at a time. + # if the number of valid samples obtained is insufficient to support rollout, start the next round of sampling. + if args.sampling_batch_size is not None: + sampling_batch_size = args.sampling_batch_size + else: + sampling_batch_size = args.rollout_batch_size + # Redundant samples: Excess valid samples and unfinished samples. + # For non-partial rollout without diversity filter, sampling only happens once. + # For non-partial rollout with diversity filter, sampling may happens multiple times. + # On-policy is required and all redundant samples are dropped, so the sampling_batch_size should not be too large. + # For partial rollout, the sampling_batch_size should be larger than the rollout_batch_size for over sampling. + # Redundant samples are stored in the buffer and will be used in the next round of sampling. state = GenerateState( remaining_batch_size=0, pendings=set(), ) + + diversity_filter = None + if args.diversity_sampling_filter_path is not None: + # diversity filter is used to filter out samples that is not suitable for training + diversity_filter = load_function(args.diversity_sampling_filter_path) + over_sampling_filter = None + if args.over_sampling_filter_path is not None: + assert args.over_sampling_filter_input_size is not None + # over sampling filter ensures over_sampling_filter_input_size samples are rollout. And pick rollout_batch_size samples from them. + over_sampling_filter = load_function(args.over_sampling_filter_path) + partial_rollout_filter = None + if args.partial_rollout and args.partial_rollout_filter_path is not None: + partial_rollout_filter = load_function(args.partial_rollout_filter_path) - def submit_generate_tasks(samples): + def submit_generate_tasks(samples: list[Sample]): for sample in samples: + if sample.status == Sample.Status.PENDING: + assert len(sample.metadata) == 0, f"Sample {sample} has metadata {sample.metadata}" state.pendings.add( asyncio.create_task( generate_and_rm( args, + rollout_id, sample, sampling_params=sampling_params, evaluation=False, @@ -140,28 +216,73 @@ def submit_generate_tasks(samples): ) 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) + def update_state(sample: Sample, is_output: bool = False, is_cached: bool = False): + if not is_output: + match sample.status: + case Sample.Status.PENDING: + state.input_pending_samples += 1 + case Sample.Status.ABORTED: + state.input_aborted_samples += 1 + case Sample.Status.COMPLETED: + state.input_completed_samples += 1 + case Sample.Status.TRUNCATED: + state.input_truncated_samples += 1 + else: + match sample.status: + case Sample.Status.ABORTED: + state.output_aborted_samples += 1 + case Sample.Status.COMPLETED: + state.output_completed_samples += 1 + case Sample.Status.TRUNCATED: + state.output_truncated_samples += 1 + if is_cached: + match sample.status: + case Sample.Status.PENDING: + state.cached_pending_samples += 1 + case Sample.Status.ABORTED: + state.cached_aborted_samples += 1 + case Sample.Status.COMPLETED: + state.cached_completed_samples += 1 + case Sample.Status.TRUNCATED: + state.cached_truncated_samples += 1 data_group = {} data = [] - do_print = True - # when doing dynamic sampling, we will use the first rollout_batch_size samples. + async def abort_rollout(): + print(f"DEBUG: Sending abort. Current data: {len(data)}, pending tasks: {len(state.pendings)}", flush=True) + try: + response = await get( + f"http://{args.sglang_router_ip}:{args.sglang_router_port}/list_workers", use_http2=args.use_http2 + ) + print(f"DEBUG: List workers: {response}", flush=True) + except Exception as e: + print(f"Error: {e}, Failed to get list_workers", flush=True) + return [] + + for url in response["urls"]: + # abort all the requests + # NOTE: Using empty string as rid to abort ALL requests by startswith() match + print(f"Abort request for {url}", flush=True) + await post(f"{url}/abort_request", {"rid": ""}, use_http2=False) + have_aborted = False + + # target_data_size is the total number of valid samples to get + # if over_sampling_filter_input_size is set, we will use it as the target data size, otherwise, we will use the rollout_batch_size target_data_size = ( - args.rollout_batch_size if do_dynamic_sampling else sampling_batch_size + args.over_sampling_filter_input_size if over_sampling_filter is not None else args.rollout_batch_size ) * args.n_samples_per_prompt + pbar = tqdm(total=target_data_size, desc="Rollout generation") while len(data) < target_data_size: + while state.remaining_batch_size < target_data_size // args.n_samples_per_prompt: + # get samples from the buffer and submit the generation requests. + samples = await data_buffer.get_samples(sampling_batch_size * args.n_samples_per_prompt) + submit_generate_tasks(samples) + for sample in samples: + update_state(sample, is_output=False, is_cached=False) + + # wait for the generation to finish 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. @@ -174,9 +295,8 @@ def submit_generate_tasks(samples): data_group[group_index] = [] data_group[group_index].append(sample) - if do_print: - print([sample.prompt + sample.response], flush=True) - do_print = False + assert sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED, f"Sample {sample.index} has status {sample.status}, but should be completed or truncated. Rollout {rollout_id}, Sample prompt: {sample.prompt}, Sample response: {sample.response}" + update_state(sample, is_output=True) if not len(data_group[group_index]) == args.n_samples_per_prompt: # wait for the data_group for this prompt finishing @@ -191,68 +311,80 @@ def submit_generate_tasks(samples): 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 + + if diversity_filter is not None and not diversity_filter(args, data_group[group_index]): + # Delete the invalid samples, don't use them in partial rollout. + del data_group[group_index] + state.remaining_batch_size -= 1 + state.diversity_filter_excluded_samples += args.n_samples_per_prompt + continue + + # add the samples to the data + if len(data) < target_data_size: data.extend(data_group[group_index]) - pbar.update(args.n_samples_per_prompt) del data_group[group_index] + pbar.update(args.n_samples_per_prompt) + + # When having enough samples, try abort the rollout and continue + if len(data) >= target_data_size and not have_aborted: + await abort_rollout() + have_aborted = True + pbar.close() - print(f"Got {len(data)} samples, sample response: {[sample.prompt + sample.response]}") + print(f"[DEBUG] Rollout {rollout_id}: Got {len(data)} samples", flush=True) - 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) + # there are still some unfinished requests, abort them + if state.pendings: + if not have_aborted: + await abort_rollout() + have_aborted = True + if args.partial_rollout: + # put unfinished samples to data group while state.pendings: done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED) for task in done: sample = task.result() + update_state(sample, is_output=True, is_cached=False) + 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) - 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: + # try cache unfinished samples and excess valid samples back to buffer + cache_num = 0 + 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}" + # Only store samples that have reasonable partial content + cached_samples = [] + for sample in samples: + # For half-completed or complete samples from incomplete groups, cache them + if sample.status != Sample.Status.PENDING: + if (partial_rollout_filter is not None and not partial_rollout_filter(args, sample)): + # If the sample cannot pass partial rollout filter, we will regenerate it in the next iteration + sample.response = None + sample.metadata = {} + sample.status = Sample.Status.PENDING + update_state(sample, is_output=True, is_cached=True) + cached_samples.append(sample) + + if len(cached_samples) > 0: + # Add cached samples back to buffer for next iteration + await data_buffer.add_samples(cached_samples) + print(f"[DEBUG] Rollout {rollout_id}: Cached {len(cached_samples)} samples back to buffer", flush=True) + + if over_sampling_filter is not None: + state.over_sampling_filter_excluded_samples += len(data) - args.rollout_batch_size * args.n_samples_per_prompt data = over_sampling_filter(args, data)[: args.rollout_batch_size * args.n_samples_per_prompt] else: data.sort(key=lambda sample: sample.index) + print(f"[DEBUG] Rollout {rollout_id}: {state.stats_to_string()}", flush=True) + assert len(data) == args.rollout_batch_size * args.n_samples_per_prompt, f"Got {len(data)} samples, expected {args.rollout_batch_size * args.n_samples_per_prompt}" return data @@ -321,6 +453,7 @@ async def eval_rollout_single_dataset(args, rollout_id, name, path): tasks.append( generate_and_rm( args, + rollout_id, sample, sampling_params=sampling_params, evaluation=True, @@ -344,7 +477,7 @@ async def eval_rollout_single_dataset(args, rollout_id, name, path): return { name: { "rewards": [sample.reward for sample in data], - "truncated": [sample.truncated for sample in data], + "truncated": [sample.status == Sample.Status.TRUNCATED for sample in data], } } diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 5b1052e300..2294ee4d91 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -196,20 +196,26 @@ def add_rollout_arguments(parser): # over sampling parser.add_argument( - "--over-sampling-batch-size", + "--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 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 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,24 +223,21 @@ def add_rollout_arguments(parser): 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 + # diversity sampling parser.add_argument( - "--dynamic-sampling-filter-path", + "--diversity-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. " - "You could use `slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std` as an example." + "This is the filter function for diversity sampling. " + "It should be able to judge whether the result of a prompt should be selected or not." + "We will do diversity filter for sampling as in DAPO. e.g. not all correct or all wrong samples." + "You could use `slime.rollout.filter_hub.diversity_sampling_filters.check_reward_nonzero_std` as an example." ), ) parser.add_argument( @@ -247,6 +250,43 @@ def add_rollout_arguments(parser): "This is useful for long responses." ), ) + parser.add_argument( + "--partial-rollout-filter-path", + type=str, + default="slime.rollout.filter_hub.partial_rollout_filters.valid_partial_sample", + help=( + "This is the filter function for partial rollout. " + "It should be able to select the samples in the buffer that are valid for partial rollout recycling." + ), + ) + parser.add_argument( + "--partial-rollout-min-response-length", + type=int, + default=10, + help=( + "Minimum response length (in characters) for a sample to be considered for partial rollout recycling. " + "Shorter responses will be discarded." + ), + ) + + parser.add_argument( + "--partial-rollout-min-tokens", + type=int, + default=5, + help=( + "Minimum number of tokens in the response for a sample to be considered for partial rollout recycling." + ), + ) + + parser.add_argument( + "--partial-rollout-mix-ratio", + type=float, + default=0.3, + help=( + "Maximum ratio of partial rollout samples to use in each batch. " + "Value between 0.0 and 1.0. Default 0.3 means up to 30% partial samples." + ), + ) parser.add_argument( "--custom-generate-function-path", @@ -810,14 +850,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 diff --git a/slime/utils/types.py b/slime/utils/types.py index 54da03be6c..d7b3273edc 100644 --- a/slime/utils/types.py +++ b/slime/utils/types.py @@ -1,8 +1,8 @@ -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Optional +from enum import Enum import torch - @dataclass class Sample: """The sample generated""" @@ -13,13 +13,17 @@ class Sample: response: Optional[str] = None tokens: Optional[list[int]] = None response_length: Optional[int] = None - truncated: Optional[bool] = None reward: Optional[float] = None loss_mask: Optional[list[int]] = None - metadata: Optional[dict] = None + metadata: dict = field(default_factory=dict) version: int = 0 - aborted: bool = False - + + class Status(Enum): + PENDING = "pending" + COMPLETED = "completed" + TRUNCATED = "truncated" + ABORTED = "aborted" + status: Status = Status.PENDING @dataclass class ParamInfo: From 3a2232db7df67a0a4a6d36a31377d0e8717d765d Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Fri, 20 Jun 2025 20:56:56 +0000 Subject: [PATCH 2/6] fix bug on data metadata initialization --- slime/utils/data.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/slime/utils/data.py b/slime/utils/data.py index 875e530963..736e043adb 100644 --- a/slime/utils/data.py +++ b/slime/utils/data.py @@ -41,7 +41,7 @@ def __init__( 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 {}, ) ) From 7f3713aa76f889ffb0884a699082ea59b74c4123 Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Fri, 20 Jun 2025 21:06:00 +0000 Subject: [PATCH 3/6] update qwen3-4B and qwen2.5-1.5B config for partial rollout, delete outdated doc, fix reward return value, cancel filter rename, delete debug output --- docs/partial.md | 38 --- ...h => run-deepseek-r1-distill-qwen-1.5B.sh} | 141 +++++------ scripts/run-moonlight.sh | 227 ------------------ scripts/run-qwen-1.5B-partial.sh | 162 ------------- scripts/run-qwen3-4B.sh | 14 +- slime/ray/buffer.py | 5 +- ...filters.py => dynamic_sampling_filters.py} | 0 slime/rollout/rm_hub/deepscaler.py | 2 +- slime/rollout/sglang_example.py | 29 +-- slime/utils/arguments.py | 10 +- 10 files changed, 99 insertions(+), 529 deletions(-) delete mode 100644 docs/partial.md rename scripts/{run-qwen-1.5B.sh => run-deepseek-r1-distill-qwen-1.5B.sh} (50%) delete mode 100644 scripts/run-moonlight.sh delete mode 100644 scripts/run-qwen-1.5B-partial.sh rename slime/rollout/filter_hub/{diversity_sampling_filters.py => dynamic_sampling_filters.py} (100%) diff --git a/docs/partial.md b/docs/partial.md deleted file mode 100644 index f035b50994..0000000000 --- a/docs/partial.md +++ /dev/null @@ -1,38 +0,0 @@ -partial rollout 开发过程中,sgl-router 需要手动重新安装 - -参考教程:[sglang/sgl-router/README.md at 971a0dfa32f7521c77c2eeb1180cc9a4fa0100aa · sgl-project/sglang](https://github.com/sgl-project/sglang/blob/971a0dfa32f7521c77c2eeb1180cc9a4fa0100aa/sgl-router/README.md)(用 Option A!!! B 应该是不行的) - -#### Build Rust Project - -```bash -cargo build -``` - -#### Build Python Binding - -##### Build and Install Wheel - -Build the wheel package: - -``` -pip install setuptools-rust wheel build -python -m build -``` - -上面的指令如果报错说 Rust 编译器版本太旧,不能构建 `icu_normalizer v2.0.0`,它需要 Rust 1.82 或更新版本,就手动升级一下 - -```bash -curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -source $HOME/.cargo/env -# 检查是否成功,应该看到 rustc 1.82.0 或更新版本。 -rustc --version -``` - -然后重新执行上面的命令 - -Install the generated wheel: - -``` -pip install --force-reinstall dist/*.whl -``` - diff --git a/scripts/run-qwen-1.5B.sh b/scripts/run-deepseek-r1-distill-qwen-1.5B.sh similarity index 50% rename from scripts/run-qwen-1.5B.sh rename to scripts/run-deepseek-r1-distill-qwen-1.5B.sh index d5346d5a6c..78fac1a388 100644 --- a/scripts/run-qwen-1.5B.sh +++ b/scripts/run-deepseek-r1-distill-qwen-1.5B.sh @@ -1,7 +1,8 @@ #!/bin/bash -# 调试时用来重启任务 +# for rerun the task pkill -9 sglang +sleep 3 ray stop --force pkill -9 ray pkill -9 python @@ -9,109 +10,89 @@ sleep 3 pkill -9 ray pkill -9 python +set -ex +# will prevent ray from buffering stdout/stderr export PYTHONBUFFERED=16 -# network -export MASTER_ADDR=${MASTER_PORT:-"127.0.0.1"} -export MASTER_PORT=${MLP_WORKER_0_PORT:-"12345"} -export no_proxy=localhost,127.0.0.1,0.0.0.0,${MASTER_ADDR} - -export TP_SIZE=1 -export PP_SIZE=1 -export CP_SIZE=1 - -# qwen2.5 1.5B -MODEL_ARGS=( - --swiglu - --num-layers 28 - --hidden-size 1536 - --ffn-hidden-size 8960 - --num-attention-heads 12 - --max-position-embeddings 32768 - --seq-length 4096 - --use-rotary-position-embeddings - --disable-bias-linear - --add-qkv-bias - --normalization "RMSNorm" - --norm-epsilon 1e-6 - --rotary-base 10000 - --attention-backend auto - --group-query-attention - --num-query-groups 2 - --vocab-size 151936 - --accumulate-allreduce-grads-in-fp32 - --attention-softmax-in-fp32 - --attention-backend flash - --moe-token-dispatcher-type alltoall - --untie-embeddings-and-output-weights - --attention-dropout 0.0 - --hidden-dropout 0.0 -) +# put wandb key here if error +# export WANDB_KEY="abcdefg" + +SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" &>/dev/null && pwd)" +source "${SCRIPT_DIR}/models/qwen2.5-1.5B.sh" CKPT_ARGS=( - --hf-checkpoint /root/DeepSeek-R1-Distill-Qwen-1.5B/ + --hf-checkpoint /root/DeepSeek-R1-Distill-Qwen-1.5B --ref-load /root/DeepSeek-R1-Distill-Qwen-1.5B_torch_dist + # --load /root/DeepSeek-R1-Distill-Qwen-1.5B_slime/ + # --save /root/DeepSeek-R1-Distill-Qwen-1.5B_slime/ + --save-interval 20 ) ROLLOUT_ARGS=( - --rm-type deepscaler - --prompt-data /root/deepscaler/deepscaler.jsonl - --apply-chat-template + --prompt-data /root/dapo-math-17k/dapo-math-17k.jsonl --input-key prompt --label-key label + --apply-chat-template + --rollout-shuffle + --rm-type deepscaler --num-rollout 3000 --rollout-batch-size 128 + --n-samples-per-prompt 8 --rollout-max-response-len 8192 --rollout-temperature 0.8 - --rollout-shuffle - --n-samples-per-prompt 8 + --global-batch-size 1024 - --micro-batch-size 1 - --ref-micro-batch-size 1 - --use-dynamic-batch-size - --max-tokens-per-gpu 9216 --balance-data --sampling-batch-size 128 + # --partial-rollout # --partial-rollout-min-response-length 20 # --partial-rollout-min-tokens 8 - # --partial-rollout-mix-ratio 0.5 + # --partial-rollout-mix-ratio 0.75 # --over-sampling-filter-path slime.rollout.filter_hub.over_sampling_filters.sort_by_reward_std - # --over-sampling-filter-input-size 160 - # --diversity-sampling-filter-path slime.rollout.filter_hub.diversity_sampling_filters.check_reward_nonzero_std + # --over-sampling-filter-input-size 192 + --dynamic-sampling-filter-path slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std ) EVAL_ARGS=( --eval-interval 20 - --eval-prompt-data /root/aime-2024/aime-2024.jsonl - --eval-max-response-len 16500 + --eval-prompt-data aime /root/aime-2024/aime-2024.jsonl --n-samples-per-eval-prompt 16 + --eval-max-response-len 16384 + --eval-top-p 0.7 ) -DISTRIBUTED_ARGS=( - --tensor-model-parallel-size ${TP_SIZE} - --pipeline-model-parallel-size ${PP_SIZE} - --context-parallel-size ${CP_SIZE} +PERF_ARGS=( + --tensor-model-parallel-size 1 --sequence-parallel -) + --pipeline-model-parallel-size 1 + --context-parallel-size 1 + --expert-model-parallel-size 1 + --expert-tensor-parallel-size 1 -PERF_ARGS=( --recompute-granularity full --recompute-method uniform --recompute-num-layers 1 + + # --micro-batch-size 1 + --use-dynamic-batch-size + --max-tokens-per-gpu 9216 ) GRPO_ARGS=( --advantage-estimator grpo --use-kl-loss - --kl-loss-coef 0.001 + --kl-loss-coef 0.00 --kl-loss-type low_var_kl --kl-coef 0.00 --entropy-coef 0.00 + --eps-clip 0.2 + --eps-clip-high 0.28 ) OPTIMIZER_ARGS=( + --optimizer adam --lr 1e-6 --lr-decay-style constant --weight-decay 0.1 @@ -120,37 +101,45 @@ OPTIMIZER_ARGS=( ) WANDB_ARGS=( - --use-wandb - --wandb-key ${WANDB_API_KEY} - --wandb-project slime-guagua-1.5B - --wandb-group slime-guagua-non-partial + #--use-wandb + # --wandb-project slime-dev + # --wandb-group qwen2.5-1.5B-test + # --wandb-key ${WANDB_KEY} +) + +SGLANG_ARGS=( + --rollout-num-gpus-per-engine 1 + --sglang-mem-fraction-static 0.7 +) + +MISC_ARGS=( + # default dropout in megatron is 0.1 + --attention-dropout 0.0 + --hidden-dropout 0.0 + # should be good for model performance + --accumulate-allreduce-grads-in-fp32 + --attention-softmax-in-fp32 + # need to comment this when using model with MLA + --attention-backend flash ) # launch the master node of ray in container +export MASTER_ADDR=${MASTER_ADDR:-"127.0.0.1"} +export MASTER_PORT=${MASTER_PORT:-"12345"} ray start --head --node-ip-address ${MASTER_ADDR} --num-gpus 8 --disable-usage-stats ray job submit --address="http://127.0.0.1:8265" \ --runtime-env-json='{ "env_vars": { - "no_proxy": "localhost,127.0.0.1,0.0.0.0,${MASTER_ADDR}", - "GLOO_SOCKET_IFNAME": "${MLP_SOCKET_IFNAME}", - "TP_SOCKET_IFNAME": "${MLP_SOCKET_IFNAME}", - "NCCL_SOCKET_IFNAME": "${MLP_SOCKET_IFNAME}", - "MASTER_ADDR": "${MASTER_ADDR}", - "MASTER_PORT": "${MASTER_PORT}", "PYTHONPATH": "/root/Megatron-LM/", "CUDA_DEVICE_MAX_CONNECTIONS": "1", - "NVTE_CUDA_INCLUDE_DIR": "/usr/local/cuda/include" + "NCCL_CUMEM_ENABLE": "0" } }' \ -- python3 train.py \ --actor-num-nodes 1 \ --actor-num-gpus-per-node 8 \ - --rollout-num-gpus 8 \ - --rollout-num-gpus-per-engine 1 \ - --sglang-mem-fraction-static 0.5 \ --colocate \ - --offload \ ${MODEL_ARGS[@]} \ ${CKPT_ARGS[@]} \ ${ROLLOUT_ARGS[@]} \ @@ -159,4 +148,6 @@ ray job submit --address="http://127.0.0.1:8265" \ ${DISTRIBUTED_ARGS[@]} \ ${WANDB_ARGS[@]} \ ${PERF_ARGS[@]} \ - ${EVAL_ARGS[@]} \ No newline at end of file + ${EVAL_ARGS[@]} \ + ${SGLANG_ARGS[@]} \ + ${MISC_ARGS[@]} diff --git a/scripts/run-moonlight.sh b/scripts/run-moonlight.sh deleted file mode 100644 index eca2b62bbb..0000000000 --- a/scripts/run-moonlight.sh +++ /dev/null @@ -1,227 +0,0 @@ -#!/bin/bash - -# 调试时用来重启任务 -pkill -9 sglang -sleep 3 -ray stop --force -pkill -9 ray -pkill -9 python -sleep 3 -pkill -9 ray -pkill -9 python - -set -ex - -export PYTHONBUFFERED=16 - -# env -export MASTER_ADDR=${MASTER_PORT:-"127.0.0.1"} -export MASTER_PORT=${MLP_WORKER_0_PORT:-"12345"} -export no_proxy=localhost,127.0.0.1,0.0.0.0,${MASTER_ADDR} - -export TP_SIZE=4 -export PP_SIZE=1 -export CP_SIZE=1 -export EP_SIZE=4 -export ETP_SIZE=2 - -TARGET_VOCAB=163840 - -EXP_NAME="moonlight" - -MOE_ROUTED_EXPERTS=64 -MOE_ACTIVE_ROUTED_EXPERTS=6 -MOE_SHARED_EXPERTS=2 - -NHIDDEN=2048 -MOE_FFN_HIDDEN=1408 -MOE_SHARED_EXPERT_INTERMEDIATE_SIZE=$(($MOE_FFN_HIDDEN * $MOE_SHARED_EXPERTS)) -MOE_ROUTER_GROUP_TOPK=1 -MOE_ROUTER_NUM_GROUPS=1 -MOE_ROUTER_TOPK_SCALING_FACTOR=2.446 -FFN_HIDDEN=11264 -NLAYERS=27 -NHEADS=16 -FIRST_K_DENSE_REPLACE=1 - -SEQ_LEN=8192 - -# 1) 构造数组 -arr=() -for ((i=0; i/dev/null && pwd)" source "${SCRIPT_DIR}/models/qwen3-4B.sh" @@ -33,9 +36,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 @@ -44,6 +45,15 @@ ROLLOUT_ARGS=( --global-batch-size 256 --balance-data + --sampling-batch-size 32 + + # --partial-rollout + # --partial-rollout-min-response-length 20 + # --partial-rollout-min-tokens 8 + # --partial-rollout-mix-ratio 0.75 + # --over-sampling-filter-path slime.rollout.filter_hub.over_sampling_filters.sort_by_reward_std + # --over-sampling-filter-input-size 48 + --dynamic-sampling-filter-path slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std ) EVAL_ARGS=( diff --git a/slime/ray/buffer.py b/slime/ray/buffer.py index 171882eb51..caf9c67048 100644 --- a/slime/ray/buffer.py +++ b/slime/ray/buffer.py @@ -169,15 +169,16 @@ async def _get_samples_from_buffer(self, num_samples) -> list[Sample]: return [] if self.args.partial_rollout: + # TODO(jiajun): Current buffer filter is too naive, and will cause performance degradation in sample selection partial_ids = [i for i, s in enumerate(self.buffer) if s.status != Sample.Status.PENDING] new_ids = [i for i, s in enumerate(self.buffer) if s.status == Sample.Status.PENDING] - # Use configurable ratio of partial samples for diversity + # Use configurable ratio of partial samples for dynamic max_partial_samples = min(len(partial_ids), int(num_samples * self.args.partial_rollout_mix_ratio / self.args.n_samples_per_prompt) * self.args.n_samples_per_prompt) selected_partial_ids = partial_ids[:max_partial_samples] selected_partial_samples = [self.buffer[i] for i in selected_partial_ids] - # Get remaining samples using regular filter + # Get remaining samples using regular filters needed_new_samples_num = min(len(new_ids), num_samples - len(selected_partial_samples)) if needed_new_samples_num > 0: selected_new_ids = new_ids[:needed_new_samples_num] diff --git a/slime/rollout/filter_hub/diversity_sampling_filters.py b/slime/rollout/filter_hub/dynamic_sampling_filters.py similarity index 100% rename from slime/rollout/filter_hub/diversity_sampling_filters.py rename to slime/rollout/filter_hub/dynamic_sampling_filters.py diff --git a/slime/rollout/rm_hub/deepscaler.py b/slime/rollout/rm_hub/deepscaler.py index 2c56492d3c..347f839e46 100644 --- a/slime/rollout/rm_hub/deepscaler.py +++ b/slime/rollout/rm_hub/deepscaler.py @@ -7,7 +7,7 @@ def get_deepscaler_rule_based_reward(response, label): elif "###Response" in response: model_solution = response.split("###Response")[1] else: - model_solution = response + return 0 model_answer = extract_answer(model_solution) if model_answer is None: diff --git a/slime/rollout/sglang_example.py b/slime/rollout/sglang_example.py index d81e77cce9..b40a023878 100644 --- a/slime/rollout/sglang_example.py +++ b/slime/rollout/sglang_example.py @@ -32,10 +32,10 @@ class GenerateState: cached_aborted_samples: int = 0 cached_completed_samples: int = 0 cached_truncated_samples: int = 0 - diversity_filter_excluded_samples: int = 0 + dynamic_filter_excluded_samples: int = 0 over_sampling_filter_excluded_samples: int = 0 def stats_to_string(self): - return f"Input pending samples: {self.input_pending_samples}, Input aborted samples: {self.input_aborted_samples}, Input completed samples: {self.input_completed_samples}, Input truncated samples: {self.input_truncated_samples}, Output aborted samples: {self.output_aborted_samples}, Output completed samples: {self.output_completed_samples}, Output truncated samples: {self.output_truncated_samples}, Cached pending samples: {self.cached_pending_samples}, Cached aborted samples: {self.cached_aborted_samples}, Cached completed samples: {self.cached_completed_samples}, Cached truncated samples: {self.cached_truncated_samples}, Diversity filter excluded samples: {self.diversity_filter_excluded_samples}, Over sampling filter excluded samples: {self.over_sampling_filter_excluded_samples}" + return f"Input pending samples: {self.input_pending_samples}, Input aborted samples: {self.input_aborted_samples}, Input completed samples: {self.input_completed_samples}, Input truncated samples: {self.input_truncated_samples}, Output aborted samples: {self.output_aborted_samples}, Output completed samples: {self.output_completed_samples}, Output truncated samples: {self.output_truncated_samples}, Cached pending samples: {self.cached_pending_samples}, Cached aborted samples: {self.cached_aborted_samples}, Cached completed samples: {self.cached_completed_samples}, Cached truncated samples: {self.cached_truncated_samples}, dynamic filter excluded samples: {self.dynamic_filter_excluded_samples}, Over sampling filter excluded samples: {self.over_sampling_filter_excluded_samples}" TOKENIZER = None SEMAPHORE = None @@ -74,10 +74,7 @@ async def generate(args, rollout_id: int, sample: Sample, sampling_params) -> Sa try: async with SEMAPHORE: output = await post(url, payload, use_http2=args.use_http2) - # print(f"Prompt: {input_text[:100]}...", flush=True) - # print(f"Output: {output['text'][:200]}...", flush=True) except Exception as e: - # TODO(jiajun): what will happen if the worker has been aborted? retry_count += 1 print(f"Error: {e}, retrying... (attempt {retry_count}/{max_retries})") if retry_count >= max_retries: @@ -170,26 +167,24 @@ async def generate_rollout_async(args, rollout_id, data_buffer) -> list[Sample]: # sampling_batch_size refers to the number of samples to get at a time. # if the number of valid samples obtained is insufficient to support rollout, start the next round of sampling. + # redundant samples: aborted samples and excess completed and truncated samples. + # for non-partial rollout with dynamic filter, on-policy is required and all redundant samples are dropped, so the sampling_batch_size should not be too large. + # for partial rollout, redundant samples are stored in the buffer and will be used in the next round of sampling. + if args.sampling_batch_size is not None: sampling_batch_size = args.sampling_batch_size else: sampling_batch_size = args.rollout_batch_size - # Redundant samples: Excess valid samples and unfinished samples. - # For non-partial rollout without diversity filter, sampling only happens once. - # For non-partial rollout with diversity filter, sampling may happens multiple times. - # On-policy is required and all redundant samples are dropped, so the sampling_batch_size should not be too large. - # For partial rollout, the sampling_batch_size should be larger than the rollout_batch_size for over sampling. - # Redundant samples are stored in the buffer and will be used in the next round of sampling. state = GenerateState( remaining_batch_size=0, pendings=set(), ) - diversity_filter = None - if args.diversity_sampling_filter_path is not None: - # diversity filter is used to filter out samples that is not suitable for training - diversity_filter = load_function(args.diversity_sampling_filter_path) + dynamic_filter = None + if args.dynamic_sampling_filter_path is not None: + # dynamic filter is used to filter out samples that is not suitable for training + dynamic_filter = load_function(args.dynamic_sampling_filter_path) over_sampling_filter = None if args.over_sampling_filter_path is not None: assert args.over_sampling_filter_input_size is not None @@ -312,11 +307,11 @@ async def abort_rollout(): else: sample.reward = rewards[i] - if diversity_filter is not None and not diversity_filter(args, data_group[group_index]): + if dynamic_filter is not None and not dynamic_filter(args, data_group[group_index]): # Delete the invalid samples, don't use them in partial rollout. del data_group[group_index] state.remaining_batch_size -= 1 - state.diversity_filter_excluded_samples += args.n_samples_per_prompt + state.dynamic_filter_excluded_samples += args.n_samples_per_prompt continue # add the samples to the data diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 2294ee4d91..c0b2ffd272 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -228,16 +228,16 @@ def add_rollout_arguments(parser): "You could use `slime.rollout.filter_hub.over_sampling_filters.sort_by_reward_std` as an example." ), ) - # diversity sampling + # dynamic sampling parser.add_argument( - "--diversity-sampling-filter-path", + "--dynamic-sampling-filter-path", type=str, default=None, help=( - "This is the filter function for diversity sampling. " + "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 diversity filter for sampling as in DAPO. e.g. not all correct or all wrong samples." - "You could use `slime.rollout.filter_hub.diversity_sampling_filters.check_reward_nonzero_std` as an example." + "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." ), ) parser.add_argument( From cebe34b0ed21f57092c84361ba59084d3df25791 Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Mon, 23 Jun 2025 06:21:49 +0000 Subject: [PATCH 4/6] update buffer read and write operation --- scripts/run-deepseek-r1-distill-qwen-1.5B.sh | 1 - slime/ray/buffer.py | 76 ++++----------- slime/rollout/buffer_hub/fifo.py | 20 ++++ slime/rollout/buffer_hub/partial_fifo.py | 94 +++++++++++++++++++ slime/rollout/filter_hub/buffer_filters.py | 10 -- .../filter_hub/partial_rollout_filters.py | 13 --- slime/rollout/sglang_example.py | 70 ++++++-------- slime/utils/arguments.py | 35 +++---- 8 files changed, 180 insertions(+), 139 deletions(-) create mode 100644 slime/rollout/buffer_hub/fifo.py create mode 100644 slime/rollout/buffer_hub/partial_fifo.py delete mode 100644 slime/rollout/filter_hub/buffer_filters.py delete mode 100644 slime/rollout/filter_hub/partial_rollout_filters.py diff --git a/scripts/run-deepseek-r1-distill-qwen-1.5B.sh b/scripts/run-deepseek-r1-distill-qwen-1.5B.sh index 78fac1a388..4943b35640 100644 --- a/scripts/run-deepseek-r1-distill-qwen-1.5B.sh +++ b/scripts/run-deepseek-r1-distill-qwen-1.5B.sh @@ -14,7 +14,6 @@ set -ex # will prevent ray from buffering stdout/stderr export PYTHONBUFFERED=16 - # put wandb key here if error # export WANDB_KEY="abcdefg" diff --git a/slime/ray/buffer.py b/slime/ray/buffer.py index caf9c67048..6477aa7ef3 100644 --- a/slime/ray/buffer.py +++ b/slime/ray/buffer.py @@ -49,7 +49,8 @@ def __init__(self, args): self.args = args self.buffer = [] - self.buffer_filter = load_function(self.args.buffer_filter_path) + self.read_function = load_function(self.args.buffer_read_function_path) + self.write_function = load_function(self.args.buffer_write_function_path) self.train_data_pool = {} self.eval_data_pool = {} @@ -123,19 +124,19 @@ def _init_wandb(self): 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, rollout_info: dict[str, Any]) -> list[Sample]: """ Return num_samples samples """ - samples = await self._get_samples_from_buffer(num_samples) + samples = await self._get_samples_from_buffer(num_samples, rollout_info) num_samples -= len(samples) - if num_samples == 0: - return 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): @@ -164,64 +165,23 @@ async def get_samples(self, num_samples) -> list[Sample]: samples.append(sample) return samples - async def _get_samples_from_buffer(self, num_samples) -> list[Sample]: + async def _get_samples_from_buffer(self, num_samples: int, rollout_info: dict[str, Any]) -> list[Sample]: if len(self.buffer) == 0 or num_samples == 0: return [] - - if self.args.partial_rollout: - # TODO(jiajun): Current buffer filter is too naive, and will cause performance degradation in sample selection - partial_ids = [i for i, s in enumerate(self.buffer) if s.status != Sample.Status.PENDING] - new_ids = [i for i, s in enumerate(self.buffer) if s.status == Sample.Status.PENDING] - - # Use configurable ratio of partial samples for dynamic - max_partial_samples = min(len(partial_ids), int(num_samples * self.args.partial_rollout_mix_ratio / self.args.n_samples_per_prompt) * self.args.n_samples_per_prompt) - selected_partial_ids = partial_ids[:max_partial_samples] - selected_partial_samples = [self.buffer[i] for i in selected_partial_ids] - - # Get remaining samples using regular filters - needed_new_samples_num = min(len(new_ids), num_samples - len(selected_partial_samples)) - if needed_new_samples_num > 0: - selected_new_ids = new_ids[:needed_new_samples_num] - selected_new_samples = [self.buffer[i] for i in selected_new_ids] - else: - selected_new_ids = [] - selected_new_samples = [] - - samples = selected_partial_samples + selected_new_samples - # Manually remove samples from buffer - to_remove = set(selected_partial_ids + selected_new_ids) - self.buffer = [s for idx, s in enumerate(self.buffer) if idx not in to_remove] - else: - # Original behavior for non-partial rollout - samples = self.buffer_filter(self.buffer, num_samples) + + samples = self.read_function(self.args, self.buffer, num_samples, rollout_info) return samples - async def add_samples(self, samples: list[Sample]): + async def add_samples(self, samples: list[Sample], rollout_info: dict[str, Any]): + """ + Add a sample group to buffer. + """ if not samples: return - - partial_samples = [] - new_samples = [] - - assert self.sample_index % self.args.n_samples_per_prompt == 0, f"Samples in a group should be continuous and divisible by n_samples_per_prompt {self.args.n_samples_per_prompt}" - for sample in samples: - sample.index = self.sample_index - self.sample_index += 1 - if (sample.status != Sample.Status.PENDING): - partial_samples.append(sample) - else: - new_samples.append(sample) - - if partial_samples: - # Add partial samples to front of buffer for priority processing - self.buffer[:0] = partial_samples - - if new_samples: - # Add new samples to end of buffer - self.buffer.extend(new_samples) + assert len(samples) % self.args.n_samples_per_prompt == 0, f"Buffer add_samples got {len(samples)} samples, expected {self.args.n_samples_per_prompt}" - assert (len(new_samples) + len(partial_samples)) % self.args.n_samples_per_prompt == 0 - print(f"Buffer size after adding samples: {len(self.buffer)} (adding {len(partial_samples)} partial samples and {len(new_samples)} new samples)", flush=True) + self.write_function(self.args, self.buffer, samples, rollout_info) + print(f"Buffer size after adding samples: {len(self.buffer)} (adding {len(samples)} samples)", flush=True) def generate(self, rollout_id, evaluation=False): if not evaluation and self.args.load_debug_rollout_data: @@ -298,4 +258,4 @@ def load(self, rollout_id=None): self.metadata = state_dict.get("metadata", {}) if self.args.rollout_global_dataset and self.args.rollout_shuffle: - self.dataset.shuffle(self.epoch_id) + self.dataset.shuffle(self.epoch_id) \ No newline at end of file diff --git a/slime/rollout/buffer_hub/fifo.py b/slime/rollout/buffer_hub/fifo.py new file mode 100644 index 0000000000..e473eaa10e --- /dev/null +++ b/slime/rollout/buffer_hub/fifo.py @@ -0,0 +1,20 @@ +from slime.utils.types import Sample +from typing import Any + +# Simple operations +def push_end(args, buffer: list[list[Sample]], samples: list[Sample], rollout_info: dict[str, Any]): + """ + Simply append the samples to the end of the buffer. + """ + buffer.append(samples) + +def pop_first(args, buffer: list[list[Sample]], num_samples: int, rollout_info: dict[str, Any]): + """ + Try to pop the first `num_samples` from the buffer. + """ + num_to_pop = min(len(buffer), num_samples // args.n_samples_per_prompt) + return_samples = [] + for i in range(num_to_pop): + return_samples.extend(buffer[i]) + del buffer[:num_to_pop] + return return_samples diff --git a/slime/rollout/buffer_hub/partial_fifo.py b/slime/rollout/buffer_hub/partial_fifo.py new file mode 100644 index 0000000000..af41222c73 --- /dev/null +++ b/slime/rollout/buffer_hub/partial_fifo.py @@ -0,0 +1,94 @@ +from slime.utils.types import Sample +from typing import Any + + +# partial filters +def clear_sample_response(sample: Sample): + """ + Clear the response, start_rollout_id, and status of a sample. + """ + sample.response = None + sample.metadata["start_rollout_id"] = None + sample.status = Sample.Status.PENDING + +def valid_partial_sample(args, sample: Sample): + """ + Check if a sample is a valid partial sample. + """ + if (sample.response and + len(sample.response.strip()) < args.partial_rollout_min_response_length): + return False + if (sample.response_length and sample.response_length < args.partial_rollout_min_tokens): + return False + return True + +def filter_partial_samples(args, samples: list[Sample], rollout_info: dict[str, Any]): + """ + Update the partial samples. + If partial samples are too short, we clear the response and status. + """ + for sample in samples: + if sample.status != Sample.Status.PENDING and not valid_partial_sample(args, sample): + clear_sample_response(sample) + # TODO(jiajun): Staleness may be handled here using rollout_info. + +def partial_push_end(args, buffer: list[list[Sample]], samples: list[Sample], rollout_info: dict[str, Any]): + """ + Push the samples to the end of the buffer. + """ + filter_partial_samples(args, samples, rollout_info) + for sample in samples: + # reset partial sample's index without start_rollout_id + if sample.status != Sample.Status.PENDING and sample.metadata.get("start_rollout_id", None) == None: + sample.metadata["start_rollout_id"] = rollout_info["rollout_id"] + buffer.append(samples) + +def partial_pop_first(args, buffer: list[list[Sample]], num_samples: int, rollout_info: dict[str, Any]): + """ + Filter for partial rollout. + This function pops the front `num_samples` from the buffer. + Partial samples are prioritized. + """ + for samples in buffer: + filter_partial_samples(args, samples, rollout_info) + + # Group samples by n_samples_per_prompt + partial_groups_idx = [] + new_groups_idx = [] + + for i in range(0, len(buffer)): + samples = buffer[i] + # Check if all samples in the group are PENDING + all_pending = all(sample.status == Sample.Status.PENDING for sample in samples) + if all_pending: + new_groups_idx.append(i) + else: + partial_groups_idx.append(i) + + # Calculate how many partial groups we can take + num_groups = num_samples // args.n_samples_per_prompt + num_partial_groups = min(int(num_groups * args.partial_rollout_mix_ratio), len(partial_groups_idx)) + selected_partial_groups_idx = partial_groups_idx[:num_partial_groups] + + # Select new groups to fill the remaining quota + num_new_groups = min(num_groups - num_partial_groups, len(new_groups_idx)) + selected_new_groups_idx = new_groups_idx[:num_new_groups] + + # Collect all selected samples + selected_samples = [] + selected_idx = [False] * len(buffer) + new_buffer = [] + for i in selected_partial_groups_idx: + selected_samples.extend(buffer[i]) + selected_idx[i] = True + for i in selected_new_groups_idx: + selected_samples.extend(buffer[i]) + selected_idx[i] = True + for i in range(len(buffer)): + if not selected_idx[i]: + new_buffer.append(buffer[i]) + + buffer = new_buffer + + return selected_samples + diff --git a/slime/rollout/filter_hub/buffer_filters.py b/slime/rollout/filter_hub/buffer_filters.py deleted file mode 100644 index ef5f92548f..0000000000 --- a/slime/rollout/filter_hub/buffer_filters.py +++ /dev/null @@ -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:] diff --git a/slime/rollout/filter_hub/partial_rollout_filters.py b/slime/rollout/filter_hub/partial_rollout_filters.py deleted file mode 100644 index c7965c3a27..0000000000 --- a/slime/rollout/filter_hub/partial_rollout_filters.py +++ /dev/null @@ -1,13 +0,0 @@ -from slime.utils.types import Sample - - -__all__ = ["valid_partial_sample"] - - -def valid_partial_sample(args, sample: Sample, **kwargs): - if (sample.response and - len(sample.response.strip()) < args.partial_rollout_min_response_length): - return False - if (sample.response_length and sample.response_length < args.partial_rollout_min_tokens): - return False - return True diff --git a/slime/rollout/sglang_example.py b/slime/rollout/sglang_example.py index b40a023878..6734fe1315 100644 --- a/slime/rollout/sglang_example.py +++ b/slime/rollout/sglang_example.py @@ -41,7 +41,7 @@ def stats_to_string(self): SEMAPHORE = None -async def generate(args, rollout_id: int, sample: Sample, sampling_params) -> Sample: +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) @@ -96,8 +96,6 @@ async def generate(args, rollout_id: int, sample: Sample, sampling_params) -> Sa response_token_ids = TOKENIZER(sample.response, add_special_tokens=False)["input_ids"] sample.tokens = prompt_tokens_ids + response_token_ids sample.response_length = len(response_token_ids) - if "rollout_id" not in sample.metadata: - sample.metadata["rollout_id"] = rollout_id match output["meta_info"]["finish_reason"]["type"]: case "length": sample.status = Sample.Status.TRUNCATED @@ -109,11 +107,10 @@ async def generate(args, rollout_id: int, sample: Sample, sampling_params) -> Sa return sample -async def generate_and_rm(args, rollout_id: int, sample: Sample, sampling_params: dict, evaluation=False) -> Sample: +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 and sample.reward is not None - assert sample.metadata.get("rollout_id", -1) != rollout_id return sample # generate @@ -121,7 +118,7 @@ async def generate_and_rm(args, rollout_id: int, sample: Sample, sampling_params custom_generate_func = load_function(args.custom_generate_function_path) sample = await custom_generate_func(args, sample, sampling_params) else: - sample = await generate(args, rollout_id, sample, sampling_params) + sample = await generate(args, sample, sampling_params) if sample.status == Sample.Status.ABORTED: return sample @@ -140,7 +137,7 @@ async def generate_and_rm(args, rollout_id: int, sample: Sample, sampling_params return sample -async def generate_rollout_async(args, rollout_id, data_buffer) -> list[Sample]: +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: @@ -190,9 +187,6 @@ async def generate_rollout_async(args, rollout_id, data_buffer) -> list[Sample]: assert args.over_sampling_filter_input_size is not None # over sampling filter ensures over_sampling_filter_input_size samples are rollout. And pick rollout_batch_size samples from them. over_sampling_filter = load_function(args.over_sampling_filter_path) - partial_rollout_filter = None - if args.partial_rollout and args.partial_rollout_filter_path is not None: - partial_rollout_filter = load_function(args.partial_rollout_filter_path) def submit_generate_tasks(samples: list[Sample]): for sample in samples: @@ -202,7 +196,6 @@ def submit_generate_tasks(samples: list[Sample]): asyncio.create_task( generate_and_rm( args, - rollout_id, sample, sampling_params=sampling_params, evaluation=False, @@ -223,23 +216,24 @@ def update_state(sample: Sample, is_output: bool = False, is_cached: bool = Fals case Sample.Status.TRUNCATED: state.input_truncated_samples += 1 else: - match sample.status: - case Sample.Status.ABORTED: - state.output_aborted_samples += 1 - case Sample.Status.COMPLETED: - state.output_completed_samples += 1 - case Sample.Status.TRUNCATED: - state.output_truncated_samples += 1 - if is_cached: - match sample.status: - case Sample.Status.PENDING: - state.cached_pending_samples += 1 - case Sample.Status.ABORTED: - state.cached_aborted_samples += 1 - case Sample.Status.COMPLETED: - state.cached_completed_samples += 1 - case Sample.Status.TRUNCATED: - state.cached_truncated_samples += 1 + if is_cached: + match sample.status: + case Sample.Status.PENDING: + state.cached_pending_samples += 1 + case Sample.Status.ABORTED: + state.cached_aborted_samples += 1 + case Sample.Status.COMPLETED: + state.cached_completed_samples += 1 + case Sample.Status.TRUNCATED: + state.cached_truncated_samples += 1 + else: + match sample.status: + case Sample.Status.ABORTED: + state.output_aborted_samples += 1 + case Sample.Status.COMPLETED: + state.output_completed_samples += 1 + case Sample.Status.TRUNCATED: + state.output_truncated_samples += 1 data_group = {} data = [] @@ -267,12 +261,16 @@ async def abort_rollout(): target_data_size = ( args.over_sampling_filter_input_size if over_sampling_filter is not None else args.rollout_batch_size ) * args.n_samples_per_prompt + # rollout_info is used for sending info to the buffer + rollout_info = { + "rollout_id": rollout_id, + } pbar = tqdm(total=target_data_size, desc="Rollout generation") while len(data) < target_data_size: while state.remaining_batch_size < target_data_size // args.n_samples_per_prompt: # get samples from the buffer and submit the generation requests. - samples = await data_buffer.get_samples(sampling_batch_size * args.n_samples_per_prompt) + samples = await data_buffer.get_samples(sampling_batch_size * args.n_samples_per_prompt, rollout_info) submit_generate_tasks(samples) for sample in samples: update_state(sample, is_output=False, is_cached=False) @@ -291,7 +289,7 @@ async def abort_rollout(): data_group[group_index].append(sample) assert sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED, f"Sample {sample.index} has status {sample.status}, but should be completed or truncated. Rollout {rollout_id}, Sample prompt: {sample.prompt}, Sample response: {sample.response}" - update_state(sample, is_output=True) + update_state(sample, is_output=True, is_cached=False) if not len(data_group[group_index]) == args.n_samples_per_prompt: # wait for the data_group for this prompt finishing @@ -349,27 +347,18 @@ async def abort_rollout(): data_group[group_index].append(sample) # try cache unfinished samples and excess valid samples back to buffer - cache_num = 0 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}" - # Only store samples that have reasonable partial content cached_samples = [] for sample in samples: - # For half-completed or complete samples from incomplete groups, cache them - if sample.status != Sample.Status.PENDING: - if (partial_rollout_filter is not None and not partial_rollout_filter(args, sample)): - # If the sample cannot pass partial rollout filter, we will regenerate it in the next iteration - sample.response = None - sample.metadata = {} - sample.status = Sample.Status.PENDING update_state(sample, is_output=True, is_cached=True) cached_samples.append(sample) if len(cached_samples) > 0: # Add cached samples back to buffer for next iteration - await data_buffer.add_samples(cached_samples) + await data_buffer.add_samples(cached_samples, rollout_info) print(f"[DEBUG] Rollout {rollout_id}: Cached {len(cached_samples)} samples back to buffer", flush=True) if over_sampling_filter is not None: @@ -448,7 +437,6 @@ async def eval_rollout_single_dataset(args, rollout_id, name, path): tasks.append( generate_and_rm( args, - rollout_id, sample, sampling_params=sampling_params, evaluation=True, diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index c0b2ffd272..31c3f36dfa 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -194,7 +194,7 @@ def add_rollout_arguments(parser): ), ) - # over sampling + # sampling parser.add_argument( "--sampling-batch-size", type=int, @@ -228,7 +228,6 @@ def add_rollout_arguments(parser): "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, @@ -240,6 +239,8 @@ def add_rollout_arguments(parser): "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", @@ -250,15 +251,6 @@ def add_rollout_arguments(parser): "This is useful for long responses." ), ) - parser.add_argument( - "--partial-rollout-filter-path", - type=str, - default="slime.rollout.filter_hub.partial_rollout_filters.valid_partial_sample", - help=( - "This is the filter function for partial rollout. " - "It should be able to select the samples in the buffer that are valid for partial rollout recycling." - ), - ) parser.add_argument( "--partial-rollout-min-response-length", type=int, @@ -268,7 +260,6 @@ def add_rollout_arguments(parser): "Shorter responses will be discarded." ), ) - parser.add_argument( "--partial-rollout-min-tokens", type=int, @@ -277,7 +268,6 @@ def add_rollout_arguments(parser): "Minimum number of tokens in the response for a sample to be considered for partial rollout recycling." ), ) - parser.add_argument( "--partial-rollout-mix-ratio", type=float, @@ -299,16 +289,25 @@ def add_rollout_arguments(parser): ) parser.add_argument( - "--buffer-filter-path", + "--buffer-read-function-path", type=str, - default="slime.rollout.filter_hub.buffer_filters.pop_first", + default="slime.rollout.buffer_hub.fifo.pop_first", 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." ), ) - + parser.add_argument( + "--buffer-write-function-path", + type=str, + default="slime.rollout.buffer_hub.fifo.push_end", + help=( + "Path to the buffer write function. " + "It should be able to write the samples to the buffer. " + "The function should take a list of samples and write them to the buffer." + ), + ) # update weight parser.add_argument( "--update-weight-buffer-size", @@ -901,6 +900,10 @@ 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.partial_rollout and args.buffer_read_function_path is None: + args.buffer_read_function_path = "slime.rollout.buffer_hub.partial_fifo.partial_pop_first" + args.buffer_write_function_path = "slime.rollout.buffer_hub.partial_fifo.partial_push_end" # placeholders args.seq_length = 4096 From 1ae9b163b2a928e077dd96247313b1bcd7b8a9ef Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Wed, 25 Jun 2025 01:22:56 +0000 Subject: [PATCH 5/6] Update stats calculation --- slime/rollout/sglang_example.py | 80 +++++++++++---------------------- 1 file changed, 27 insertions(+), 53 deletions(-) diff --git a/slime/rollout/sglang_example.py b/slime/rollout/sglang_example.py index 6734fe1315..8c98ec0857 100644 --- a/slime/rollout/sglang_example.py +++ b/slime/rollout/sglang_example.py @@ -1,6 +1,10 @@ import asyncio import copy -from dataclasses import dataclass +from dataclasses import dataclass, field +from typing import Any, DefaultDict, Set, Counter +from collections import defaultdict +import json +from enum import Enum from tqdm import tqdm from transformers import AutoTokenizer @@ -21,21 +25,22 @@ class GenerateState: remaining_batch_size: int = 0 # the number of samples that are done/running/pending pendings: set = None # the set of pending tasks - input_pending_samples: int = 0 - input_aborted_samples: int = 0 - input_completed_samples: int = 0 - input_truncated_samples: int = 0 - output_aborted_samples: int = 0 - output_completed_samples: int = 0 - output_truncated_samples: int = 0 - cached_pending_samples: int = 0 - cached_aborted_samples: int = 0 - cached_completed_samples: int = 0 - cached_truncated_samples: int = 0 - dynamic_filter_excluded_samples: int = 0 - over_sampling_filter_excluded_samples: int = 0 - def stats_to_string(self): - return f"Input pending samples: {self.input_pending_samples}, Input aborted samples: {self.input_aborted_samples}, Input completed samples: {self.input_completed_samples}, Input truncated samples: {self.input_truncated_samples}, Output aborted samples: {self.output_aborted_samples}, Output completed samples: {self.output_completed_samples}, Output truncated samples: {self.output_truncated_samples}, Cached pending samples: {self.cached_pending_samples}, Cached aborted samples: {self.cached_aborted_samples}, Cached completed samples: {self.cached_completed_samples}, Cached truncated samples: {self.cached_truncated_samples}, dynamic filter excluded samples: {self.dynamic_filter_excluded_samples}, Over sampling filter excluded samples: {self.over_sampling_filter_excluded_samples}" + num_stats: DefaultDict[str, Counter] = field( + default_factory=lambda: defaultdict(Counter) + ) + def bump(self, group: str, value: str, n: int): + self.num_stats[group][value] += n + def stats_to_string(self) -> str: + flat = { + f"{g}_{k}_samples": v + for g, ctr in self.num_stats.items() + for k, v in ctr.items() + } + return json.dumps(flat, ensure_ascii=False) + +def update_num_stats(state: GenerateState, group: str, value: Any, n: int = 1): + vstr = value.name.lower() if isinstance(value, Enum) else str(value) + state.bump(group, vstr, n) TOKENIZER = None SEMAPHORE = None @@ -204,37 +209,6 @@ def submit_generate_tasks(samples: list[Sample]): ) state.remaining_batch_size += len(samples) // args.n_samples_per_prompt - def update_state(sample: Sample, is_output: bool = False, is_cached: bool = False): - if not is_output: - match sample.status: - case Sample.Status.PENDING: - state.input_pending_samples += 1 - case Sample.Status.ABORTED: - state.input_aborted_samples += 1 - case Sample.Status.COMPLETED: - state.input_completed_samples += 1 - case Sample.Status.TRUNCATED: - state.input_truncated_samples += 1 - else: - if is_cached: - match sample.status: - case Sample.Status.PENDING: - state.cached_pending_samples += 1 - case Sample.Status.ABORTED: - state.cached_aborted_samples += 1 - case Sample.Status.COMPLETED: - state.cached_completed_samples += 1 - case Sample.Status.TRUNCATED: - state.cached_truncated_samples += 1 - else: - match sample.status: - case Sample.Status.ABORTED: - state.output_aborted_samples += 1 - case Sample.Status.COMPLETED: - state.output_completed_samples += 1 - case Sample.Status.TRUNCATED: - state.output_truncated_samples += 1 - data_group = {} data = [] @@ -273,7 +247,7 @@ async def abort_rollout(): samples = await data_buffer.get_samples(sampling_batch_size * args.n_samples_per_prompt, rollout_info) submit_generate_tasks(samples) for sample in samples: - update_state(sample, is_output=False, is_cached=False) + update_num_stats(state, "input", sample.status) # wait for the generation to finish done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED) @@ -289,7 +263,7 @@ async def abort_rollout(): data_group[group_index].append(sample) assert sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED, f"Sample {sample.index} has status {sample.status}, but should be completed or truncated. Rollout {rollout_id}, Sample prompt: {sample.prompt}, Sample response: {sample.response}" - update_state(sample, is_output=True, is_cached=False) + update_num_stats(state, "output", sample.status) if not len(data_group[group_index]) == args.n_samples_per_prompt: # wait for the data_group for this prompt finishing @@ -309,7 +283,7 @@ async def abort_rollout(): # Delete the invalid samples, don't use them in partial rollout. del data_group[group_index] state.remaining_batch_size -= 1 - state.dynamic_filter_excluded_samples += args.n_samples_per_prompt + update_num_stats(state, "excluded", "dynamic_sampling_filter", args.n_samples_per_prompt) continue # add the samples to the data @@ -339,7 +313,7 @@ async def abort_rollout(): done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED) for task in done: sample = task.result() - update_state(sample, is_output=True, is_cached=False) + update_num_stats(state, "output", sample.status) group_index = sample.index // args.n_samples_per_prompt if group_index not in data_group: @@ -353,7 +327,7 @@ async def abort_rollout(): ), f"Got {len(samples)} samples, expected {args.n_samples_per_prompt}" cached_samples = [] for sample in samples: - update_state(sample, is_output=True, is_cached=True) + update_num_stats(state, "cached", sample.status) cached_samples.append(sample) if len(cached_samples) > 0: @@ -362,7 +336,7 @@ async def abort_rollout(): print(f"[DEBUG] Rollout {rollout_id}: Cached {len(cached_samples)} samples back to buffer", flush=True) if over_sampling_filter is not None: - state.over_sampling_filter_excluded_samples += len(data) - args.rollout_batch_size * args.n_samples_per_prompt + update_num_stats(state, "excluded", "over_sampling_filter", len(data) - args.rollout_batch_size * args.n_samples_per_prompt) data = over_sampling_filter(args, data)[: args.rollout_batch_size * args.n_samples_per_prompt] else: data.sort(key=lambda sample: sample.index) From 1d1a0ed56433e6cee435a551a2d3fb0f79f16b9a Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Wed, 25 Jun 2025 01:35:03 +0000 Subject: [PATCH 6/6] format code updates using pre-commit --- .gitignore | 4 +- scripts/run-deepseek-r1-distill-qwen-1.5B.sh | 10 ++- scripts/run-qwen3-4B.sh | 5 +- slime/backends/megatron_utils/data.py | 1 - slime/backends/megatron_utils/model.py | 1 + slime/ray/buffer.py | 10 +-- slime/rollout/buffer_hub/fifo.py | 5 +- slime/rollout/buffer_hub/partial_fifo.py | 25 ++++--- slime/rollout/sglang_example.py | 72 +++++++++++--------- slime/utils/arguments.py | 4 +- slime/utils/types.py | 8 ++- 11 files changed, 79 insertions(+), 66 deletions(-) diff --git a/.gitignore b/.gitignore index c840800627..1305da50f8 100644 --- a/.gitignore +++ b/.gitignore @@ -5,7 +5,6 @@ __pycache__/ # C extensions *.so -*.out # Distribution / packaging .Python @@ -179,4 +178,5 @@ outputs/ tests/ local/ **/rollout_data/ -**/buffer_stats/ \ No newline at end of file +**/buffer_stats/ +*.out diff --git a/scripts/run-deepseek-r1-distill-qwen-1.5B.sh b/scripts/run-deepseek-r1-distill-qwen-1.5B.sh index 4943b35640..426424befe 100644 --- a/scripts/run-deepseek-r1-distill-qwen-1.5B.sh +++ b/scripts/run-deepseek-r1-distill-qwen-1.5B.sh @@ -14,8 +14,6 @@ set -ex # will prevent ray from buffering stdout/stderr export PYTHONBUFFERED=16 -# put wandb key here if error -# export WANDB_KEY="abcdefg" SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" &>/dev/null && pwd)" source "${SCRIPT_DIR}/models/qwen2.5-1.5B.sh" @@ -23,8 +21,8 @@ source "${SCRIPT_DIR}/models/qwen2.5-1.5B.sh" CKPT_ARGS=( --hf-checkpoint /root/DeepSeek-R1-Distill-Qwen-1.5B --ref-load /root/DeepSeek-R1-Distill-Qwen-1.5B_torch_dist - # --load /root/DeepSeek-R1-Distill-Qwen-1.5B_slime/ - # --save /root/DeepSeek-R1-Distill-Qwen-1.5B_slime/ + --load /root/DeepSeek-R1-Distill-Qwen-1.5B_slime/ + --save /root/DeepSeek-R1-Distill-Qwen-1.5B_slime/ --save-interval 20 ) @@ -51,7 +49,7 @@ ROLLOUT_ARGS=( # --partial-rollout-mix-ratio 0.75 # --over-sampling-filter-path slime.rollout.filter_hub.over_sampling_filters.sort_by_reward_std # --over-sampling-filter-input-size 192 - --dynamic-sampling-filter-path slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std + # --dynamic-sampling-filter-path slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std ) EVAL_ARGS=( @@ -100,7 +98,7 @@ OPTIMIZER_ARGS=( ) WANDB_ARGS=( - #--use-wandb + # --use-wandb # --wandb-project slime-dev # --wandb-group qwen2.5-1.5B-test # --wandb-key ${WANDB_KEY} diff --git a/scripts/run-qwen3-4B.sh b/scripts/run-qwen3-4B.sh index 724341a8c7..85a5c883ef 100644 --- a/scripts/run-qwen3-4B.sh +++ b/scripts/run-qwen3-4B.sh @@ -15,9 +15,6 @@ set -ex # will prevent ray from buffering stdout/stderr export PYTHONBUFFERED=16 -# put wandb key here if error -# export WANDB_KEY="abcdefg" - SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" &>/dev/null && pwd)" source "${SCRIPT_DIR}/models/qwen3-4B.sh" @@ -53,7 +50,7 @@ ROLLOUT_ARGS=( # --partial-rollout-mix-ratio 0.75 # --over-sampling-filter-path slime.rollout.filter_hub.over_sampling_filters.sort_by_reward_std # --over-sampling-filter-input-size 48 - --dynamic-sampling-filter-path slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std + # --dynamic-sampling-filter-path slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std ) EVAL_ARGS=( diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index 750ffda6fa..2b769bf25a 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -300,7 +300,6 @@ def get_partition(val): "total_lengths", "response_lengths", "rewards", - "raw_reward", "truncated", "loss_masks", "round_number", diff --git a/slime/backends/megatron_utils/model.py b/slime/backends/megatron_utils/model.py index b84147c861..c00215ae35 100644 --- a/slime/backends/megatron_utils/model.py +++ b/slime/backends/megatron_utils/model.py @@ -17,6 +17,7 @@ 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 diff --git a/slime/ray/buffer.py b/slime/ray/buffer.py index 6477aa7ef3..94d8a2afad 100644 --- a/slime/ray/buffer.py +++ b/slime/ray/buffer.py @@ -131,10 +131,10 @@ async def get_samples(self, num_samples: int, rollout_info: dict[str, Any]) -> l samples = await self._get_samples_from_buffer(num_samples, rollout_info) 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 @@ -178,7 +178,9 @@ async def add_samples(self, samples: list[Sample], rollout_info: dict[str, Any]) """ if not samples: return - assert len(samples) % self.args.n_samples_per_prompt == 0, f"Buffer add_samples got {len(samples)} samples, expected {self.args.n_samples_per_prompt}" + assert ( + len(samples) % self.args.n_samples_per_prompt == 0 + ), f"Buffer add_samples got {len(samples)} samples, expected {self.args.n_samples_per_prompt}" self.write_function(self.args, self.buffer, samples, rollout_info) print(f"Buffer size after adding samples: {len(self.buffer)} (adding {len(samples)} samples)", flush=True) @@ -258,4 +260,4 @@ def load(self, rollout_id=None): self.metadata = state_dict.get("metadata", {}) if self.args.rollout_global_dataset and self.args.rollout_shuffle: - self.dataset.shuffle(self.epoch_id) \ No newline at end of file + self.dataset.shuffle(self.epoch_id) diff --git a/slime/rollout/buffer_hub/fifo.py b/slime/rollout/buffer_hub/fifo.py index e473eaa10e..866907e236 100644 --- a/slime/rollout/buffer_hub/fifo.py +++ b/slime/rollout/buffer_hub/fifo.py @@ -1,6 +1,8 @@ -from slime.utils.types import Sample from typing import Any +from slime.utils.types import Sample + + # Simple operations def push_end(args, buffer: list[list[Sample]], samples: list[Sample], rollout_info: dict[str, Any]): """ @@ -8,6 +10,7 @@ def push_end(args, buffer: list[list[Sample]], samples: list[Sample], rollout_in """ buffer.append(samples) + def pop_first(args, buffer: list[list[Sample]], num_samples: int, rollout_info: dict[str, Any]): """ Try to pop the first `num_samples` from the buffer. diff --git a/slime/rollout/buffer_hub/partial_fifo.py b/slime/rollout/buffer_hub/partial_fifo.py index af41222c73..62a605a421 100644 --- a/slime/rollout/buffer_hub/partial_fifo.py +++ b/slime/rollout/buffer_hub/partial_fifo.py @@ -1,6 +1,7 @@ -from slime.utils.types import Sample from typing import Any +from slime.utils.types import Sample + # partial filters def clear_sample_response(sample: Sample): @@ -11,17 +12,18 @@ def clear_sample_response(sample: Sample): sample.metadata["start_rollout_id"] = None sample.status = Sample.Status.PENDING + def valid_partial_sample(args, sample: Sample): """ Check if a sample is a valid partial sample. """ - if (sample.response and - len(sample.response.strip()) < args.partial_rollout_min_response_length): + if sample.response and len(sample.response.strip()) < args.partial_rollout_min_response_length: return False - if (sample.response_length and sample.response_length < args.partial_rollout_min_tokens): + if sample.response_length and sample.response_length < args.partial_rollout_min_tokens: return False return True + def filter_partial_samples(args, samples: list[Sample], rollout_info: dict[str, Any]): """ Update the partial samples. @@ -32,6 +34,7 @@ def filter_partial_samples(args, samples: list[Sample], rollout_info: dict[str, clear_sample_response(sample) # TODO(jiajun): Staleness may be handled here using rollout_info. + def partial_push_end(args, buffer: list[list[Sample]], samples: list[Sample], rollout_info: dict[str, Any]): """ Push the samples to the end of the buffer. @@ -43,6 +46,7 @@ def partial_push_end(args, buffer: list[list[Sample]], samples: list[Sample], ro sample.metadata["start_rollout_id"] = rollout_info["rollout_id"] buffer.append(samples) + def partial_pop_first(args, buffer: list[list[Sample]], num_samples: int, rollout_info: dict[str, Any]): """ Filter for partial rollout. @@ -51,11 +55,11 @@ def partial_pop_first(args, buffer: list[list[Sample]], num_samples: int, rollou """ for samples in buffer: filter_partial_samples(args, samples, rollout_info) - + # Group samples by n_samples_per_prompt partial_groups_idx = [] new_groups_idx = [] - + for i in range(0, len(buffer)): samples = buffer[i] # Check if all samples in the group are PENDING @@ -64,16 +68,16 @@ def partial_pop_first(args, buffer: list[list[Sample]], num_samples: int, rollou new_groups_idx.append(i) else: partial_groups_idx.append(i) - + # Calculate how many partial groups we can take num_groups = num_samples // args.n_samples_per_prompt num_partial_groups = min(int(num_groups * args.partial_rollout_mix_ratio), len(partial_groups_idx)) selected_partial_groups_idx = partial_groups_idx[:num_partial_groups] - + # Select new groups to fill the remaining quota num_new_groups = min(num_groups - num_partial_groups, len(new_groups_idx)) selected_new_groups_idx = new_groups_idx[:num_new_groups] - + # Collect all selected samples selected_samples = [] selected_idx = [False] * len(buffer) @@ -89,6 +93,5 @@ def partial_pop_first(args, buffer: list[list[Sample]], num_samples: int, rollou new_buffer.append(buffer[i]) buffer = new_buffer - - return selected_samples + return selected_samples diff --git a/slime/rollout/sglang_example.py b/slime/rollout/sglang_example.py index 8c98ec0857..92f01a4ffc 100644 --- a/slime/rollout/sglang_example.py +++ b/slime/rollout/sglang_example.py @@ -1,14 +1,13 @@ import asyncio import copy -from dataclasses import dataclass, field -from typing import Any, DefaultDict, Set, Counter -from collections import defaultdict import json +from collections import defaultdict +from dataclasses import dataclass, field from enum import Enum +from typing import Any, Counter, DefaultDict from tqdm import tqdm from transformers import AutoTokenizer -import wandb from slime.utils.async_utils import run from slime.utils.data import JsonlDataset @@ -23,25 +22,23 @@ @dataclass class GenerateState: - remaining_batch_size: int = 0 # the number of samples that are done/running/pending - pendings: set = None # the set of pending tasks - num_stats: DefaultDict[str, Counter] = field( - default_factory=lambda: defaultdict(Counter) - ) + remaining_batch_size: int = 0 # the number of samples that are done/running/pending + pendings: set = None # the set of pending tasks + num_stats: DefaultDict[str, Counter] = field(default_factory=lambda: defaultdict(Counter)) + def bump(self, group: str, value: str, n: int): self.num_stats[group][value] += n + def stats_to_string(self) -> str: - flat = { - f"{g}_{k}_samples": v - for g, ctr in self.num_stats.items() - for k, v in ctr.items() - } + flat = {f"{g}_{k}_samples": v for g, ctr in self.num_stats.items() for k, v in ctr.items()} return json.dumps(flat, ensure_ascii=False) + def update_num_stats(state: GenerateState, group: str, value: Any, n: int = 1): vstr = value.name.lower() if isinstance(value, Enum) else str(value) state.bump(group, vstr, n) + TOKENIZER = None SEMAPHORE = None @@ -57,9 +54,11 @@ async def generate(args, sample: Sample, sampling_params) -> Sample: ) 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}" - is_partial_sample = (sample.status == Sample.Status.ABORTED) + + assert ( + sample.status == Sample.Status.PENDING or sample.status == Sample.Status.ABORTED + ), f"Sample status is {sample.status}" + is_partial_sample = sample.status == Sample.Status.ABORTED # Handle partial rollout samples: continue generation from existing response if is_partial_sample: # Continue generation from partial response @@ -67,13 +66,13 @@ async def generate(args, sample: Sample, sampling_params) -> Sample: else: # Regular generation from prompt input_text = sample.prompt - + payload = { "text": input_text, "sampling_params": sampling_params, } - max_retries = 5 + max_retries = 60 retry_count = 0 while retry_count < max_retries: try: @@ -90,7 +89,7 @@ async def generate(args, sample: Sample, sampling_params) -> Sample: break prompt_tokens_ids = TOKENIZER(sample.prompt, add_special_tokens=False)["input_ids"] - + if is_partial_sample: # For partial rollout: combine existing response with new generation sample.response = sample.response + output["text"] @@ -117,7 +116,7 @@ async def generate_and_rm(args, sample: Sample, sampling_params: dict, evaluatio if sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED: assert sample.response is not None and sample.reward is not None return sample - + # generate if args.custom_generate_function_path is not None: custom_generate_func = load_function(args.custom_generate_function_path) @@ -182,7 +181,7 @@ async def generate_rollout_async(args, rollout_id: int, data_buffer) -> list[Sam remaining_batch_size=0, pendings=set(), ) - + dynamic_filter = None if args.dynamic_sampling_filter_path is not None: # dynamic filter is used to filter out samples that is not suitable for training @@ -228,8 +227,9 @@ async def abort_rollout(): # NOTE: Using empty string as rid to abort ALL requests by startswith() match print(f"Abort request for {url}", flush=True) await post(f"{url}/abort_request", {"rid": ""}, use_http2=False) + have_aborted = False - + # target_data_size is the total number of valid samples to get # if over_sampling_filter_input_size is set, we will use it as the target data size, otherwise, we will use the rollout_batch_size target_data_size = ( @@ -248,7 +248,7 @@ async def abort_rollout(): submit_generate_tasks(samples) for sample in samples: update_num_stats(state, "input", sample.status) - + # wait for the generation to finish 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. @@ -262,7 +262,9 @@ async def abort_rollout(): data_group[group_index] = [] data_group[group_index].append(sample) - assert sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED, f"Sample {sample.index} has status {sample.status}, but should be completed or truncated. Rollout {rollout_id}, Sample prompt: {sample.prompt}, Sample response: {sample.response}" + assert ( + sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED + ), f"Sample {sample.index} has status {sample.status}, but should be completed or truncated. Rollout {rollout_id}, Sample prompt: {sample.prompt}, Sample response: {sample.response}" update_num_stats(state, "output", sample.status) if not len(data_group[group_index]) == args.n_samples_per_prompt: @@ -278,25 +280,25 @@ async def abort_rollout(): sample.reward = rewards[i][args.reward_key] else: sample.reward = rewards[i] - + if dynamic_filter is not None and not dynamic_filter(args, data_group[group_index]): # Delete the invalid samples, don't use them in partial rollout. del data_group[group_index] state.remaining_batch_size -= 1 update_num_stats(state, "excluded", "dynamic_sampling_filter", args.n_samples_per_prompt) continue - + # add the samples to the data if len(data) < target_data_size: data.extend(data_group[group_index]) del data_group[group_index] pbar.update(args.n_samples_per_prompt) - + # When having enough samples, try abort the rollout and continue if len(data) >= target_data_size and not have_aborted: await abort_rollout() have_aborted = True - + pbar.close() print(f"[DEBUG] Rollout {rollout_id}: Got {len(data)} samples", flush=True) @@ -329,20 +331,24 @@ async def abort_rollout(): for sample in samples: update_num_stats(state, "cached", sample.status) cached_samples.append(sample) - + if len(cached_samples) > 0: # Add cached samples back to buffer for next iteration await data_buffer.add_samples(cached_samples, rollout_info) print(f"[DEBUG] Rollout {rollout_id}: Cached {len(cached_samples)} samples back to buffer", flush=True) if over_sampling_filter is not None: - update_num_stats(state, "excluded", "over_sampling_filter", len(data) - args.rollout_batch_size * args.n_samples_per_prompt) + update_num_stats( + state, "excluded", "over_sampling_filter", len(data) - args.rollout_batch_size * args.n_samples_per_prompt + ) data = over_sampling_filter(args, data)[: args.rollout_batch_size * args.n_samples_per_prompt] else: data.sort(key=lambda sample: sample.index) - print(f"[DEBUG] Rollout {rollout_id}: {state.stats_to_string()}", flush=True) - assert len(data) == args.rollout_batch_size * args.n_samples_per_prompt, f"Got {len(data)} samples, expected {args.rollout_batch_size * args.n_samples_per_prompt}" + print(f"[DEBUG] Rollout {rollout_id}: {state.stats_to_string()}", flush=True) + assert ( + len(data) == args.rollout_batch_size * args.n_samples_per_prompt + ), f"Got {len(data)} samples, expected {args.rollout_batch_size * args.n_samples_per_prompt}" return data diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 31c3f36dfa..6318fb27c5 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -239,7 +239,7 @@ def add_rollout_arguments(parser): "You could use `slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std` as an example." ), ) - + # partial rollout parser.add_argument( "--partial-rollout", @@ -900,7 +900,7 @@ 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.partial_rollout and args.buffer_read_function_path is None: args.buffer_read_function_path = "slime.rollout.buffer_hub.partial_fifo.partial_pop_first" args.buffer_write_function_path = "slime.rollout.buffer_hub.partial_fifo.partial_push_end" diff --git a/slime/utils/types.py b/slime/utils/types.py index d7b3273edc..0a8d1bdd9c 100644 --- a/slime/utils/types.py +++ b/slime/utils/types.py @@ -1,8 +1,10 @@ from dataclasses import dataclass, field -from typing import Optional from enum import Enum +from typing import Optional + import torch + @dataclass class Sample: """The sample generated""" @@ -17,14 +19,16 @@ class Sample: loss_mask: Optional[list[int]] = None metadata: dict = field(default_factory=dict) version: int = 0 - + class Status(Enum): PENDING = "pending" COMPLETED = "completed" TRUNCATED = "truncated" ABORTED = "aborted" + status: Status = Status.PENDING + @dataclass class ParamInfo: name: str