Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -178,4 +178,5 @@ outputs/
tests/
local/
**/rollout_data/
**/buffer_stats/
**/buffer_stats/
*.out
150 changes: 150 additions & 0 deletions scripts/run-deepseek-r1-distill-qwen-1.5B.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,150 @@
#!/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

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
--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=(
--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

--global-batch-size 1024
--balance-data
--sampling-batch-size 128

# --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 192
# --dynamic-sampling-filter-path slime.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std
)

EVAL_ARGS=(
--eval-interval 20
--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
)

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

--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.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 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": {
"PYTHONPATH": "/root/Megatron-LM/",
"CUDA_DEVICE_MAX_CONNECTIONS": "1",
"NCCL_CUMEM_ENABLE": "0"
}
}' \
-- python3 train.py \
--actor-num-nodes 1 \
--actor-num-gpus-per-node 8 \
--colocate \
${MODEL_ARGS[@]} \
${CKPT_ARGS[@]} \
${ROLLOUT_ARGS[@]} \
${OPTIMIZER_ARGS[@]} \
${GRPO_ARGS[@]} \
${DISTRIBUTED_ARGS[@]} \
${WANDB_ARGS[@]} \
${PERF_ARGS[@]} \
${EVAL_ARGS[@]} \
${SGLANG_ARGS[@]} \
${MISC_ARGS[@]}
11 changes: 9 additions & 2 deletions scripts/run-qwen3-4B.sh
Original file line number Diff line number Diff line change
Expand Up @@ -33,9 +33,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
Expand All @@ -44,6 +42,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=(
Expand Down
1 change: 1 addition & 0 deletions slime/backends/megatron_utils/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
33 changes: 21 additions & 12 deletions slime/ray/buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ def convert_samples_to_train_data(samples: list[Sample]):
"tokens": [sample.tokens for sample in samples],
"response_lengths": [sample.response_length for sample in samples],
"rewards": [sample.reward for sample in samples],
"truncated": [1 if sample.truncated else 0 for sample in samples],
"truncated": [1 if sample.status == Sample.Status.TRUNCATED else 0 for sample in samples],
}
if samples[0].loss_mask:
train_data["loss_masks"] = []
Expand All @@ -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 = {}
Expand Down Expand Up @@ -123,13 +124,14 @@ 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)

assert num_samples % self.args.n_samples_per_prompt == 0
num_prompts = num_samples // self.args.n_samples_per_prompt

Expand Down Expand Up @@ -161,20 +163,27 @@ 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]:
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 []
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]):
# TODO: we can save some partial rollout data here.
assert len(samples) % self.args.n_samples_per_prompt == 0
self.buffer.extend(samples)
async def add_samples(self, samples: list[Sample], rollout_info: dict[str, Any]):
"""
Add a sample group to buffer.
"""
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}"

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:
Expand Down
23 changes: 23 additions & 0 deletions slime/rollout/buffer_hub/fifo.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
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]):
"""
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
97 changes: 97 additions & 0 deletions slime/rollout/buffer_hub/partial_fifo.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
from typing import Any

from slime.utils.types import Sample


# 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
10 changes: 0 additions & 10 deletions slime/rollout/filter_hub/buffer_filters.py

This file was deleted.

Loading