diff --git a/docs/en/advanced/mooncake-store-transfer.md b/docs/en/advanced/mooncake-store-transfer.md new file mode 100644 index 000000000..b7fe1c663 --- /dev/null +++ b/docs/en/advanced/mooncake-store-transfer.md @@ -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). diff --git a/docs/en/index.rst b/docs/en/index.rst index a564c52ee..6b12d6fbe 100644 --- a/docs/en/index.rst +++ b/docs/en/index.rst @@ -45,6 +45,7 @@ vime is built on `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 diff --git a/tests/run_disaggregate_2gpu.py b/tests/run_disaggregate_2gpu.py new file mode 100644 index 000000000..ff9fd1e35 --- /dev/null +++ b/tests/run_disaggregate_2gpu.py @@ -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() diff --git a/tests/utils/test_mooncake_store_transfer.py b/tests/utils/test_mooncake_store_transfer.py new file mode 100644 index 000000000..4ba99c540 --- /dev/null +++ b/tests/utils/test_mooncake_store_transfer.py @@ -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" diff --git a/train.py b/train.py index d9f9b2af9..bc0beda31 100644 --- a/train.py +++ b/train.py @@ -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 @@ -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) diff --git a/train_async.py b/train_async.py index da191396d..99531d282 100644 --- a/train_async.py +++ b/train_async.py @@ -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 @@ -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: diff --git a/vime/ray/rollout.py b/vime/ray/rollout.py index 7499cda19..06d252d97 100644 --- a/vime/ray/rollout.py +++ b/vime/ray/rollout.py @@ -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): diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index 704aeebe5..678e945fa 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -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( @@ -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.") diff --git a/vime/utils/data.py b/vime/utils/data.py index a2b8c50c4..d44977346 100644 --- a/vime/utils/data.py +++ b/vime/utils/data.py @@ -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"] diff --git a/vime/utils/external_utils/command_utils.py b/vime/utils/external_utils/command_utils.py index 4b73c4666..3480aded4 100644 --- a/vime/utils/external_utils/command_utils.py +++ b/vime/utils/external_utils/command_utils.py @@ -6,6 +6,7 @@ import json import os import random +import sys import time from dataclasses import dataclass from pathlib import Path @@ -18,6 +19,11 @@ repo_base_dir = Path(os.path.abspath(__file__)).resolve().parents[3] +def _ray_cli() -> str: + """Use ``python -m ray`` so broken venv entrypoint shebangs do not break CI scripts.""" + return f"{sys.executable} -m ray.scripts.scripts" + + def convert_checkpoint( model_name, megatron_model_type, @@ -73,6 +79,20 @@ def hf_download_dataset(full_name: str): exec_command(f"hf download --repo-type dataset {full_name} --local-dir /root/datasets/{partial_name}") +def ms_download_model(model_id: str, local_dir: str): + exec_command( + "python -c " + f"\"from modelscope import snapshot_download; snapshot_download('{model_id}', local_dir='{local_dir}')\"" + ) + + +def ms_download_dataset(dataset_id: str, local_dir: str | None = None): + partial_name = dataset_id.split("/")[-1] + local_dir = local_dir or f"/root/datasets/{partial_name}" + exec_command(f"mkdir -p {local_dir}") + exec_command(f"modelscope download --dataset {dataset_id} --local_dir {local_dir}") + + def fp8_cast_bf16(path_src, path_dst): if Path(path_dst).exists(): print(f"fp8_cast_bf16 skip {path_dst} since exists") @@ -110,7 +130,7 @@ def execute_train( exec_command( "pkill -9 -f '[v]llm serve|VLL[M]::'; " "sleep 3; " - f"{'' if external_ray else 'ray stop --force; '}" + f"{'' if external_ray else f'{_ray_cli()} stop --force; '}" f"{'' if external_ray else 'pkill -9 ray; '}" # cannot be run in CI, o/w kill the parent script # TODO: do we really need this kill? (or can we instead kill vime) @@ -128,7 +148,7 @@ def execute_train( exec_command( # will prevent ray from buffering stdout/stderr f"export PYTHONUNBUFFERED=1 && " - f"ray start --head --node-ip-address {master_addr} --num-gpus {num_gpus_per_node} --disable-usage-stats" + f"{_ray_cli()} start --head --node-ip-address {master_addr} --num-gpus {num_gpus_per_node} --port 6380 --disable-usage-stats" ) if (f := before_ray_job_submit) is not None: @@ -169,9 +189,9 @@ def execute_train( exec_command( f"export no_proxy=127.0.0.1 && export PYTHONUNBUFFERED=1 && " f"{cmd_megatron_model_source}" - f'ray job submit --address="http://127.0.0.1:8265" ' + f'{_ray_cli()} job submit --address="http://127.0.0.1:8265" ' f"--runtime-env-json='{runtime_env_json}' " - f"-- python3 {train_script} " + f"-- {sys.executable} {train_script} " f"{'${MODEL_ARGS[@]}' if megatron_model_type is not None else ''} " f"{train_args}" ) diff --git a/vime/utils/mooncake_store_service.py b/vime/utils/mooncake_store_service.py new file mode 100644 index 000000000..393e0e70e --- /dev/null +++ b/vime/utils/mooncake_store_service.py @@ -0,0 +1,43 @@ +"""Start mooncake_master for local smoke tests.""" + +from __future__ import annotations + +import shutil +import signal +import socket +import subprocess +import time + + +def ensure_mooncake_master(host: str = "127.0.0.1", rpc_port: int = 50051, metadata_port: int = 18080) -> None: + if _port_open(host, rpc_port): + return + mooncake_master = shutil.which("mooncake_master") + if mooncake_master is None: + raise RuntimeError("mooncake_master not found; pip install mooncake-transfer-engine") + proc = subprocess.Popen( + [ + mooncake_master, + "--enable_http_metadata_server=true", + f"--http_metadata_server_host={host}", + f"--http_metadata_server_port={metadata_port}", + ], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + start_new_session=True, + ) + deadline = time.time() + 15.0 + while time.time() < deadline: + if proc.poll() is not None: + raise RuntimeError("mooncake_master exited early") + if _port_open(host, rpc_port) and _port_open(host, metadata_port): + return + time.sleep(0.2) + proc.send_signal(signal.SIGTERM) + raise RuntimeError("mooncake_master did not become ready") + + +def _port_open(host: str, port: int) -> bool: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.settimeout(0.5) + return sock.connect_ex((host, port)) == 0 diff --git a/vime/utils/remote_batch.py b/vime/utils/remote_batch.py new file mode 100644 index 000000000..058f59e4c --- /dev/null +++ b/vime/utils/remote_batch.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import ctypes +import os +import re +from dataclasses import dataclass, field +from typing import Any + +import numpy as np +import torch + +_STORE_CACHE: dict[tuple[tuple[str, str], ...], Any] = {} +_FIELD_NAME_RE = re.compile(r"^[A-Za-z0-9_.-]{1,128}$") + + +def default_store_init_kwargs() -> dict[str, Any]: + return { + "local_hostname": os.getenv("MOONCAKE_LOCAL_HOSTNAME", "127.0.0.1"), + "metadata_server": os.getenv("MOONCAKE_TE_META_DATA_SERVER", "http://127.0.0.1:18080/metadata"), + "global_segment_size": int(os.getenv("MOONCAKE_GLOBAL_SEGMENT_SIZE", str(512 * 1024 * 1024))), + "local_buffer_size": int(os.getenv("MOONCAKE_LOCAL_BUFFER_SIZE", str(128 * 1024 * 1024))), + "protocol": os.getenv("MOONCAKE_PROTOCOL", "tcp"), + "rdma_devices": os.getenv("MOONCAKE_DEVICE", ""), + "master_server_addr": os.getenv("MOONCAKE_MASTER", "127.0.0.1:50051"), + } + + +def normalize_store_init_kwargs(store_init_kwargs: dict[str, Any] | None) -> dict[str, Any]: + if not store_init_kwargs: + return default_store_init_kwargs() + return default_store_init_kwargs() | dict(store_init_kwargs) + + +def create_mooncake_store(store_init_kwargs: dict[str, Any] | None = None) -> Any: + from mooncake.store import MooncakeDistributedStore # type: ignore + + kwargs = normalize_store_init_kwargs(store_init_kwargs) + store = MooncakeDistributedStore() + if store.setup(**kwargs) != 0: + raise RuntimeError("Mooncake Store setup failed") + return store + + +def get_cached_mooncake_store(store_init_kwargs: dict[str, Any] | None = None) -> Any: + kwargs = normalize_store_init_kwargs(store_init_kwargs) + cache_key = tuple(sorted((key, repr(val)) for key, val in kwargs.items())) + if cache_key not in _STORE_CACHE: + _STORE_CACHE[cache_key] = create_mooncake_store(kwargs) + return _STORE_CACHE[cache_key] + + +def remove_mooncake_keys(store: Any, keys: list[str]) -> None: + errors = [] + for key in sorted(set(keys)): + ret = store.remove(key, True) + if ret != 0: + errors.append((key, ret)) + if errors: + raise RuntimeError(f"Mooncake key cleanup failed: {errors}") + + +@dataclass +class MooncakeRemoteBatch: + """Remote tensor batch stored in Mooncake Store.""" + + fields: dict[str, tuple[str, tuple[int, ...], str]] # name -> (store_key, shape, dtype) + batch_size: int + store_init_kwargs: dict[str, Any] = field(default_factory=dict) + keys_to_cleanup: tuple[str, ...] = () + + @classmethod + def from_tensors( + cls, + tensors: dict[str, torch.Tensor], + store: Any, + prefix: str, + store_init_kwargs: dict[str, Any] | None = None, + ) -> MooncakeRemoteBatch: + if not prefix or ".." in prefix: + raise ValueError(f"invalid Mooncake key prefix: {prefix!r}") + config = _hard_pin_config(store) + fields: dict[str, tuple[str, tuple[int, ...], str]] = {} + written_keys: list[str] = [] + batch_size = None + try: + for name, tensor in tensors.items(): + if _FIELD_NAME_RE.fullmatch(name) is None: + raise ValueError(f"invalid Mooncake tensor field name: {name!r}") + cpu_tensor = tensor.detach().contiguous().cpu() + if batch_size is None: + batch_size = int(cpu_tensor.shape[0]) + elif int(cpu_tensor.shape[0]) != batch_size: + raise ValueError(f"tensor {name} batch size mismatch") + key = f"{prefix}/{name}" + if _put_tensor(store, key, cpu_tensor, config) != 0: + raise RuntimeError(f"Mooncake put failed for {key}") + written_keys.append(key) + fields[name] = (key, tuple(cpu_tensor.shape), str(cpu_tensor.dtype).removeprefix("torch.")) + except Exception: + remove_mooncake_keys(store, written_keys) + raise + return cls( + fields=fields, + batch_size=batch_size or 0, + store_init_kwargs=store_init_kwargs or {}, + keys_to_cleanup=tuple(written_keys), + ) + + def __len__(self) -> int: + return self.batch_size + + def materialize(self, fields: list[str] | None = None) -> dict[str, torch.Tensor]: + store = get_cached_mooncake_store(self.store_init_kwargs) + selected = list(self.fields) if fields is None else fields + return {name: _get_tensor(store, *self.fields[name]) for name in selected} + + def cleanup(self) -> None: + if self.keys_to_cleanup: + remove_mooncake_keys(get_cached_mooncake_store(self.store_init_kwargs), list(self.keys_to_cleanup)) + + +def _put_tensor(store: Any, key: str, tensor: torch.Tensor, config: Any) -> int: + arr = np.ascontiguousarray(tensor.numpy()) + region = _WritableRegion(memoryview(arr)) + try: + if store.register_buffer(region.ptr, region.size) != 0: + raise RuntimeError("register_buffer failed for put_from") + try: + return store.put_from(key=key, buffer_ptr=region.ptr, size=region.size, config=config) + finally: + if store.unregister_buffer(region.ptr) != 0: + raise RuntimeError("unregister_buffer failed for put_from") + finally: + region.close() + + +def _get_tensor(store: Any, key: str, shape: tuple[int, ...], dtype_name: str) -> torch.Tensor: + torch_dtype = getattr(torch, dtype_name.lower()) + nbytes = int(np.prod(shape, dtype=np.int64)) * torch_dtype.itemsize + if nbytes == 0: + return torch.empty(shape, dtype=torch_dtype) + region = _WritableRegion(bytearray(nbytes)) + try: + if store.register_buffer(region.ptr, region.size) != 0: + raise RuntimeError("register_buffer failed for get_into") + try: + got = store.get_into(key, region.ptr, nbytes) + if got != nbytes: + raise RuntimeError(f"get_into failed for {key}: expected {nbytes}, got {got}") + finally: + if store.unregister_buffer(region.ptr) != 0: + raise RuntimeError("unregister_buffer failed for get_into") + count = int(np.prod(shape, dtype=np.int64)) + return torch.frombuffer(region.buffer, dtype=torch_dtype, count=count).reshape(shape).clone() + finally: + region.close() + + +class _WritableRegion: + def __init__(self, buffer: Any) -> None: + self.buffer = buffer + self.view = memoryview(buffer).cast("B") + self.c_buffer = (ctypes.c_ubyte * self.view.nbytes).from_buffer(self.view) + self.ptr = ctypes.addressof(self.c_buffer) + self.size = self.view.nbytes + + def close(self) -> None: + self.c_buffer = None + self.view.release() + + +def _hard_pin_config(store: Any) -> Any: + from mooncake.store import ReplicateConfig # type: ignore + + config = ReplicateConfig() + config.preferred_segments = [store.get_hostname()] + config.with_hard_pin = True + return config diff --git a/vime/utils/rollout_store_transfer.py b/vime/utils/rollout_store_transfer.py new file mode 100644 index 000000000..c5627bec3 --- /dev/null +++ b/vime/utils/rollout_store_transfer.py @@ -0,0 +1,190 @@ +from __future__ import annotations + +import uuid +from dataclasses import dataclass, field +from typing import Any + +import numpy as np +import torch +from torch.nn.utils.rnn import pad_sequence + +from vime.utils.remote_batch import ( + MooncakeRemoteBatch, + get_cached_mooncake_store, + normalize_store_init_kwargs, + remove_mooncake_keys, +) + +REMOTE_TENSOR_KEYS = ("tokens", "loss_masks") +PARTITIONED_KEYS = ( + "tokens", + "multimodal_train_inputs", + "response_lengths", + "rewards", + "truncated", + "loss_masks", + "round_number", + "sample_indices", + "rollout_ids", + "rollout_mask_sums", + "rollout_log_probs", + "rollout_routed_experts", + "prompt", + "teacher_log_probs", + "metadata", +) +GLOBAL_KEYS = ("raw_reward", "total_lengths") + + +@dataclass +class RolloutStoreBatch: + non_tensor_batch: dict[str, np.ndarray] = field(default_factory=dict) + meta_info: dict = field(default_factory=dict) + remote_batch: MooncakeRemoteBatch | None = None + tensors: dict[str, torch.Tensor] | None = None + + def materialize_remote_batch(self) -> RolloutStoreBatch: + if self.remote_batch is not None: + self.tensors = self.remote_batch.materialize() + self.remote_batch = None + return self + + +def split_rollout_data_by_dp_mooncake_store( + args: Any, + data: dict, + dp_size: int, + partitions: list, + micro_batch_indices: list | None = None, + num_microbatches: int | None = None, + global_batch_sizes: list | None = None, + dynamic_global_batch_size: int | None = None, +) -> list[RolloutStoreBatch]: + if len(partitions) != dp_size: + raise ValueError(f"expected {dp_size} partitions, got {len(partitions)}") + + store_init_kwargs = normalize_store_init_kwargs(getattr(args, "mooncake_store_init_kwargs", None)) + store = get_cached_mooncake_store(store_init_kwargs) + transfer_id = uuid.uuid4().hex + refs: list[RolloutStoreBatch] = [] + try: + for dp_rank, partition in enumerate(partitions): + indices = [int(idx) for idx in partition] + shard = {key: [data[key][idx] for idx in indices] for key in PARTITIONED_KEYS if key in data} + shard["partition"] = np.asarray(indices, dtype=np.int64) + + meta_info = {key: data[key] for key in GLOBAL_KEYS if key in data} + if dynamic_global_batch_size is not None: + meta_info["dynamic_global_batch_size"] = dynamic_global_batch_size + if global_batch_sizes is not None: + meta_info["global_batch_sizes"] = global_batch_sizes + if num_microbatches is not None: + meta_info["num_microbatches"] = num_microbatches + if micro_batch_indices is not None: + meta_info["micro_batch_indices"] = micro_batch_indices[dp_rank] + + remote_tensors, remote_lengths = _extract_remote_tensors(shard) + meta_info.update(remote_lengths) + remote_batch = None + if remote_tensors: + remote_batch = MooncakeRemoteBatch.from_tensors( + remote_tensors, + store, + prefix=f"vime-rollout/{transfer_id}/dp{dp_rank}", + store_init_kwargs=store_init_kwargs, + ) + meta_info["mooncake_cleanup_keys"] = list(remote_batch.keys_to_cleanup) + meta_info["mooncake_cleanup_store_kwargs"] = dict(store_init_kwargs) + + try: + refs.append( + RolloutStoreBatch( + non_tensor_batch=_dict_to_non_tensors(shard), + meta_info=meta_info, + remote_batch=remote_batch, + ) + ) + except Exception: + if remote_batch is not None: + remote_batch.cleanup() + raise + except Exception: + cleanup_mooncake_store_refs(refs) + raise + return refs + + +def maybe_cleanup_mooncake_store_refs( + args: Any, refs: list[RolloutStoreBatch] | RolloutStoreBatch, suppress_errors: bool = False +) -> None: + if getattr(args, "transfer_backend", "ray") != "mooncake_store": + return + batches = refs if isinstance(refs, list) else [refs] + try: + cleanup_mooncake_store_refs(batches) + except Exception: + if not suppress_errors: + raise + + +def cleanup_mooncake_store_refs(refs: list[RolloutStoreBatch]) -> None: + keys: set[str] = set() + store_init_kwargs = None + for batch in refs: + keys.update(batch.meta_info.get("mooncake_cleanup_keys", [])) + if store_init_kwargs is None: + store_init_kwargs = batch.meta_info.get("mooncake_cleanup_store_kwargs") + if keys and store_init_kwargs is not None: + remove_mooncake_keys(get_cached_mooncake_store(store_init_kwargs), sorted(keys)) + + +def rollout_store_batch_to_data(batch: RolloutStoreBatch) -> dict: + batch.materialize_remote_batch() + rollout_data = {key: val.tolist() for key, val in batch.non_tensor_batch.items()} + rollout_data.update( + {key: val for key, val in batch.meta_info.items() if not key.startswith("mooncake_cleanup_")} + ) + if batch.tensors: + for key, tensor in batch.tensors.items(): + lengths = batch.meta_info.get(f"{key}_lengths") + rollout_data[key] = _tensor_to_row_tensors(tensor, lengths) + return rollout_data + + +def _extract_remote_tensors(shard: dict) -> tuple[dict[str, torch.Tensor], dict[str, list[int]]]: + tensors = {} + lengths = {} + for key in REMOTE_TENSOR_KEYS: + if key not in shard: + continue + values = shard.pop(key) + tensor, field_lengths = _list_to_padded_tensor(values, torch.long if key == "tokens" else torch.int) + tensors[key] = tensor + lengths[f"{key}_lengths"] = field_lengths + return tensors, lengths + + +def _list_to_padded_tensor(values: list, dtype: torch.dtype) -> tuple[torch.Tensor, list[int]]: + if not values: + return torch.empty((0, 0), dtype=dtype), [] + tensors = [torch.as_tensor(value, dtype=dtype).reshape(-1) for value in values] + lengths = [int(tensor.numel()) for tensor in tensors] + return pad_sequence(tensors, batch_first=True, padding_value=0), lengths + + +def _tensor_to_row_tensors(tensor: torch.Tensor, lengths: list[int] | None) -> list[torch.Tensor]: + if tensor.ndim == 2 and lengths is not None: + return [tensor[idx, : int(length)] for idx, length in enumerate(lengths)] + return [tensor[idx] for idx in range(tensor.shape[0])] + + +def _dict_to_non_tensors(data: dict) -> dict[str, np.ndarray]: + result = {} + for key, val in data.items(): + if isinstance(val, np.ndarray): + result[key] = val + elif isinstance(val, (int, float, bool, np.number)): + result[key] = np.asarray([val]) + else: + result[key] = np.asarray(val, dtype=object) + return result