update code

This commit is contained in:
miles.pr.bot
2025-12-02 11:26:52 +08:00
parent c67babcfad
commit eedc4c4c86
4 changed files with 363 additions and 2 deletions
+2 -2
View File
@@ -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)
+213
View File
@@ -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
+6
View File
@@ -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,
+142
View File
@@ -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[@]}