Skip to content
Open
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
14 changes: 14 additions & 0 deletions docs/en/advanced/mooncake-store-transfer.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
# Mooncake Store Rollout Transfer

Transfer rollout tensors (`tokens`, `loss_masks`) via [Mooncake Store](https://github.com/kvcache-ai/Mooncake) instead of Ray object refs when rollout and training run on separate GPUs.

```bash
pip install mooncake-transfer-engine
mooncake_master --enable_http_metadata_server=true \
--http_metadata_server_host=127.0.0.1 --http_metadata_server_port=18080
python train_async.py --transfer-backend mooncake_store
```

Optional: `--mooncake-store-init-kwargs '{"master_server_addr":"127.0.0.1:50051"}'` or `MOONCAKE_*` env vars.

Ported from [slime PR #1709](https://github.com/THUDM/slime/pull/1709).
1 change: 1 addition & 0 deletions docs/en/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@ vime is built on `slime <https://github.com/THUDM/slime>`_, the RL framework beh
advanced/fault-tolerance.md
advanced/observability.md
advanced/pd-disaggregation.md
advanced/mooncake-store-transfer.md
advanced/external-rollout-engines.md
advanced/delta-weight-sync.md
advanced/vllm-config.md
Expand Down
133 changes: 133 additions & 0 deletions tests/run_disaggregate_2gpu.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,133 @@
"""Minimal 2-GPU train/inference disaggregation smoke test with Mooncake rollout transfer.

GPU layout (non-colocate):
- GPU 0: Megatron training (actor)
- GPU 1: vLLM rollout (inference)

Rollout tensors are transferred via Mooncake instead of Ray object store.
"""

import os

import vime.utils.external_utils.command_utils as U

MODEL_NAME = "Qwen2.5-0.5B-Instruct"
MODEL_TYPE = "qwen2.5-0.5B"
MODELSCOPE_MODEL_ID = "Qwen/Qwen2.5-0.5B-Instruct"
MODELSCOPE_DATASET_ID = "AI-ModelScope/gsm8k"
DATASET_TRAIN_PATH = "/root/datasets/gsm8k/main/train-00000-of-00001.parquet"
NUM_GPUS = 2


def prepare():
U.exec_command("mkdir -p /root/models /root/datasets")
model_dir = f"/root/models/{MODEL_NAME}"
if not os.path.exists(f"{model_dir}/config.json"):
U.ms_download_model(MODELSCOPE_MODEL_ID, model_dir)
dataset_dir = "/root/datasets/gsm8k"
if not os.path.exists(DATASET_TRAIN_PATH):
U.ms_download_dataset(MODELSCOPE_DATASET_ID, dataset_dir)


def execute():
from vime.utils.mooncake_store_service import ensure_mooncake_master

ensure_mooncake_master()

ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/models/{MODEL_NAME}/ "

rollout_args = (
f"--prompt-data {DATASET_TRAIN_PATH} "
"--input-key question "
"--label-key answer "
"--rollout-shuffle "
"--rm-type math "
"--num-rollout 2 "
"--rollout-batch-size 2 "
"--n-samples-per-prompt 2 "
"--rollout-max-response-len 512 "
"--rollout-temperature 0.8 "
"--global-batch-size 4 "
)

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 "
"--use-dynamic-batch-size "
"--max-tokens-per-gpu 4096 "
)

grpo_args = (
"--advantage-estimator grpo "
"--use-kl-loss "
"--kl-loss-coef 0.00 "
"--kl-loss-type low_var_kl "
"--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 "
)

vllm_args = (
"--rollout-num-gpus-per-engine 1 "
"--vllm-gpu-memory-utilization 0.25 "
"--vllm-max-cudagraph-capture-size 8 "
"--vllm-max-num-seqs 16 "
)

mooncake_args = "--transfer-backend mooncake_store "

ci_args = "--ci-test "

misc_args = (
"--attention-dropout 0.0 "
"--hidden-dropout 0.0 "
"--accumulate-allreduce-grads-in-fp32 "
"--attention-softmax-in-fp32 "
"--no-gradient-accumulation-fusion "
"--attention-backend flash "
"--actor-num-nodes 1 "
"--actor-num-gpus-per-node 1 "
"--rollout-num-gpus 1 "
"--megatron-to-hf-mode bridge "
)

train_args = (
f"{ckpt_args} "
f"{rollout_args} "
f"{optimizer_args} "
f"{grpo_args} "
f"{perf_args} "
f"{vllm_args} "
f"{mooncake_args} "
f"{ci_args} "
f"{misc_args} "
)

U.execute_train(
train_args=train_args,
num_gpus_per_node=NUM_GPUS,
megatron_model_type=MODEL_TYPE,
train_script="train_async.py",
)


if __name__ == "__main__":
prepare()
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
os.environ.pop(proxy_var, None)
# Use free GPUs when 0/1 are occupied (e.g. sglang). Override: CUDA_VISIBLE_DEVICES=4,5
os.environ.setdefault("CUDA_VISIBLE_DEVICES", "4,5")
execute()
58 changes: 58 additions & 0 deletions tests/utils/test_mooncake_store_transfer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
import types

import numpy as np
import torch

from vime.utils.mooncake_store_service import ensure_mooncake_master
from vime.utils.remote_batch import MooncakeRemoteBatch, create_mooncake_store, normalize_store_init_kwargs
from vime.utils.rollout_store_transfer import rollout_store_batch_to_data, split_rollout_data_by_dp_mooncake_store


def test_mooncake_remote_batch_roundtrip():
ensure_mooncake_master()
store = create_mooncake_store()
remote = MooncakeRemoteBatch.from_tensors(
{"tokens": torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.long)},
store,
prefix="vime-test/roundtrip",
)
try:
assert remote.materialize()["tokens"].tolist() == [[1, 2, 3], [4, 5, 6]]
finally:
remote.cleanup()


