diff --git a/docs/en/advanced/mooncake-rollout-transfer.md b/docs/en/advanced/mooncake-rollout-transfer.md new file mode 100644 index 000000000..ab66dc4c2 --- /dev/null +++ b/docs/en/advanced/mooncake-rollout-transfer.md @@ -0,0 +1,273 @@ +# Mooncake Rollout Data Transfer + +slime normally passes rollout data from the rollout manager to the trainer through +Ray's object store. Mooncake Store is an alternative for this handoff. It keeps the +rollout dictionary contract unchanged while moving the payload through Mooncake's +data plane. + +This setting only controls rollout data transfer: + +- `--rollout-data-transport` selects how rollout data reaches the trainer. +- Weight synchronization is configured separately by the `--update-weight-*` + options. +- SGLang KV-cache disaggregation uses its own transfer settings. + +Both `train.py` and `train_async.py` support the Mooncake path. + +## Requirements + +Before starting a slime job: + +- Use the same slime revision and Mooncake wheel on every Ray node. The Mooncake + package must provide the structured-object `put`/`get` APIs used by slime. +- Start `mooncake_master`, or provide a Mooncake HA endpoint. slime connects as a + client and does not manage the endpoint lifecycle. +- Use data-network addresses that are reachable from every Ray node. +- Reserve enough host memory for `MOONCAKE_GLOBAL_SEGMENT_SIZE` and + `MOONCAKE_LOCAL_BUFFER_SIZE`. +- For RDMA, expose the RDMA device to the runtime, allow memory locking, and set the + node-local device name before starting Ray. +- Make the model, Megatron checkpoint, dataset, and Python environment available on + every node that may run the corresponding worker. + +Follow the [Mooncake installation guide](https://kvcache-ai.github.io/Mooncake/getting_started/build.html) +for packages and platform requirements. + +## Configure the backend + +Choose TCP or RDMA before starting Ray. TCP only needs a routable data network. +RDMA also needs a local device on each node; device names may differ across nodes. + +Add this option to an existing slime recipe: + +```bash +--rollout-data-transport mooncake +``` + +Mooncake reads its connection settings from the environment. Export them before +starting Ray so that Ray workers inherit the node-local values: + +```bash +export MOONCAKE_MASTER=":50051" +export MOONCAKE_PROTOCOL="" +export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE" +export MOONCAKE_GLOBAL_SEGMENT_SIZE="2gb" +export MOONCAKE_LOCAL_BUFFER_SIZE="2gb" + +# Optional when the Ray node IP is already the desired data-network address. +# export MOONCAKE_LOCAL_HOSTNAME="" + +# RDMA only. Set the device attached to this node's data-network address. +# export MOONCAKE_DEVICE="" +``` + +`MOONCAKE_LOCAL_HOSTNAME` and `MOONCAKE_DEVICE` are node-local. Do not put one +node's values into a cluster-wide Ray runtime environment. + +## Two-node walkthrough + +The following example runs two synchronous rollout and training iterations with +Qwen3-4B. It uses one eight-GPU node for training and one eight-GPU node for rollout. +Set the variables for your cluster, then complete the steps in order. + +Both nodes must use the same Python environment, slime revision, and Mooncake wheel. + +### 1. Choose the protocol and set node-local values + +On the head node: + +```bash +export HEAD_IP="" +export MOONCAKE_MASTER="${HEAD_IP}:50051" +export MOONCAKE_PROTOCOL="" +export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE" +export MOONCAKE_GLOBAL_SEGMENT_SIZE="2gb" +export MOONCAKE_LOCAL_BUFFER_SIZE="2gb" +export MOONCAKE_LOCAL_HOSTNAME="${HEAD_IP}" + +# RDMA only. +# export MOONCAKE_DEVICE="" +``` + +On the worker node: + +```bash +export HEAD_IP="" +export WORKER_IP="" +export MOONCAKE_MASTER="${HEAD_IP}:50051" +export MOONCAKE_PROTOCOL="" +export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE" +export MOONCAKE_GLOBAL_SEGMENT_SIZE="2gb" +export MOONCAKE_LOCAL_BUFFER_SIZE="2gb" +export MOONCAKE_LOCAL_HOSTNAME="${WORKER_IP}" + +# RDMA only. +# export MOONCAKE_DEVICE="" +``` + +The 2 GiB settings keep this small example easy to run. For a production job, +size both values for the rollout payloads and the number of partitions that may +remain in flight at the same time. + +### 2. Start Mooncake and Ray on the head + +Activate the slime environment, then run: + +```bash +mooncake_master --rpc_address=0.0.0.0 --rpc_port=50051 \ + >mooncake_master.log 2>&1 & + +ray stop --force +ray start --head \ + --node-ip-address="${HEAD_IP}" \ + --port=6379 \ + --num-gpus=8 \ + --disable-usage-stats \ + --dashboard-host=0.0.0.0 \ + --dashboard-port=8265 +``` + +An existing Mooncake master can be replaced with a configured HA endpoint; in that +case, do not launch another local master. + +### 3. Join the worker + +Activate the same slime environment on the worker, then run: + +```bash +ray stop --force +ray start \ + --address="${HEAD_IP}:6379" \ + --node-ip-address="${WORKER_IP}" \ + --num-gpus=8 \ + --disable-usage-stats +``` + +Check `ray status` on the head before submitting training. The cluster should report +16 GPUs. + +### 4. Submit training from the head + +Set paths that are valid on the nodes where the corresponding workers run: + +```bash +export SLIME_HOME="" +export MEGATRON_HOME="" +export HF_CHECKPOINT="" +export REF_LOAD="" +export PROMPT_DATA="" +export RAY_DASHBOARD_ADDR="http://127.0.0.1:8265" + +cd "${SLIME_HOME}" +source scripts/models/qwen3-4B.sh + +RUNTIME_ENV_JSON="{ + \"env_vars\": { + \"PYTHONPATH\": \"${MEGATRON_HOME}\", + \"CUDA_DEVICE_MAX_CONNECTIONS\": \"1\", + \"NCCL_NVLS_ENABLE\": \"0\", + \"MOONCAKE_MASTER\": \"${MOONCAKE_MASTER}\", + \"MOONCAKE_PROTOCOL\": \"${MOONCAKE_PROTOCOL}\", + \"MOONCAKE_TE_META_DATA_SERVER\": \"${MOONCAKE_TE_META_DATA_SERVER}\", + \"MOONCAKE_GLOBAL_SEGMENT_SIZE\": \"${MOONCAKE_GLOBAL_SEGMENT_SIZE}\", + \"MOONCAKE_LOCAL_BUFFER_SIZE\": \"${MOONCAKE_LOCAL_BUFFER_SIZE}\" + } +}" + +ray job submit --address="${RAY_DASHBOARD_ADDR}" \ + --working-dir="${SLIME_HOME}" \ + --runtime-env-json="${RUNTIME_ENV_JSON}" \ + -- python3 train.py \ + --actor-num-nodes 1 \ + --actor-num-gpus-per-node 8 \ + --rollout-num-gpus 8 \ + "${MODEL_ARGS[@]}" \ + --hf-checkpoint "${HF_CHECKPOINT}" \ + --ref-load "${REF_LOAD}" \ + --load "${REF_LOAD}" \ + --prompt-data "${PROMPT_DATA}" \ + --input-key prompt \ + --label-key label \ + --apply-chat-template \ + --rollout-shuffle \ + --rm-type deepscaler \ + --start-rollout-id 0 \ + --num-rollout 2 \ + --rollout-batch-size 2 \ + --n-samples-per-prompt 2 \ + --rollout-max-response-len 256 \ + --rollout-temperature 1 \ + --global-batch-size 4 \ + --balance-data \ + --rollout-data-transport mooncake \ + --tensor-model-parallel-size 8 \ + --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 1024 \ + --advantage-estimator grpo \ + --use-kl-loss \ + --kl-loss-coef 0.0 \ + --kl-loss-type low_var_kl \ + --entropy-coef 0.0 \ + --eps-clip 0.2 \ + --eps-clip-high 0.28 \ + --optimizer adam \ + --lr 1e-6 \ + --lr-decay-style constant \ + --weight-decay 0.1 \ + --adam-beta1 0.9 \ + --adam-beta2 0.98 \ + --rollout-num-gpus-per-engine 8 \ + --sglang-mem-fraction-static 0.35 \ + --attention-dropout 0.0 \ + --hidden-dropout 0.0 \ + --accumulate-allreduce-grads-in-fp32 \ + --attention-softmax-in-fp32 \ + --attention-backend flash \ + --no-gradient-accumulation-fusion \ + --bf16 \ + --distributed-backend nccl +``` + +For fully asynchronous training, keep the same Mooncake environment and transport +option, then use the regular async entrypoint and settings: + +```diff +- -- python3 train.py ... ++ -- python3 train_async.py ... +``` + +See the [fully asynchronous example](../_examples_synced/fully_async/README.md) for +the remaining async arguments. + +## Configuration reference + +| Setting | Default | Purpose | +|---|---|---| +| `--rollout-data-transport` | `object-store` | Set to `mooncake` to use Mooncake for rollout data. | +| `MOONCAKE_MASTER` | none | Address of `mooncake_master` or the HA metadata endpoint. | +| `MOONCAKE_LOCAL_HOSTNAME` | Ray node IP | Data-network address advertised by the local client. | +| `MOONCAKE_TE_META_DATA_SERVER` | `P2PHANDSHAKE` | Transfer Engine metadata service. | +| `MOONCAKE_PROTOCOL` | `rdma` | Transfer protocol, normally `tcp` or `rdma`. | +| `MOONCAKE_DEVICE` | auto-discovery | Local RDMA device name. | +| `MOONCAKE_GLOBAL_SEGMENT_SIZE` | `8gb` | Store capacity contributed by rollout writers. | +| `MOONCAKE_LOCAL_BUFFER_SIZE` | `32gb` | Local transfer and staging capacity. | + +The trainer-side GET clients do not contribute a global Store segment. slime still +uses `MOONCAKE_LOCAL_BUFFER_SIZE` there for transfer and pool-backed results, and +releases those buffers after the training step consumes them. + +## Troubleshooting + +- **Mooncake import fails during argument parsing:** install a compatible Mooncake + wheel in the Python environment used by every Ray worker. +- **Store setup fails:** verify `MOONCAKE_MASTER`, endpoint reachability, and that all + nodes use the same Mooncake version. +- **RDMA setup fails:** verify `MOONCAKE_DEVICE`, locked-memory limits, device access, + and that `MOONCAKE_LOCAL_HOSTNAME` belongs to the selected RDMA network. +- **Allocation fails:** increase available host memory or lower the two Mooncake size + settings. Account for concurrent rollout partitions and in-flight steps. diff --git a/docs/en/index.rst b/docs/en/index.rst index 83230af1f..0b6002074 100644 --- a/docs/en/index.rst +++ b/docs/en/index.rst @@ -84,6 +84,7 @@ Start by Use Case advanced/pd-disaggregation.md advanced/external-rollout-engines.md advanced/delta-weight-sync.md + advanced/mooncake-rollout-transfer.md advanced/sglang-config.md advanced/megatron-config.md advanced/arch-support-beyond-megatron.md diff --git a/docs/zh/advanced/mooncake-rollout-transfer.md b/docs/zh/advanced/mooncake-rollout-transfer.md new file mode 100644 index 000000000..0729a2397 --- /dev/null +++ b/docs/zh/advanced/mooncake-rollout-transfer.md @@ -0,0 +1,265 @@ +# 使用 Mooncake 传输 Rollout Data + +slime 默认通过 Ray object store 将 rollout data 从 rollout manager 传给 trainer。 +Mooncake Store 可以替换这段传输路径:rollout dict 的接口和训练语义保持不变, +数据 payload 则通过 Mooncake data plane 传输。 + +这个配置只影响 rollout data: + +- `--rollout-data-transport` 选择 rollout data 到达 trainer 的方式; +- 模型权重同步由 `--update-weight-*` 参数单独配置; +- SGLang KV cache disaggregation 也使用独立的传输配置。 + +同步入口 `train.py` 和异步入口 `train_async.py` 都支持 Mooncake。 + +## 环境要求 + +启动 slime 任务前,请确认: + +- 所有 Ray 节点使用相同的 slime revision 和 Mooncake wheel。Mooncake 包需要 + 提供 slime 使用的 structured-object `put`/`get` 接口; +- 已经启动 `mooncake_master`,或者准备好 Mooncake HA endpoint。slime 只作为 + client 连接,不负责管理 endpoint 生命周期; +- 所有 Ray 节点都能访问用于数据传输的网络地址; +- 为 `MOONCAKE_GLOBAL_SEGMENT_SIZE` 和 `MOONCAKE_LOCAL_BUFFER_SIZE` + 预留足够的主机内存; +- 使用 RDMA 时,运行环境能够访问 RDMA device、允许锁定内存,并且在启动 + Ray 前配置好每个节点自己的 device name; +- 模型、Megatron checkpoint、数据集和 Python 环境在使用它们的节点上均可用。 + +Mooncake 的包和平台要求请参考 +[安装文档](https://kvcache-ai.github.io/Mooncake/getting_started/build.html)。 + +## 配置传输后端 + +启动 Ray 前先选择 TCP 或 RDMA。TCP 只要求数据网络可达;RDMA 还需要每个 +节点配置本地 device,不同节点的 device name 可以不同。 + +在已有 slime recipe 中增加: + +```bash +--rollout-data-transport mooncake +``` + +Mooncake 从环境变量中读取连接配置。请在启动 Ray 前导出这些变量,使 Ray +worker 继承节点本地配置: + +```bash +export MOONCAKE_MASTER=":50051" +export MOONCAKE_PROTOCOL="" +export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE" +export MOONCAKE_GLOBAL_SEGMENT_SIZE="2gb" +export MOONCAKE_LOCAL_BUFFER_SIZE="2gb" + +# Ray node IP 已经是数据网络地址时可以省略。 +# export MOONCAKE_LOCAL_HOSTNAME="" + +# 仅 RDMA 需要。填写当前节点数据网卡对应的 device。 +# export MOONCAKE_DEVICE="" +``` + +`MOONCAKE_LOCAL_HOSTNAME` 和 `MOONCAKE_DEVICE` 是节点本地配置。不要把 +某一个节点的值写入整个集群共用的 Ray runtime environment。 + +## 双机运行示例 + +下面的示例使用 Qwen3-4B 完成两轮同步 rollout 和训练。一个八卡节点用于 +训练,另一个八卡节点用于 rollout。请先根据集群环境设置变量,再按顺序执行。 + +两个节点必须使用相同的 Python 环境、slime revision 和 Mooncake wheel。 + +### 1. 选择协议并配置节点变量 + +在 head 节点执行: + +```bash +export HEAD_IP="" +export MOONCAKE_MASTER="${HEAD_IP}:50051" +export MOONCAKE_PROTOCOL="" +export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE" +export MOONCAKE_GLOBAL_SEGMENT_SIZE="2gb" +export MOONCAKE_LOCAL_BUFFER_SIZE="2gb" +export MOONCAKE_LOCAL_HOSTNAME="${HEAD_IP}" + +# 仅 RDMA 需要。 +# export MOONCAKE_DEVICE="" +``` + +在 worker 节点执行: + +```bash +export HEAD_IP="" +export WORKER_IP="" +export MOONCAKE_MASTER="${HEAD_IP}:50051" +export MOONCAKE_PROTOCOL="" +export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE" +export MOONCAKE_GLOBAL_SEGMENT_SIZE="2gb" +export MOONCAKE_LOCAL_BUFFER_SIZE="2gb" +export MOONCAKE_LOCAL_HOSTNAME="${WORKER_IP}" + +# 仅 RDMA 需要。 +# export MOONCAKE_DEVICE="" +``` + +这里使用 2 GiB 是为了方便运行这个小规模示例。生产任务需要根据 rollout +payload 大小和同时处于 in-flight 状态的 partition 数量重新评估这两个值。 + +### 2. 在 head 节点启动 Mooncake 和 Ray + +激活 slime Python 环境后执行: + +```bash +mooncake_master --rpc_address=0.0.0.0 --rpc_port=50051 \ + >mooncake_master.log 2>&1 & + +ray stop --force +ray start --head \ + --node-ip-address="${HEAD_IP}" \ + --port=6379 \ + --num-gpus=8 \ + --disable-usage-stats \ + --dashboard-host=0.0.0.0 \ + --dashboard-port=8265 +``` + +如果使用已有的 Mooncake master 或 HA endpoint,不要再启动本地 master。 + +### 3. 将 worker 加入 Ray 集群 + +在 worker 节点激活相同的 Python 环境,然后执行: + +```bash +ray stop --force +ray start \ + --address="${HEAD_IP}:6379" \ + --node-ip-address="${WORKER_IP}" \ + --num-gpus=8 \ + --disable-usage-stats +``` + +提交训练前,在 head 节点运行 `ray status`,确认集群中有 16 张 GPU。 + +### 4. 从 head 节点提交训练 + +以下路径必须在实际使用它们的节点上有效: + +```bash +export SLIME_HOME="" +export MEGATRON_HOME="" +export HF_CHECKPOINT="" +export REF_LOAD="" +export PROMPT_DATA="" +export RAY_DASHBOARD_ADDR="http://127.0.0.1:8265" + +cd "${SLIME_HOME}" +source scripts/models/qwen3-4B.sh + +RUNTIME_ENV_JSON="{ + \"env_vars\": { + \"PYTHONPATH\": \"${MEGATRON_HOME}\", + \"CUDA_DEVICE_MAX_CONNECTIONS\": \"1\", + \"NCCL_NVLS_ENABLE\": \"0\", + \"MOONCAKE_MASTER\": \"${MOONCAKE_MASTER}\", + \"MOONCAKE_PROTOCOL\": \"${MOONCAKE_PROTOCOL}\", + \"MOONCAKE_TE_META_DATA_SERVER\": \"${MOONCAKE_TE_META_DATA_SERVER}\", + \"MOONCAKE_GLOBAL_SEGMENT_SIZE\": \"${MOONCAKE_GLOBAL_SEGMENT_SIZE}\", + \"MOONCAKE_LOCAL_BUFFER_SIZE\": \"${MOONCAKE_LOCAL_BUFFER_SIZE}\" + } +}" + +ray job submit --address="${RAY_DASHBOARD_ADDR}" \ + --working-dir="${SLIME_HOME}" \ + --runtime-env-json="${RUNTIME_ENV_JSON}" \ + -- python3 train.py \ + --actor-num-nodes 1 \ + --actor-num-gpus-per-node 8 \ + --rollout-num-gpus 8 \ + "${MODEL_ARGS[@]}" \ + --hf-checkpoint "${HF_CHECKPOINT}" \ + --ref-load "${REF_LOAD}" \ + --load "${REF_LOAD}" \ + --prompt-data "${PROMPT_DATA}" \ + --input-key prompt \ + --label-key label \ + --apply-chat-template \ + --rollout-shuffle \ + --rm-type deepscaler \ + --start-rollout-id 0 \ + --num-rollout 2 \ + --rollout-batch-size 2 \ + --n-samples-per-prompt 2 \ + --rollout-max-response-len 256 \ + --rollout-temperature 1 \ + --global-batch-size 4 \ + --balance-data \ + --rollout-data-transport mooncake \ + --tensor-model-parallel-size 8 \ + --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 1024 \ + --advantage-estimator grpo \ + --use-kl-loss \ + --kl-loss-coef 0.0 \ + --kl-loss-type low_var_kl \ + --entropy-coef 0.0 \ + --eps-clip 0.2 \ + --eps-clip-high 0.28 \ + --optimizer adam \ + --lr 1e-6 \ + --lr-decay-style constant \ + --weight-decay 0.1 \ + --adam-beta1 0.9 \ + --adam-beta2 0.98 \ + --rollout-num-gpus-per-engine 8 \ + --sglang-mem-fraction-static 0.35 \ + --attention-dropout 0.0 \ + --hidden-dropout 0.0 \ + --accumulate-allreduce-grads-in-fp32 \ + --attention-softmax-in-fp32 \ + --attention-backend flash \ + --no-gradient-accumulation-fusion \ + --bf16 \ + --distributed-backend nccl +``` + +使用 fully async 训练时,保留相同的 Mooncake 环境变量和传输参数,切换到 +slime 原有的异步入口和配置: + +```diff +- -- python3 train.py ... ++ -- python3 train_async.py ... +``` + +其余异步参数请参考 [fully async 示例](../_examples_synced/fully_async/README.md)。 + +## 配置项说明 + +| 配置 | 默认值 | 作用 | +|---|---|---| +| `--rollout-data-transport` | `object-store` | 设为 `mooncake` 后使用 Mooncake 传输 rollout data。 | +| `MOONCAKE_MASTER` | 无 | `mooncake_master` 或 HA metadata endpoint 的地址。 | +| `MOONCAKE_LOCAL_HOSTNAME` | Ray node IP | 当前 client 对外发布的数据网络地址。 | +| `MOONCAKE_TE_META_DATA_SERVER` | `P2PHANDSHAKE` | Transfer Engine metadata service。 | +| `MOONCAKE_PROTOCOL` | `rdma` | 传输协议,通常使用 `tcp` 或 `rdma`。 | +| `MOONCAKE_DEVICE` | 自动发现 | 当前节点的 RDMA device name。 | +| `MOONCAKE_GLOBAL_SEGMENT_SIZE` | `8gb` | rollout writer 贡献的 Store 容量。 | +| `MOONCAKE_LOCAL_BUFFER_SIZE` | `32gb` | 本地传输和 staging buffer 容量。 | + +trainer 侧的 GET client 不会贡献 global Store segment,但仍会使用 +`MOONCAKE_LOCAL_BUFFER_SIZE` 完成传输和管理 pool-backed result。训练消费完 +数据后,slime 会释放这些 buffer。 + +## 常见问题 + +- **参数解析时无法导入 Mooncake:**所有 Ray worker 的 Python 环境都需要安装 + 兼容的 Mooncake wheel; +- **Store setup 失败:**检查 `MOONCAKE_MASTER`、endpoint 可达性,以及所有 + 节点使用的 Mooncake 版本是否一致; +- **RDMA 初始化失败:**检查 `MOONCAKE_DEVICE`、locked-memory limit、device + 权限,以及 `MOONCAKE_LOCAL_HOSTNAME` 是否属于所选 RDMA 网络; +- **内存分配失败:**释放主机内存或下调两个 Mooncake size 参数,并为并发的 + rollout partition 和 in-flight step 预留容量。 diff --git a/docs/zh/index.rst b/docs/zh/index.rst index 747deddf6..a5763e741 100644 --- a/docs/zh/index.rst +++ b/docs/zh/index.rst @@ -84,6 +84,7 @@ slime 的设计目标,是让这两大能力彼此强化,同时避免把系 advanced/pd-disaggregation.md advanced/external-rollout-engines.md advanced/delta-weight-sync.md + advanced/mooncake-rollout-transfer.md advanced/sglang-config.md advanced/megatron-config.md advanced/arch-support-beyond-megatron.md diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index ea2601a9d..d181a8396 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -387,6 +387,11 @@ def train(self, rollout_id: int, rollout_data_ref: Box, external_data=None): self.train_actor(rollout_id, rollout_data, external_data=external_data) result = None + if getattr(self.args, "rollout_data_transport", "object-store") == "mooncake": + from slime.utils.data_transfer import release_mooncake_rollout_data + + release_mooncake_rollout_data(self.args, rollout_data) + if self.args.offload_train: del rollout_data self.sleep() diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py index abd8e35d7..cac5b98e9 100644 --- a/slime/ray/rollout.py +++ b/slime/ray/rollout.py @@ -929,7 +929,11 @@ def _split_train_data_by_dp(self, data): rollout_data["micro_batch_indices"] = micro_batch_indices[r] _tensorize_rollout_data_for_training(rollout_data) transport = getattr(self.args, "rollout_data_transport", "object-store") - if transport == "nixl": + if transport == "mooncake": + from slime.utils.data_transfer import put_mooncake_rollout_data + + rollout_data_refs.append(put_mooncake_rollout_data(self.args, rollout_data, partition=f"dp{r}")) + elif transport == "nixl": rollout_data_refs.append(Box(ray.put(rollout_data, _tensor_transport="nixl"))) elif transport == "object-store": rollout_data_refs.append(Box(ray.put(rollout_data))) diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 96be54457..40e00a063 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -557,12 +557,13 @@ def add_rollout_arguments(parser): parser.add_argument( "--rollout-data-transport", type=str, - choices=["object-store", "nixl"], + choices=["object-store", "nixl", "mooncake"], default="object-store", help=( "Transport for rollout data refs sent from rollout manager to trainer. Large rollout " "fields are tensorized on CPU before the refs are stored. Set to nixl to transfer " - "those torch tensors via Ray NIXL." + "those torch tensors via Ray NIXL, or mooncake to transfer the rollout dictionary " + "through Mooncake Store." ), ) parser.add_argument( @@ -2065,3 +2066,8 @@ def slime_validate_args(args): "--update-weight-mode=delta requires --update-weight-local-checkpoint-dir " "(a rollout-host-local NVMe directory)." ) + + if getattr(args, "rollout_data_transport", "object-store") == "mooncake": + from slime.utils.data_transfer import check_mooncake_available + + check_mooncake_available() diff --git a/slime/utils/data.py b/slime/utils/data.py index adac4b95a..ae245170f 100644 --- a/slime/utils/data.py +++ b/slime/utils/data.py @@ -304,7 +304,13 @@ 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) + ref = rollout_data_ref[dp_rank] + if getattr(args, "rollout_data_transport", "object-store") == "mooncake": + from slime.utils.data_transfer import get_mooncake_rollout_data + + rollout_data = get_mooncake_rollout_data(args, ref) + else: + rollout_data = ray.get(ref.inner) # Keep `partition` in rollout_data: each local sample's position in the # flattened rollout batch (== its index in the rollout debug dump's diff --git a/slime/utils/data_transfer.py b/slime/utils/data_transfer.py new file mode 100644 index 000000000..c57dd6f3c --- /dev/null +++ b/slime/utils/data_transfer.py @@ -0,0 +1,159 @@ +"""Thin Mooncake rollout data transport: put/get/cleanup.""" + +import logging +import os +from functools import cache +from typing import Any + +from slime.utils.misc import Box + +try: + from mooncake.store import MooncakeDistributedStore + from mooncake.structured_object_store import FieldSchema, MooncakeBundleTransfer, export_ref, import_ref + + _MOONCAKE_AVAILABLE = True +except ImportError: + _MOONCAKE_AVAILABLE = False + +logger = logging.getLogger(__name__) + +_ROLLOUT_FIELD_SCHEMA_SPECS = { + # rollout.py tensorizes these row-aligned fields before transport. + "tokens": ("ragged_tensor", None, "non_tensor_batch"), + "loss_masks": ("ragged_tensor", None, "non_tensor_batch"), + "rollout_log_probs": ("ragged_tensor", None, "non_tensor_batch"), + "rollout_top_p_token_ids": ("ragged_tensor", None, "non_tensor_batch"), + "rollout_top_p_token_offsets": ("ragged_tensor", None, "non_tensor_batch"), + "teacher_log_probs": ("ragged_tensor", None, "non_tensor_batch"), + "rollout_routed_experts": ("ragged_tensor", None, "non_tensor_batch"), + # Row-aligned scalar fields. + "partition": ("ndarray", "int64", "non_tensor_batch"), + "response_lengths": ("ndarray", "int64", "non_tensor_batch"), + "rewards": ("ndarray", "float32", "non_tensor_batch"), + "truncated": ("ndarray", "int64", "non_tensor_batch"), + "round_number": ("ndarray", "int64", "non_tensor_batch"), + "sample_indices": ("ndarray", "int64", "non_tensor_batch"), + "rollout_ids": ("ndarray", "int64", "non_tensor_batch"), + "rollout_mask_sums": ("tensor", "float32", "batch"), + # Optional row-aligned text fields. + "prompt": ("utf8_ragged", None, "non_tensor_batch"), + # Metadata fields carried with each DP partition. + "raw_reward": ("auto", None, "meta_info"), + "total_lengths": ("auto", None, "meta_info"), + "global_batch_sizes": ("auto", None, "meta_info"), + "num_microbatches": ("auto", None, "meta_info"), + "micro_batch_indices": ("auto", None, "meta_info"), +} + +_ROLLOUT_FIELD_SCHEMAS = ( + { + key: FieldSchema( + codec=codec, + nullable=False, + metadata={"section": section, **({"dtype": dtype} if dtype else {})}, + ) + for key, (codec, dtype, section) in _ROLLOUT_FIELD_SCHEMA_SPECS.items() + } + if _MOONCAKE_AVAILABLE + else {} +) + + +def check_mooncake_available() -> None: + """Call during argument parsing to fail fast if mooncake is not installed.""" + if not _MOONCAKE_AVAILABLE: + raise ImportError( + "rollout-data-transport='mooncake' requires the mooncake package. " "Install it with: pip install mooncake" + ) + + +def put_mooncake_rollout_data(args: Any, data: dict[str, Any], partition: str) -> Box: + ref = _mooncake_transfer(args, contribute_segment=True).put( + data, + type="dict", + namespace="slime", + partition=partition, + stage="rollout", + field_schemas=_rollout_field_schemas_for_data(data), + ) + return Box(export_ref(ref)) + + +def _rollout_field_schemas_for_data(data: dict[str, Any]) -> dict: + return {key: schema for key, schema in _ROLLOUT_FIELD_SCHEMAS.items() if key in data} + + +def get_mooncake_rollout_data(args: Any, ref: Box) -> dict[str, Any]: + return _mooncake_transfer(args, contribute_segment=False).get(import_ref(ref.inner), type="dict") + + +def release_mooncake_rollout_data(args: Any, data: dict[str, Any]) -> None: + """Release pool-backed buffers after training has fully consumed the data.""" + from mooncake.structured_object_store import MooncakeBundleTransfer + + MooncakeBundleTransfer.release_result(data) + + +def cleanup_mooncake_rollout_data(args: Any, ref: Box) -> None: + _mooncake_transfer(args, contribute_segment=False).cleanup_dataproto(import_ref(ref.inner)) + + +def cleanup_mooncake_rollout_refs(args: Any, refs: list[Box]) -> None: + for ref in refs: + cleanup_mooncake_rollout_data(args, ref) + + +def _mooncake_transfer(args: Any, contribute_segment: bool): + config = _mooncake_store_config(args, contribute_segment=contribute_segment) + return _cached_mooncake_transfer(tuple(sorted(config.items()))) + + +@cache +def _cached_mooncake_transfer(config_items: tuple[tuple[str, Any], ...]): + store = MooncakeDistributedStore() + ret = store.setup(dict(config_items)) + if ret: + raise RuntimeError(f"Mooncake store setup failed: {ret}") + return MooncakeBundleTransfer(store, key_prefix="slime-rollout") + + +def _mooncake_store_config(args: Any, contribute_segment: bool) -> dict[str, Any]: + mc_kwargs = getattr(args, "mooncake_store_init_kwargs", None) or {} + global_segment_size = _parse_size( + mc_kwargs.get("global_segment_size") or os.getenv("MOONCAKE_GLOBAL_SEGMENT_SIZE", str(8 * 1024**3)) + ) + if not contribute_segment: + global_segment_size = 0 + return { + "local_hostname": str(mc_kwargs.get("local_hostname") or _local_hostname()), + "metadata_server": str( + mc_kwargs.get("metadata_server") or os.getenv("MOONCAKE_TE_META_DATA_SERVER", "P2PHANDSHAKE") + ), + "global_segment_size": global_segment_size, + "local_buffer_size": _parse_size( + mc_kwargs.get("local_buffer_size") or os.getenv("MOONCAKE_LOCAL_BUFFER_SIZE", str(32 * 1024**3)) + ), + "protocol": str(mc_kwargs.get("protocol") or os.getenv("MOONCAKE_PROTOCOL", "rdma")), + "rdma_devices": str(mc_kwargs.get("device_name") or os.getenv("MOONCAKE_DEVICE", "")), + "master_server_addr": str(mc_kwargs.get("master_server_address") or os.getenv("MOONCAKE_MASTER", "")), + } + + +def _parse_size(value: Any) -> int: + if isinstance(value, int): + return value + text = str(value).strip().lower() + units = {"kb": 1024, "mb": 1024**2, "gb": 1024**3, "k": 1024, "m": 1024**2, "g": 1024**3} + for suffix, multiplier in units.items(): + if text.endswith(suffix): + return int(float(text[: -len(suffix)]) * multiplier) + return int(text) + + +def _local_hostname() -> str: + value = os.getenv("MOONCAKE_LOCAL_HOSTNAME") or os.getenv("LOCAL_HOSTNAME") + if value: + return value + import ray + + return ray.util.get_node_ip_address() diff --git a/train.py b/train.py index 9cb296886..258f19bdf 100644 --- a/train.py +++ b/train.py @@ -2,6 +2,7 @@ from slime.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models from slime.utils.arguments import parse_args +from slime.utils.data_transfer import cleanup_mooncake_rollout_refs from slime.utils.logging_utils import configure_logger, finish_tracking, init_tracking from slime.utils.misc import should_run_periodic_action @@ -68,6 +69,9 @@ def offload_train(actor_trains_this_step): else: ray.get(actor_model.async_train(rollout_id, rollout_data_ref)) + if getattr(args, "rollout_data_transport", "object-store") == "mooncake": + cleanup_mooncake_rollout_refs(args, rollout_data_ref) + if release_train or should_run_periodic_action( rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout ): diff --git a/train_async.py b/train_async.py index 9141b612f..d511d1e67 100644 --- a/train_async.py +++ b/train_async.py @@ -2,6 +2,7 @@ from slime.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models from slime.utils.arguments import parse_args +from slime.utils.data_transfer import cleanup_mooncake_rollout_refs from slime.utils.logging_utils import configure_logger, finish_tracking, init_tracking from slime.utils.misc import should_run_periodic_action @@ -52,6 +53,9 @@ def train(args): else: ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref)) + if getattr(args, "rollout_data_transport", "object-store") == "mooncake": + cleanup_mooncake_rollout_refs(args, rollout_data_curr_ref) + if release_train or should_run_periodic_action( rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout ):