mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
[refactor] rename buffer to rollout controller and extract buffer to data source (#189)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user