mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
update code
This commit is contained in:
@@ -13,7 +13,6 @@ import torch
|
||||
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
|
||||
|
||||
from miles.backends.sglang_utils.sglang_engine import SGLangEngine
|
||||
from miles.ray.rollout_data_source import RolloutDataSourceWithBuffer
|
||||
from miles.rollout.base_types import call_rollout_fn
|
||||
from miles.utils import tracking_utils
|
||||
from miles.utils.health_monitor import RolloutHealthMonitor
|
||||
@@ -50,7 +49,8 @@ class RolloutManager:
|
||||
init_tracking(args, primary=False, router_addr=f"http://{args.sglang_router_ip}:{args.sglang_router_port}")
|
||||
init_http_client(args)
|
||||
|
||||
self.data_source = RolloutDataSourceWithBuffer(args)
|
||||
data_source_cls = load_function(self.args.data_source_path)
|
||||
self.data_source = data_source_cls(args)
|
||||
|
||||
self.generate_rollout = load_function(self.args.rollout_function_path)
|
||||
self.eval_generate_rollout = load_function(self.args.eval_function_path)
|
||||
|
||||
@@ -0,0 +1,213 @@
|
||||
import abc
|
||||
import copy
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from miles.utils.data import Dataset
|
||||
from miles.utils.misc import load_function
|
||||
from miles.utils.types import Sample
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DataSource(abc.ABC):
|
||||
@abc.abstractmethod
|
||||
def get_samples(self, num_samples: int) -> list[list[Sample]]:
|
||||
"""
|
||||
Return num_samples samples
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def add_samples(self, samples: list[list[Sample]]):
|
||||
"""
|
||||
Add samples to the data source
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def save(self, rollout_id):
|
||||
"""
|
||||
Save the state of the data source
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def load(self, rollout_id=None):
|
||||
"""
|
||||
Load the state of the data source
|
||||
"""
|
||||
|
||||
|
||||
# TODO may further refactor data-loading part later
|
||||
class RolloutDataSource(DataSource):
|
||||
def __init__(self, args):
|
||||
self.args = args
|
||||
|
||||
self.epoch_id = 0
|
||||
self.sample_group_index = 0
|
||||
self.sample_index = 0
|
||||
self.sample_offset = 0
|
||||
# TODO remove this
|
||||
self.metadata = {}
|
||||
|
||||
if args.rollout_global_dataset:
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.hf_checkpoint, trust_remote_code=True)
|
||||
|
||||
# TODO move (during the refactor)
|
||||
if (d := args.dump_details) is not None:
|
||||
tokenizer.save_pretrained(Path(d) / "tokenizer")
|
||||
|
||||
self.dataset = Dataset(
|
||||
args.prompt_data,
|
||||
tokenizer=tokenizer,
|
||||
max_length=args.rollout_max_prompt_len,
|
||||
prompt_key=args.input_key,
|
||||
label_key=args.label_key,
|
||||
metadata_key=args.metadata_key,
|
||||
tool_key=args.tool_key,
|
||||
apply_chat_template=args.apply_chat_template,
|
||||
apply_chat_template_kwargs=args.apply_chat_template_kwargs,
|
||||
seed=args.rollout_seed,
|
||||
)
|
||||
if self.args.rollout_shuffle:
|
||||
self.dataset.shuffle(self.epoch_id)
|
||||
else:
|
||||
self.dataset = None
|
||||
|
||||
def get_samples(self, num_samples):
|
||||
# TODO further improve code
|
||||
if self.dataset is not None:
|
||||
if self.sample_offset + num_samples <= len(self.dataset):
|
||||
prompt_samples = self.dataset.samples[self.sample_offset : self.sample_offset + num_samples]
|
||||
self.sample_offset += num_samples
|
||||
else:
|
||||
prompt_samples = self.dataset.samples[self.sample_offset :]
|
||||
num_samples -= len(prompt_samples)
|
||||
self.epoch_id += 1
|
||||
if self.args.rollout_shuffle:
|
||||
self.dataset.shuffle(self.epoch_id)
|
||||
prompt_samples += self.dataset.samples[:num_samples]
|
||||
self.sample_offset = num_samples
|
||||
else:
|
||||
prompt_samples = [Sample() for _ in range(num_samples)]
|
||||
|
||||
samples = []
|
||||
for prompt_sample in prompt_samples:
|
||||
group = []
|
||||
for _ in range(self.args.n_samples_per_prompt):
|
||||
sample = copy.deepcopy(prompt_sample)
|
||||
sample.group_index = self.sample_group_index
|
||||
sample.index = self.sample_index
|
||||
self.sample_index += 1
|
||||
group.append(sample)
|
||||
self.sample_group_index += 1
|
||||
samples.append(group)
|
||||
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
|
||||
|
||||
state_dict = {
|
||||
"sample_offset": self.sample_offset,
|
||||
"epoch_id": self.epoch_id,
|
||||
"sample_group_index": self.sample_group_index,
|
||||
"sample_index": self.sample_index,
|
||||
"metadata": self.metadata,
|
||||
}
|
||||
path = os.path.join(self.args.save, f"rollout/global_dataset_state_dict_{rollout_id}.pt")
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
torch.save(state_dict, path)
|
||||
|
||||
def load(self, rollout_id=None):
|
||||
if not self.args.rollout_global_dataset:
|
||||
return
|
||||
|
||||
if self.args.load is None:
|
||||
return
|
||||
|
||||
path = os.path.join(self.args.load, f"rollout/global_dataset_state_dict_{rollout_id}.pt")
|
||||
if not os.path.exists(path):
|
||||
logger.info(f"Checkpoint {path} does not exist.")
|
||||
return
|
||||
|
||||
logger.info(f"load metadata from {path}")
|
||||
logger.info(f"load metadata: {self.metadata}")
|
||||
state_dict = torch.load(path)
|
||||
self.sample_offset = state_dict.get("sample_offset", 0)
|
||||
self.epoch_id = state_dict.get("epoch_id", 0)
|
||||
self.sample_group_index = state_dict.get("sample_group_index", 0)
|
||||
self.sample_index = state_dict.get("sample_index", 0)
|
||||
self.metadata = state_dict.get("metadata", {})
|
||||
|
||||
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, None, 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
|
||||
@@ -444,6 +444,12 @@ def get_miles_extra_args_provider(add_custom_arguments=None):
|
||||
),
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--data-source-path",
|
||||
type=str,
|
||||
default="miles.rollout.data_source.RolloutDataSourceWithBuffer",
|
||||
help="The data source class for rollout data.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-data",
|
||||
type=str,
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
#!/bin/bash
|
||||
|
||||
# for rerun the task
|
||||
pkill -9 sglang
|
||||
sleep 3
|
||||
ray stop --force
|
||||
pkill -9 ray
|
||||
pkill -9 python
|
||||
sleep 3
|
||||
pkill -9 ray
|
||||
pkill -9 python
|
||||
|
||||
|
||||
|
||||
|
||||
set -ex
|
||||
|
||||
# will prevent ray from buffering stdout/stderr
|
||||
export PYTHONBUFFERED=16
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7
|
||||
NVLINK_COUNT=$(nvidia-smi | grep -o "NVLink" | wc -l)
|
||||
if [ "$NVLINK_COUNT" -gt 0 ]; then
|
||||
HAS_NVLINK=1
|
||||
else
|
||||
HAS_NVLINK=0
|
||||
fi
|
||||
echo "HAS_NVLINK: $HAS_NVLINK (detected $NVLINK_COUNT NVLink references)"
|
||||
|
||||
|
||||
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" &>/dev/null && pwd)"
|
||||
|
||||
RUN_ID=${RUN_ID:-"run_$(date +%Y%m%d_%H%M%S)"}
|
||||
LOAD_SAVE_PATH="/root/shared_data/${RUN_ID}/checkpoints"
|
||||
|
||||
CKPT_ARGS=(
|
||||
--hf-checkpoint /root/Qwen3-4B
|
||||
--load /root/Qwen3-4B
|
||||
--ref-load /root/Qwen3-4B
|
||||
)
|
||||
|
||||
ROLLOUT_ARGS=(
|
||||
--prompt-data /root/dapo-math-17k/dapo-math-17k.jsonl
|
||||
--input-key prompt
|
||||
--label-key label
|
||||
--apply-chat-template
|
||||
--rollout-shuffle
|
||||
--balance-data
|
||||
--rm-type deepscaler
|
||||
--num-rollout 100
|
||||
--rollout-batch-size 8
|
||||
--n-samples-per-prompt 8
|
||||
--rollout-max-response-len 4096
|
||||
--rollout-temperature 0.8
|
||||
--global-batch-size 64
|
||||
)
|
||||
|
||||
GRPO_ARGS=(
|
||||
--use-kl-loss
|
||||
--advantage-estimator grpo
|
||||
--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
|
||||
--adam-beta1 0.9
|
||||
--adam-beta2 0.98
|
||||
)
|
||||
|
||||
WANDB_ARGS=(
|
||||
--use-wandb
|
||||
--wandb-project miles-dev-mcore-fsdp
|
||||
--wandb-group qwen3-4B-fsdp-1130-ref
|
||||
--wandb-key ${WANDB_API_KEY}
|
||||
)
|
||||
|
||||
SGLANG_ARGS=(
|
||||
--rollout-num-gpus-per-engine 1
|
||||
--sglang-mem-fraction-static 0.75
|
||||
--sglang-decode-log-interval 1000
|
||||
--sglang-chunked-prefill-size 4096
|
||||
--sglang-attention-backend fa3
|
||||
)
|
||||
|
||||
TRAIN_BACKEND_ARGS=(
|
||||
--train-backend fsdp
|
||||
--update-weight-buffer-size 536870912
|
||||
--gradient-checkpointing
|
||||
--attn-implementation flash_attention_3
|
||||
--train-env-vars '{"PYTORCH_CUDA_ALLOC_CONF":"expandable_segments:True"}'
|
||||
)
|
||||
|
||||
PERF_ARGS=(
|
||||
--use-dynamic-batch-size
|
||||
--max-tokens-per-gpu 9216
|
||||
)
|
||||
|
||||
MISC_ARGS=(
|
||||
--actor-num-nodes 1
|
||||
--actor-num-gpus-per-node 8
|
||||
--colocate
|
||||
--use-fault-tolerance
|
||||
--dump-details /root/shared_data/qwen3-4B-fsdp-1116-noref/dump_details
|
||||
# --fsdp-cpu-offload
|
||||
)
|
||||
|
||||
# launch the master node of ray in container - 8 GPUs for training
|
||||
export MASTER_ADDR=${MASTER_ADDR:-"127.0.0.1"}
|
||||
ray start --head --node-ip-address ${MASTER_ADDR} --num-gpus 8 --disable-usage-stats
|
||||
|
||||
|
||||
RUNTIME_ENV_JSON="{
|
||||
\"env_vars\": {
|
||||
\"PYTHONPATH\": \"/root/Megatron-LM/:${SCRIPT_DIR}\",
|
||||
\"CUDA_DEVICE_MAX_CONNECTIONS\": \"1\"
|
||||
}
|
||||
}"
|
||||
|
||||
|
||||
ray job submit --address="http://127.0.0.1:8265" \
|
||||
--runtime-env-json="${RUNTIME_ENV_JSON}" \
|
||||
-- python3 train.py \
|
||||
${CKPT_ARGS[@]} \
|
||||
${ROLLOUT_ARGS[@]} \
|
||||
${OPTIMIZER_ARGS[@]} \
|
||||
${GRPO_ARGS[@]} \
|
||||
${WANDB_ARGS[@]} \
|
||||
${SGLANG_ARGS[@]} \
|
||||
${TRAIN_BACKEND_ARGS[@]} \
|
||||
${PERF_ARGS[@]} \
|
||||
${MISC_ARGS[@]}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user