def test_split_and_materialize_roundtrip():
ensure_mooncake_master()
args = types.SimpleNamespace(transfer_backend="mooncake_store", mooncake_store_init_kwargs=None)
data = {
"tokens": [[1, 2, 3], [4, 5], [6, 7, 8, 9]],
"loss_masks": [[1, 1, 1], [1, 1], [1, 1, 1, 1]],
"response_lengths": [3, 2, 4],
"rewards": [0.1, 0.2, 0.3],
"truncated": [0, 0, 1],
"rollout_ids": [0, 0, 1],
"total_lengths": [3, 2, 4],
"global_batch_sizes": [2, 1],
"num_microbatches": 2,
}
refs = split_rollout_data_by_dp_mooncake_store(
args,
data,
dp_size=2,
partitions=[[0, 1], [2]],
micro_batch_indices=[[0, 1], [0]],
num_microbatches=2,
global_batch_sizes=[2, 1],
)
assert len(refs) == 2
rollout = rollout_store_batch_to_data(refs[0])
assert rollout["partition"] == [0, 1]
assert all(isinstance(row, torch.Tensor) for row in rollout["tokens"])
assert rollout["tokens"][0].tolist() == [1, 2, 3]


def test_normalize_store_init_kwargs_uses_env_defaults():
kwargs = normalize_store_init_kwargs(None)
assert kwargs["protocol"] == "tcp"
assert kwargs["master_server_addr"] == "127.0.0.1:50051"
21 changes: 14 additions & 7 deletions train.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import ray

from vime.utils.rollout_store_transfer import maybe_cleanup_mooncake_store_refs

from vime.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models
from vime.utils.arguments import parse_args
from vime.utils.logging_utils import configure_logger, finish_tracking, init_tracking
Expand Down Expand Up @@ -71,14 +73,19 @@ def save(rollout_id):

actor_trains_this_step = (not args.use_critic) or rollout_id >= args.num_critic_only_steps

if args.use_critic:
value_refs = critic_model.async_train(rollout_id, rollout_data_ref)
if actor_trains_this_step:
ray.get(actor_model.async_train(rollout_id, rollout_data_ref, external_data=value_refs))
train_succeeded = False
try:
if args.use_critic:
value_refs = critic_model.async_train(rollout_id, rollout_data_ref)
if actor_trains_this_step:
ray.get(actor_model.async_train(rollout_id, rollout_data_ref, external_data=value_refs))
else:
ray.get(value_refs)
else:
ray.get(value_refs)
else:
ray.get(actor_model.async_train(rollout_id, rollout_data_ref))
ray.get(actor_model.async_train(rollout_id, rollout_data_ref))
train_succeeded = True
finally:
maybe_cleanup_mooncake_store_refs(args, rollout_data_ref, suppress_errors=not train_succeeded)

