[refactor] rename buffer to rollout controller and extract buffer to data source (#189)

This commit is contained in:
Zilin Zhu
2025-08-14 16:06:26 +08:00
committed by GitHub
parent dfd91a58ef
commit 3b4eeabf3f
10 changed files with 86 additions and 95 deletions
-3
View File
@@ -16,9 +16,6 @@ class FSDPTrainRayActor(TrainRayActor):
def connect_rollout_engines(self, rollout_engines, rollout_engine_lock):
raise NotImplementedError
def set_data_buffer(self, data_buffer):
raise NotImplementedError
def train(self, rollout_id, with_data_fetching=True):
raise NotImplementedError
-4
View File
@@ -84,7 +84,6 @@ class MegatronTrainRayActor(TrainRayActor):
self.sleep(("model"))
self.rollout_engines = None
self.data_buffer = None
self.rollout_data_postprocess = None
if self.args.rollout_data_postprocess_path is not None:
@@ -150,9 +149,6 @@ class MegatronTrainRayActor(TrainRayActor):
mpu.reload_process_groups()
print_memory("after wake_up model")
def set_data_buffer(self, data_buffer):
self.data_buffer = data_buffer
def _get_rollout_data(self, rollout_data_ref):
# Fetch data through ray on CPU, not sure if this will be performance bottleneck.
# Both first pp stage and the last pp stage will recieve the data.
-2
View File
@@ -109,8 +109,6 @@ class RayTrainGroup:
to update weights after each training stage.
"""
self.rollout = rollout
ray.get([actor.set_data_buffer.remote(rollout.data_buffer) for actor in self._actor_handlers])
return [
actor.connect_rollout_engines.remote(
rollout.rollout_engines,
+7 -68
View File
@@ -8,7 +8,7 @@ import torch
from slime.utils.misc import load_function
from slime.utils.types import Sample
from slime.ray.rollout_data_source import RolloutDataSource
from slime.ray.rollout_data_source import RolloutDataSourceWithBuffer
from slime.utils.ray_utils import Box
from slime.utils.wandb_utils import init_wandb_secondary
@@ -16,13 +16,6 @@ logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)
def pop_first(args, rollout_id, buffer: list[list[Sample]], num_samples: int) -> list[list[Sample]]:
num_to_pop = min(len(buffer), num_samples)
samples = buffer[:num_to_pop]
del buffer[:num_to_pop]
return samples
def log_eval_data(rollout_id, args, data):
log_dict = {}
for key in data.keys():
@@ -43,20 +36,14 @@ def log_eval_data(rollout_id, args, data):
@ray.remote
class Buffer:
class RolloutController:
"""The class to run rollout and convert rollout data to training data."""
def __init__(self, args, wandb_run_id):
self.args = args
init_wandb_secondary(args, wandb_run_id)
self.data_source = RolloutDataSource(args)
# a list of sample group.
# each group has n_samples_per_prompt samples, all of them has the same prompt.
self.buffer: list[list[Sample]] = []
if self.args.buffer_filter_path is None:
self.buffer_filter = pop_first
else:
self.buffer_filter = load_function(self.args.buffer_filter_path)
self.data_source = RolloutDataSourceWithBuffer(args)
self.generate_rollout = load_function(self.args.rollout_function_path)
self.eval_generate_rollout = load_function(self.args.eval_function_path)
@@ -67,43 +54,6 @@ class Buffer:
assert self.args.rollout_global_dataset
return len(self.data_source.dataset) // self.args.rollout_batch_size
# TODO simplify remaining logic
def get_samples(self, num_samples: int) -> list[list[Sample]]:
"""
Return num_samples samples
"""
samples = self._get_samples_from_buffer(num_samples)
num_samples -= len(samples)
if num_samples == 0:
return samples
samples += self.data_source.get_samples(num_samples=num_samples)
return samples
def _get_samples_from_buffer(self, num_samples: int) -> list[list[Sample]]:
if len(self.buffer) == 0 or num_samples == 0:
return []
samples = self.buffer_filter(self.args, self.rollout_id, self.buffer, num_samples)
return samples
def add_samples(self, samples: list[list[Sample]]):
"""
Add a sample group to buffer.
"""
if not samples:
return
assert isinstance(samples, list), f"samples must be a list, got {type(samples)}"
assert isinstance(samples[0], list), f"the elements of samples must be list, got {type(samples[0])}"
for i in range(0, len(samples)):
assert (
len(samples[i]) == self.args.n_samples_per_prompt
), f"the length of the elements of samples must be equal to n_samples_per_prompt, got {len(samples[i])} != {self.args.n_samples_per_prompt}"
group = samples[i] # type: ignore
self.buffer.append(group)
def generate(self, rollout_id):
self.rollout_id = rollout_id
@@ -113,7 +63,7 @@ class Buffer:
)["samples"]
data = [Sample.from_dict(sample) for sample in data]
else:
data = self.generate_rollout(self.args, rollout_id, self, evaluation=False)
data = self.generate_rollout(self.args, rollout_id, self.data_source, evaluation=False)
# flatten the data if it is a list of lists
if isinstance(data[0], list):
data = sum(data, [])
@@ -139,7 +89,7 @@ class Buffer:
# if debug train only, we don't generate evaluation data
return
data = self.eval_generate_rollout(self.args, rollout_id, self, evaluation=True)
data = self.eval_generate_rollout(self.args, rollout_id, self.data_source, evaluation=True)
log_eval_data(rollout_id, self.args, data)
def _convert_samples_to_train_data(self, samples: Union[list[Sample], list[list[Sample]]]):
@@ -178,17 +128,6 @@ class Buffer:
train_data["round_number"] = [sample.metadata["round_number"] for sample in samples]
return train_data
# TODO remove
def update_metadata(self, metadata: dict):
self.data_source.metadata.update(metadata)
# TODO remove
def get_metadata(self):
return self.data_source.metadata
def get_buffer_length(self):
return len(self.buffer)
def save(self, rollout_id):
self.data_source.save(rollout_id)
-4
View File
@@ -67,10 +67,6 @@ class TrainRayActor(RayActor):
def connect_rollout_engines(self, rollout_engines, rollout_engine_lock):
raise NotImplementedError
@abc.abstractmethod
def set_data_buffer(self, data_buffer):
raise NotImplementedError
@abc.abstractmethod
def train(self, rollout_id, rollout_data_ref):
raise NotImplementedError
+4 -4
View File
@@ -6,7 +6,7 @@ import ray
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
from slime.backends.sglang_utils.sglang_engine import SGLangEngine
from slime.ray.buffer import Buffer
from slime.ray.buffer import RolloutController
from slime.utils.http_utils import find_available_port, get_host_info, run_router
from .utils import Lock, NOSET_VISIBLE_DEVICES_ENV_VARS_LIST
from typing import List
@@ -146,7 +146,7 @@ class RolloutManager:
def __init__(self, args, pg, wandb_run_id):
self.args = args
_start_router(args)
self.data_buffer = Buffer.options(
self.controller = RolloutController.options(
num_cpus=1,
num_gpus=0,
).remote(args, wandb_run_id=wandb_run_id)
@@ -161,10 +161,10 @@ class RolloutManager:
).remote()
def async_generate(self, rollout_id):
return self.data_buffer.generate.remote(rollout_id)
return self.controller.generate.remote(rollout_id)
def async_eval(self, rollout_id):
return self.data_buffer.eval.remote(rollout_id)
return self.controller.eval.remote(rollout_id)
def async_offload(self):
return [engine.release_memory_occupation.remote() for engine in self.rollout_engines]
+68
View File
@@ -3,6 +3,7 @@ import os
from pathlib import Path
import torch
from slime.utils.misc import load_function
from slime.utils.data import Dataset
from transformers import AutoTokenizer
from slime.utils.types import Sample
@@ -79,6 +80,9 @@ class RolloutDataSource:
return samples
def add_samples(self, samples: list[list[Sample]]):
raise RuntimeError(f"Cannot add samples to {self.__class__.__name__}. This is a read-only data source.")
def save(self, rollout_id):
if not self.args.rollout_global_dataset:
return
@@ -115,3 +119,67 @@ class RolloutDataSource:
if self.args.rollout_global_dataset and self.args.rollout_shuffle:
self.dataset.shuffle(self.epoch_id)
class RolloutDataSourceWithBuffer(RolloutDataSource):
def __init__(self, args):
super().__init__(args)
self.buffer = []
if self.args.buffer_filter_path is None:
self.buffer_filter = pop_first
else:
self.buffer_filter = load_function(self.args.buffer_filter_path)
def get_samples(self, num_samples: int) -> list[list[Sample]]:
"""
Return num_samples samples
"""
samples = self._get_samples_from_buffer(num_samples)
num_samples -= len(samples)
if num_samples == 0:
return samples
samples += super().get_samples(num_samples=num_samples)
return samples
def _get_samples_from_buffer(self, num_samples: int) -> list[list[Sample]]:
if len(self.buffer) == 0 or num_samples == 0:
return []
samples = self.buffer_filter(self.args, self.rollout_id, self.buffer, num_samples)
return samples
def add_samples(self, samples: list[list[Sample]]):
"""
Add a sample group to buffer.
"""
if not samples:
return
assert isinstance(samples, list), f"samples must be a list, got {type(samples)}"
assert isinstance(samples[0], list), f"the elements of samples must be list, got {type(samples[0])}"
for i in range(0, len(samples)):
assert (
len(samples[i]) == self.args.n_samples_per_prompt
), f"the length of the elements of samples must be equal to n_samples_per_prompt, got {len(samples[i])} != {self.args.n_samples_per_prompt}"
group = samples[i] # type: ignore
self.buffer.append(group)
# TODO remove
def update_metadata(self, metadata: dict):
self.metadata.update(metadata)
# TODO remove
def get_metadata(self):
return self.metadata
def get_buffer_length(self):
return len(self.buffer)
def pop_first(args, rollout_id, buffer: list[list[Sample]], num_samples: int) -> list[list[Sample]]:
num_to_pop = min(len(buffer), num_samples)
samples = buffer[:num_to_pop]
del buffer[:num_to_pop]
return samples
@@ -7,7 +7,6 @@ import requests
from transformers import AutoTokenizer
import wandb
from slime.ray.buffer import Buffer
from slime.utils.async_utils import run
from slime.utils.mask_utils import MultiTurnLossMaskGenerator
from slime.utils.types import Sample
@@ -201,9 +200,7 @@ def start_rollout(api_base_url: str, args, metadata):
print(f"[start_rollout] Failed to send rollout config: {e}")
async def generate_rollout_async(
args, rollout_id: int, data_buffer: Buffer, evaluation: bool = False
) -> Dict[str, Any]:
async def generate_rollout_async(args, rollout_id: int, data_buffer, evaluation: bool = False) -> Dict[str, Any]:
global START_ROLLOUT
if evaluation:
+3 -3
View File
@@ -19,7 +19,7 @@ def train(args):
# calculate num_rollout from num_epoch
num_rollout_per_epoch = None
if args.num_rollout is None:
num_rollout_per_epoch = ray.get(rollout_manager.data_buffer.get_num_rollout_per_epoch.remote())
num_rollout_per_epoch = ray.get(rollout_manager.controller.get_num_rollout_per_epoch.remote())
args.num_rollout = num_rollout_per_epoch * args.num_epoch
assert args.num_rollout > 0
@@ -32,7 +32,7 @@ def train(args):
args.start_rollout_id = start_rollout_ids[0]
if args.rollout_global_dataset:
ray.get(rollout_manager.data_buffer.load.remote(args.start_rollout_id - 1))
ray.get(rollout_manager.controller.load.remote(args.start_rollout_id - 1))
# initialize the connection for weight update during training
ray.get(actor_model.async_init_weight_update_connections(rollout_manager))
@@ -66,7 +66,7 @@ def train(args):
):
ray.get(actor_model.async_save_model(rollout_id))
if args.rollout_global_dataset:
ray.get(rollout_manager.data_buffer.save.remote(rollout_id))
ray.get(rollout_manager.controller.save.remote(rollout_id))
if args.offload:
ray.get(actor_model.async_offload())
+3 -3
View File
@@ -21,7 +21,7 @@ def train(args):
# calculate num_rollout from num_epoch
num_rollout_per_epoch = None
if args.num_rollout is None:
num_rollout_per_epoch = ray.get(rollout_manager.data_buffer.get_num_rollout_per_epoch.remote())
num_rollout_per_epoch = ray.get(rollout_manager.controller.get_num_rollout_per_epoch.remote())
args.num_rollout = num_rollout_per_epoch * args.num_epoch
assert args.num_rollout > 0
@@ -35,7 +35,7 @@ def train(args):
args.start_rollout_id = start_rollout_ids[0]
if args.rollout_global_dataset:
ray.get(rollout_manager.data_buffer.load.remote(args.start_rollout_id - 1))
ray.get(rollout_manager.controller.load.remote(args.start_rollout_id - 1))
# initialize the connection for weight update during training
ray.get(actor_model.async_init_weight_update_connections(rollout_manager))
@@ -62,7 +62,7 @@ def train(args):
):
ray.get(actor_model.async_save_model(rollout_id))
if args.rollout_global_dataset:
ray.get(rollout_manager.data_buffer.save.remote(rollout_id))
ray.get(rollout_manager.controller.save.remote(rollout_id))
if (rollout_id + 1) % args.update_weights_interval == 0:
# sync generate before update weights to prevent update weight in the middle of generation