if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout):
save(rollout_id)
Expand Down
23 changes: 15 additions & 8 deletions train_async.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import ray

from vime.utils.rollout_store_transfer import maybe_cleanup_mooncake_store_refs

from vime.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models
from vime.utils.arguments import parse_args
from vime.utils.logging_utils import configure_logger, finish_tracking, init_tracking
Expand Down Expand Up @@ -38,15 +40,20 @@ def train(args):
if rollout_id + 1 < args.num_rollout:
rollout_data_next_future = rollout_manager.generate.remote(rollout_id + 1)

if args.use_critic:
actor_trains_this_step = rollout_id >= args.num_critic_only_steps
value_refs = critic_model.async_train(rollout_id, rollout_data_curr_ref)
if actor_trains_this_step:
ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref, external_data=value_refs))
train_succeeded = False
try:
if args.use_critic:
actor_trains_this_step = rollout_id >= args.num_critic_only_steps
value_refs = critic_model.async_train(rollout_id, rollout_data_curr_ref)
if actor_trains_this_step:
ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref, external_data=value_refs))
else:
ray.get(value_refs)
else:
ray.get(value_refs)
else:
ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref))
ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref))
train_succeeded = True
finally:
maybe_cleanup_mooncake_store_refs(args, rollout_data_curr_ref, suppress_errors=not train_succeeded)

if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout):
if (not args.use_critic) or rollout_id >= args.num_critic_only_steps:
Expand Down
13 changes: 13 additions & 0 deletions vime/ray/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -860,6 +860,19 @@ def _split_train_data_by_dp(self, data):
rollout_indices=data["rollout_ids"],
)

if getattr(self.args, "transfer_backend", "ray") == "mooncake_store":
from vime.utils.rollout_store_transfer import split_rollout_data_by_dp_mooncake_store

return split_rollout_data_by_dp_mooncake_store(
self.args,
data,
dp_size,
partitions,
micro_batch_indices=micro_batch_indices,
num_microbatches=num_microbatches,
global_batch_sizes=global_batch_sizes,
)

# Package per-rank rollout_data
rollout_data_refs = []
for r in range(dp_size):
Expand Down
17 changes: 17 additions & 0 deletions vime/utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,6 +403,18 @@ def add_rollout_arguments(parser):
"This is used to shuffle the prompts and also for the random sampling of the prompts."
),
)
parser.add_argument(
"--transfer-backend",
choices=["ray", "mooncake_store"],
default="ray",
help="Rollout data transfer backend. Keep ray as the default; mooncake_store uses Mooncake Store.",
)
parser.add_argument(
"--mooncake-store-init-kwargs",
type=json.loads,
default=None,
help="JSON kwargs passed to MooncakeDistributedStore.setup for mooncake_store transfer.",
)

# sampling
parser.add_argument(
Expand Down Expand Up @@ -1767,6 +1779,11 @@ def _validate_update_weight_args(args) -> None:
def vime_validate_args(args):
args.eval_datasets = _resolve_eval_datasets(args)

if getattr(args, "transfer_backend", "ray") == "mooncake_store":
from vime.utils.remote_batch import normalize_store_init_kwargs

args.mooncake_store_init_kwargs = normalize_store_init_kwargs(args.mooncake_store_init_kwargs)

if args.kl_coef != 0 or args.use_kl_loss:
if not os.path.exists(args.ref_load):
raise FileNotFoundError(f"ref_load {args.ref_load} does not exist, please check the path.")
Expand Down
9 changes: 8 additions & 1 deletion vime/utils/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -291,7 +291,14 @@ def __len__(self):

def process_rollout_data(args, rollout_data_ref, dp_rank, dp_size):
assert len(rollout_data_ref) == dp_size
rollout_data = ray.get(rollout_data_ref[dp_rank].inner)
if getattr(args, "transfer_backend", "ray") == "mooncake_store":
from vime.utils.rollout_store_transfer import RolloutStoreBatch, rollout_store_batch_to_data

batch = rollout_data_ref[dp_rank]
assert isinstance(batch, RolloutStoreBatch), f"expected RolloutStoreBatch, got {type(batch)}"
rollout_data = rollout_store_batch_to_data(batch)
else:
rollout_data = ray.get(rollout_data_ref[dp_rank].inner)

partition = rollout_data.pop("partition")
total_lengths = rollout_data["total_lengths"]
Expand Down
Loading