From f2c3657a57d6fb504341854509fbf061f13a2670 Mon Sep 17 00:00:00 2001 From: Hyunjae Woo Date: Thu, 10 Sep 2026 13:33:30 -0700 Subject: [PATCH 1/2] feat(modelexpress): add S3 delta refit transport Signed-off-by: Hyunjae Woo --- tests/test_megatron_argument_validation.py | 66 ++ tests/test_modelexpress_vllm.py | 70 ++ tests/test_modelexpress_weight_update.py | 599 ++++++++++++++++++ tests/test_update_weight_factory.py | 8 + .../megatron_utils/update_weight/__init__.py | 4 + .../update_weight_from_modelexpress.py | 315 +++++++++ vime/backends/vllm_utils/vllm_engine.py | 4 +- vime/utils/arguments.py | 28 +- 8 files changed, 1088 insertions(+), 6 deletions(-) create mode 100644 tests/test_modelexpress_vllm.py create mode 100644 tests/test_modelexpress_weight_update.py create mode 100644 vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py diff --git a/tests/test_megatron_argument_validation.py b/tests/test_megatron_argument_validation.py index fd124719..5392bbed 100644 --- a/tests/test_megatron_argument_validation.py +++ b/tests/test_megatron_argument_validation.py @@ -259,6 +259,7 @@ def make_vime_validate_args(**overrides): update_weight_local_checkpoint_dir=None, update_weight_mode="full", rollout_temperature=1.0, + modelexpress_config={}, ) values.update(overrides) return types.SimpleNamespace(**values) @@ -430,5 +431,70 @@ def test_force_fp8_ue8m0_scale_argument(monkeypatch): assert configured.force_fp8_ue8m0_scale is True +@pytest.mark.unit +def test_modelexpress_does_not_require_native_disk_configuration(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + args = make_vime_validate_args( + update_weight_transport="modelexpress", + modelexpress_config={ + "model_name": "policy", + "server_url": "dns:///mx:50051", + "initial_base_version_id": "base-uid", + "s3_uri_prefix": "s3://weights/run/policy", + "seed_checkpoint_path": "/models/seed", + "refit_checkpoint_dir": "/mxdelta/refit", + }, + ) + + module.vime_validate_args(args) + + +@pytest.mark.unit +def test_modelexpress_uses_existing_transfer_selector_and_one_json_config(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + parser = argparse.ArgumentParser() + module.get_vime_extra_args_provider()(parser) + + args = parser.parse_args( + [ + "--update-weight-transport", + "modelexpress", + "--modelexpress-config", + '{"model_name":"policy","future_option":{"enabled":true}}', + "--rollout-batch-size", + "1", + ] + ) + + assert args.update_weight_transport == "modelexpress" + assert args.modelexpress_config == { + "model_name": "policy", + "future_option": {"enabled": True}, + } + assert not hasattr(args, "update_weight_backend") + assert not hasattr(args, "modelexpress_model_id") + + +@pytest.mark.unit +def test_modelexpress_config_rejects_non_object_json(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + args = make_vime_validate_args( + update_weight_transport="modelexpress", + modelexpress_config=["not", "an", "object"], + ) + + with pytest.raises(ValueError, match="must be a JSON object"): + module.vime_validate_args(args) + + +@pytest.mark.unit +def test_modelexpress_config_requires_modelexpress_transport(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + args = make_vime_validate_args(modelexpress_config={"future_option": True}) + + with pytest.raises(ValueError, match="requires --update-weight-transport=modelexpress"): + module.vime_validate_args(args) + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/tests/test_modelexpress_vllm.py b/tests/test_modelexpress_vllm.py new file mode 100644 index 00000000..e0846ae5 --- /dev/null +++ b/tests/test_modelexpress_vllm.py @@ -0,0 +1,70 @@ +import sys +import types +from pathlib import Path +from types import SimpleNamespace + +import pytest + +_tests_root = Path(__file__).resolve().parent +if str(_tests_root) not in sys.path: + sys.path.insert(0, str(_tests_root)) + +import _unit_stubs + +if "cloudpickle" not in sys.modules: + cloudpickle = types.ModuleType("cloudpickle") + cloudpickle.dumps = lambda value: b"stub" + sys.modules["cloudpickle"] = cloudpickle +_unit_stubs.install_vllm_cli_stubs() + +from vime.backends.vllm_utils import vllm_engine +from vime.backends.vllm_utils.vllm_engine import VLLMEngine + +pytestmark = pytest.mark.unit + + +def test_modelexpress_proxy_sends_exact_target_through_vllm_weight_transfer(monkeypatch): + engine = VLLMEngine.__new__(VLLMEngine) + engine.node_rank = 0 + calls = [] + monkeypatch.setattr( + engine, + "_make_request", + lambda endpoint, payload=None: calls.append((endpoint, payload)), + ) + + engine.update_weights({"version_id": "a1b2c3d4"}) + + assert calls == [("update_weights", {"update_info": {"version_id": "a1b2c3d4"}})] + + +def test_modelexpress_selects_vllm_backend(monkeypatch): + args = SimpleNamespace( + actor_num_gpus_per_node=8, + actor_num_nodes=1, + colocate=False, + debug_rollout_only=False, + fp16=False, + num_gpus_per_node=8, + offload_rollout=False, + rollout_num_gpus_per_engine=1, + seed=1, + update_weight_transport="modelexpress", + use_critic=False, + use_rollout_routing_replay=False, + vllm_data_parallel_size=1, + vllm_dp_size=1, + vllm_pipeline_parallel_size=1, + ) + vars(args)["hf_checkpoint"] = "/models/model" + monkeypatch.setattr(vllm_engine, "_VLLM_SERVER_FIELDS", frozenset()) + + server_args, _ = vllm_engine._compute_server_args( + args, + rank=0, + dist_init_addr=None, + host="127.0.0.1", + port=30000, + ) + + assert server_args["weight_transfer_config"] == {"backend": "modelexpress"} diff --git a/tests/test_modelexpress_weight_update.py b/tests/test_modelexpress_weight_update.py new file mode 100644 index 00000000..4071af0d --- /dev/null +++ b/tests/test_modelexpress_weight_update.py @@ -0,0 +1,599 @@ +import sys +import types +from argparse import Namespace +from dataclasses import dataclass +from enum import Enum +from types import SimpleNamespace + +import pytest +import torch + +ray_module = types.ModuleType("ray") +ray_module.get = lambda refs: refs +ray_actor_module = types.ModuleType("ray.actor") +ray_actor_module.ActorHandle = object +sys.modules.setdefault("ray", ray_module) +sys.modules.setdefault("ray.actor", ray_actor_module) + +megatron_module = types.ModuleType("megatron") +megatron_core_module = types.ModuleType("megatron.core") +megatron_core_module.mpu = types.SimpleNamespace() +megatron_module.core = megatron_core_module +sys.modules.setdefault("megatron", megatron_module) +sys.modules.setdefault("megatron.core", megatron_core_module) + +distributed_module = types.ModuleType("vime.utils.distributed_utils") +distributed_module.get_gloo_group = lambda: object() +sys.modules.setdefault("vime.utils.distributed_utils", distributed_module) + +iterator_module = types.ModuleType("vime.backends.megatron_utils.update_weight.hf_weight_iterator_direct") + + +class HfWeightIteratorDirect: + def __init__(self, **_kwargs): + pass + + def get_hf_weight_chunks(self, _weights, progress_desc, should_convert_chunk): + assert progress_desc == "Stage ModelExpress weights" + chunks = [[("weight", object())]] + return iter(chunk for index, chunk in enumerate(chunks) if should_convert_chunk(index)) + + +iterator_module.HfWeightIteratorDirect = HfWeightIteratorDirect +sys.modules.setdefault(iterator_module.__name__, iterator_module) + + +class WeightPayloadFormat(Enum): + FULL_TENSOR = "FULL_TENSOR" + FULL_HF_CHECKPOINT = "FULL_HF_CHECKPOINT" + XOR_DELTA = "XOR_DELTA" + + +class WeightVersionState(Enum): + STAGING = "STAGING" + READY = "READY" + + +class ObjectStorageType(Enum): + S3 = "S3" + + +class ObjectStorageConfig(SimpleNamespace): + pass + + +class ModelExpressTrainerConfig(SimpleNamespace): + pass + + +@dataclass(frozen=True) +class ObjectStorageSource: + storage_type: ObjectStorageType + uri: str + + +@dataclass(frozen=True) +class WeightVersionRef: + version_id: str + + +modelexpress_rl = types.ModuleType("modelexpress_rl") +modelexpress_rl.ModelExpressControlClient = object +modelexpress_rl.ModelExpressTrainerClient = object +modelexpress_rl.ModelExpressTrainerConfig = ModelExpressTrainerConfig +modelexpress_rl.ObjectStorageConfig = ObjectStorageConfig +modelexpress_rl.ObjectStorageSource = ObjectStorageSource +modelexpress_rl.ObjectStorageType = ObjectStorageType +modelexpress_rl.TrainerStagingMode = SimpleNamespace(WRITE_TO_STORAGE="WRITE_TO_STORAGE") +modelexpress_rl.WeightPayloadFormat = WeightPayloadFormat +modelexpress_rl.WeightVersionRef = WeightVersionRef +modelexpress_rl.WeightVersionState = WeightVersionState +sys.modules.setdefault("modelexpress_rl", modelexpress_rl) + +from vime.backends.megatron_utils.update_weight import update_weight_from_modelexpress as mx_module +from vime.backends.megatron_utils.update_weight.update_weight_from_modelexpress import UpdateWeightFromModelExpress + +pytestmark = pytest.mark.unit + + +class RemoteMethod: + def __init__(self, fn): + self._fn = fn + + def remote(self, *args, **kwargs): + return self._fn(*args, **kwargs) + + +class FakeControl: + def __init__(self): + self.created = [] + self._next_version = 1 + self.create_error = None + self.state_updates = [] + self.state_error = None + + def create_weight_version(self, **kwargs): + self.created.append(kwargs) + if self.create_error is not None: + raise RuntimeError(self.create_error) + version_id = kwargs.get("uid") + if version_id is None: + version_id = f"opaque-{self._next_version}" + self._next_version += 1 + return SimpleNamespace(version_id=version_id) + + def update_weight_version_state(self, version_id, state): + self.state_updates.append((version_id, state)) + if self.state_error is not None: + error, self.state_error = self.state_error, None + raise RuntimeError(error) + return SimpleNamespace(version_id=version_id, state=state) + + +class FakeStaged: + def __init__(self, trainer, version_id, buckets): + self._trainer = trainer + self._version_id = version_id + self._buckets = buckets + + def publish(self): + self._trainer.publishes.append((self._version_id, self._buckets)) + + +class FakeTrainer: + def __init__(self): + self.model_name = "policy" + self.server_url = "dns:///mx:50051" + self.baselines = [] + self.stages = [] + self.publishes = [] + self.metrics = { + "changed_bytes": 25, + "total_bytes": 100, + "wire_bytes": 123, + "stage_delta_time": 7.0, + "publish_object_storage_time": 8.0, + } + + def prepare_delta_base(self, *, hf_tensor_iter): + self.baselines.append(list(hf_tensor_iter)) + + def stage_shard(self, *, version, hf_tensor_iter): + buckets = list(hf_tensor_iter) + self.stages.append((version.version_id, buckets)) + return FakeStaged(self, version.version_id, buckets) + + def pop_metrics(self): + metrics, self.metrics = self.metrics, {} + return metrics + + +class FakeEngine: + def __init__(self, events, update_error=None, init_error=None): + self.events = events + self.update_error = update_error + self.init_error = init_error + self.active = False + self.init_weight_transfer_engine = RemoteMethod(self._init) + self.pause_generation = RemoteMethod(lambda: self._event("pause")) + self.flush_cache = RemoteMethod(lambda: self._event("flush")) + self.start_weight_update = RemoteMethod(self._start) + self.update_weights = RemoteMethod(self._update) + self.finish_weight_update = RemoteMethod(self._finish) + self.continue_generation = RemoteMethod(lambda: self._event("continue")) + + def _event(self, name): + self.events.append((name, None)) + return {"ok": True} + + def _init(self, payload): + self.events.append(("init", payload)) + if self.init_error: + raise RuntimeError(self.init_error) + return {"ok": True} + + def _start(self): + if self.active: + raise RuntimeError("already active") + self.active = True + return self._event("start") + + def _update(self, update_info): + version_id = update_info["version_id"] + self._event(f"update:{version_id}") + if self.update_error: + self.active = False + raise RuntimeError(self.update_error) + return {"ok": True} + + def _finish(self): + if not self.active: + raise RuntimeError("not active") + self.active = False + return self._event("finish") + + +def args(**config_overrides): + modelexpress_config = { + "model_name": "policy", + "server_url": "dns:///mx:50051", + "initial_base_version_id": "base-uid", + "seed_checkpoint_path": "/models/seed", + "refit_checkpoint_dir": "/mxdelta/refit", + "s3_uri_prefix": "s3://weights/run/policy", + "s3_endpoint_url": "http://minio:9000", + "s3_region_name": "us-west-2", + "rpc_timeout_seconds": 321.0, + "max_transfer_attempts": 4, + } + modelexpress_config.update(config_overrides) + values = dict( + modelexpress_config=modelexpress_config, + ) + out = Namespace(**values) + vars(out)["hf_checkpoint"] = "/models/not-used-by-modelexpress" + return out + + +@pytest.fixture(autouse=True) +def patch_runtime(monkeypatch): + monkeypatch.setattr(mx_module.ray, "get", lambda refs: refs) + monkeypatch.setattr(mx_module.dist, "get_rank", lambda: 0) + monkeypatch.setattr(mx_module.dist, "get_world_size", lambda: 1) + monkeypatch.setattr(mx_module.dist, "barrier", lambda group=None: None) + monkeypatch.setattr( + mx_module.dist, + "all_reduce", + lambda value, op=None, group=None: None, + ) + monkeypatch.setattr( + mx_module.dist, + "broadcast_object_list", + lambda values, src, group=None: None, + ) + monkeypatch.setattr(mx_module, "get_gloo_group", lambda: object()) + + +def updater(monkeypatch, control, trainer, **config_overrides): + monkeypatch.setattr( + mx_module, + "ModelExpressControlClient", + SimpleNamespace(connect=lambda **_kwargs: control), + ) + monkeypatch.setattr( + mx_module, + "ModelExpressTrainerClient", + SimpleNamespace(initialize=lambda _config: trainer), + ) + instance = UpdateWeightFromModelExpress( + args(**config_overrides), + model=[], + weights_getter=lambda: {}, + model_name="qwen3", + quantization_config=None, + ) + return instance + + +def test_vime_builds_generic_trainer_object_storage_config_and_normalizes_nulls(monkeypatch): + captured = {} + + class Config(SimpleNamespace): + pass + + class TrainerClient: + @staticmethod + def initialize(config): + captured["config"] = config + trainer = FakeTrainer() + trainer.model_name = config.model_name + trainer.server_url = config.server_url + return trainer + + monkeypatch.setattr(mx_module, "ModelExpressTrainerClient", TrainerClient) + monkeypatch.setattr(mx_module, "ModelExpressTrainerConfig", Config) + monkeypatch.setattr( + mx_module, + "ModelExpressControlClient", + SimpleNamespace(connect=lambda **_kwargs: FakeControl()), + ) + + UpdateWeightFromModelExpress( + args(), + model=[], + weights_getter=lambda: {}, + model_name="qwen3", + quantization_config=None, + ) + + config = captured["config"] + assert config.object_storage.storage_type is ObjectStorageType.S3 + assert config.object_storage.uri_prefix == "s3://weights/run/policy" + assert config.object_storage.seed_checkpoint_path == "/models/seed" + assert not hasattr(config.object_storage, "process_group") + assert config.process_group is not None + + UpdateWeightFromModelExpress( + args( + s3_uri_prefix=None, + initial_base_version_id=None, + seed_checkpoint_path=None, + ), + model=[], + weights_getter=lambda: {}, + model_name="qwen3", + quantization_config=None, + ) + + object_storage = captured["config"].object_storage + assert object_storage.uri_prefix == "" + assert object_storage.initial_base_version_id == "" + assert object_storage.seed_checkpoint_path == "" + + +@pytest.mark.parametrize(("rank", "expected"), [(0, [0, 2]), (1, [1, 3])]) +def test_vime_assigns_each_hf_bucket_to_one_trainer_rank(monkeypatch, rank, expected): + instance = updater(monkeypatch, FakeControl(), FakeTrainer()) + + class Iterator: + def get_hf_weight_chunks(self, _weights, progress_desc, should_convert_chunk): + assert progress_desc == "Stage ModelExpress weights" + return iter([("weight", index)] for index in range(4) if should_convert_chunk(index)) + + instance._weight_iterator = Iterator() + monkeypatch.setattr(mx_module.dist, "get_rank", lambda: rank) + monkeypatch.setattr(mx_module.dist, "get_world_size", lambda: 2) + + assert [chunk[0][1] for chunk in instance._iter_hf_buckets()] == expected + + +def test_vime_initializes_vllm_and_publishes_version_owned_s3_delta(monkeypatch): + control = FakeControl() + trainer = FakeTrainer() + instance = updater(monkeypatch, control, trainer) + events = [] + instance.connect_rollout_engines([FakeEngine(events)], object()) + base_version = { + "uid": "base-uid", + "model_name": "policy", + "idempotency_key": "vime:s3://weights/run/policy/v0/model.safetensors.index.json", + "payload_format": WeightPayloadFormat.FULL_TENSOR, + "object_storage": ObjectStorageSource( + storage_type=ObjectStorageType.S3, + uri="s3://weights/run/policy/v0/model.safetensors.index.json", + ), + "state": WeightVersionState.READY, + } + assert control.created == [base_version] + assert not trainer.baselines + + instance.update_weights() + clock = iter([0.0, 5.0]) + reductions = [] + monkeypatch.setattr(mx_module, "perf_counter", lambda: next(clock)) + + def reduce_metrics(value, op=None, group=None): + reductions.append((value.tolist(), op)) + if value.dtype == torch.int64: + value.copy_(torch.tensor([50, 200, 246], dtype=value.dtype)) + else: + value.copy_(torch.tensor([17.0, 18.0, 15.0], dtype=value.dtype)) + + monkeypatch.setattr(mx_module.dist, "all_reduce", reduce_metrics) + instance.update_weights() + + assert trainer.baselines + assert trainer.stages[0][0] == "opaque-1" + assert trainer.publishes[0][0] == "opaque-1" + assert control.created == [ + base_version, + { + "model_name": "policy", + "idempotency_key": "vime:s3://weights/run/policy/v1/model.safetensors.index.json", + "payload_format": WeightPayloadFormat.XOR_DELTA, + "base_version_id": "base-uid", + "object_storage": ObjectStorageSource( + storage_type=ObjectStorageType.S3, + uri="s3://weights/run/policy/v1/model.safetensors.index.json", + ), + "state": WeightVersionState.STAGING, + }, + ] + assert control.state_updates == [("opaque-1", WeightVersionState.READY)] + assert events[0] == ( + "init", + { + "init_info": { + "model_name": "policy", + "server_url": "dns:///mx:50051", + "initial_base_version_id": "base-uid", + "seed_checkpoint_path": "/models/seed", + "refit_checkpoint_dir": "/mxdelta/refit", + "object_storage_type": "S3", + "object_storage_endpoint_url": "http://minio:9000", + "object_storage_region_name": "us-west-2", + "registration_ttl_seconds": None, + "lease_ttl_seconds": None, + "max_transfer_attempts": 4, + "rpc_timeout_seconds": 321.0, + } + }, + ) + assert [event for event, _payload in events[1:]] == [ + "pause", + "flush", + "start", + "update:opaque-1", + "finish", + "continue", + ] + assert instance.weight_version == 1 + assert instance._current_version_id == "opaque-1" + assert instance.pop_metrics() == { + "perf/update_weights_density": 0.25, + "perf/update_weights_wire_bytes": 246, + "perf/mx_stage_delta_time": 17.0, + "perf/mx_publish_object_storage_time": 18.0, + "perf/mx_update_engine_weights_time": 15.0, + } + assert reductions == [ + ([25, 100, 123], torch.distributed.ReduceOp.SUM), + ( + [7.0, 8.0, 5.0], + torch.distributed.ReduceOp.MAX, + ), + ] + + +def test_vime_publishes_periodic_full_hf_checkpoints(monkeypatch): + control = FakeControl() + instance = updater( + monkeypatch, + control, + FakeTrainer(), + full_hf_checkpoint_interval=2, + ) + instance.connect_rollout_engines([FakeEngine([])], object()) + + instance.update_weights() + for _ in range(3): + instance.update_weights() + + assert [created["payload_format"] for created in control.created] == [ + WeightPayloadFormat.FULL_TENSOR, + WeightPayloadFormat.XOR_DELTA, + WeightPayloadFormat.FULL_HF_CHECKPOINT, + WeightPayloadFormat.XOR_DELTA, + ] + assert control.created[1]["base_version_id"] == "base-uid" + assert "base_version_id" not in control.created[2] + assert control.created[3]["base_version_id"] == "opaque-2" + + +@pytest.mark.parametrize("interval", [0, -1, True, 1.5, "2"]) +def test_vime_rejects_invalid_full_hf_checkpoint_intervals(monkeypatch, interval): + with pytest.raises(ValueError, match="full_hf_checkpoint_interval must be a positive integer"): + updater( + monkeypatch, + FakeControl(), + FakeTrainer(), + full_hf_checkpoint_interval=interval, + ) + + +def test_reconnecting_the_same_vllm_cohort_is_a_noop(monkeypatch): + instance = updater(monkeypatch, FakeControl(), FakeTrainer()) + events = [] + engine = FakeEngine(events) + + instance.connect_rollout_engines([engine], object()) + instance.connect_rollout_engines([engine], object()) + + assert [event for event, _payload in events if event == "init"] == ["init"] + + +def test_failed_vllm_initialization_is_retried_for_the_same_cohort(monkeypatch): + instance = updater(monkeypatch, FakeControl(), FakeTrainer()) + events = [] + engine = FakeEngine(events, init_error="init failed") + + with pytest.raises(RuntimeError, match="init failed"): + instance.connect_rollout_engines([engine], object()) + + assert instance.rollout_engines is None + engine.init_error = None + instance.connect_rollout_engines([engine], object()) + + assert instance.rollout_engines == (engine,) + assert [event for event, _payload in events if event == "init"] == ["init", "init"] + + +def test_failed_baseline_registration_is_retried_before_vllm_init(monkeypatch): + control = FakeControl() + control.create_error = "catalog unavailable" + instance = updater(monkeypatch, control, FakeTrainer()) + events = [] + engine = FakeEngine(events) + + with pytest.raises(RuntimeError, match="catalog unavailable"): + instance.connect_rollout_engines([engine], object()) + + assert events == [] + assert instance.rollout_engines is None + control.create_error = None + instance.connect_rollout_engines([engine], object()) + + assert len(control.created) == 2 + assert control.created[1] == control.created[0] + assert [event for event, _payload in events] == ["init"] + + +def test_full_checkpoint_ready_failure_does_not_advance_version(monkeypatch): + control = FakeControl() + control.state_error = "ready failed" + trainer = FakeTrainer() + instance = updater( + monkeypatch, + control, + trainer, + full_hf_checkpoint_interval=1, + ) + engine = FakeEngine([]) + instance.connect_rollout_engines([engine], object()) + instance.update_weights() + + with pytest.raises(RuntimeError, match="ready failed"): + instance.update_weights() + + assert instance.weight_version == 0 + assert instance._current_version_id == "base-uid" + assert len(control.created) == 2 + assert len(trainer.publishes) == 1 + assert control.created[1]["payload_format"] is WeightPayloadFormat.FULL_HF_CHECKPOINT + assert "base_version_id" not in control.created[1] + + +def test_failed_vllm_update_does_not_resume_or_advance_version(monkeypatch): + control = FakeControl() + trainer = FakeTrainer() + instance = updater(monkeypatch, control, trainer) + events = [] + engine = FakeEngine(events, update_error="install failed") + instance.connect_rollout_engines([engine], object()) + instance.update_weights() + + with pytest.raises(RuntimeError, match="install failed"): + instance.update_weights() + + assert instance.weight_version == 0 + assert instance._current_version_id == "base-uid" + assert len(control.created) == 2 + assert len(trainer.publishes) == 1 + assert not any(event == "continue" for event, _payload in events[1:]) + + +def test_update_engine_weights_runs_in_bulk_phases(monkeypatch): + instance = updater(monkeypatch, FakeControl(), FakeTrainer()) + events = [] + first = FakeEngine(events) + second = FakeEngine(events) + instance.connect_rollout_engines([first, second], object()) + events.clear() + + instance._update_engine_weights("opaque-target") + + assert [event for event, _payload in events] == [ + "pause", + "pause", + "flush", + "flush", + "start", + "start", + "update:opaque-target", + "update:opaque-target", + "finish", + "finish", + "continue", + "continue", + ] + assert instance._update_engine_weights_time >= 0 diff --git a/tests/test_update_weight_factory.py b/tests/test_update_weight_factory.py index 0259cd89..b879f246 100644 --- a/tests/test_update_weight_factory.py +++ b/tests/test_update_weight_factory.py @@ -25,6 +25,14 @@ def __init__(self, args, model, weights_getter, *, model_name, quantization_conf [ pytest.param("delta", "disk", False, "update_weight_from_disk_delta", "UpdateWeightFromDiskDelta", id="delta"), pytest.param("full", "disk", False, "update_weight_from_disk", "UpdateWeightFromDisk", id="disk"), + pytest.param( + "full", + "modelexpress", + False, + "update_weight_from_modelexpress", + "UpdateWeightFromModelExpress", + id="modelexpress", + ), pytest.param("full", "nccl", True, "update_weight_from_tensor", "UpdateWeightFromTensor", id="colocated"), pytest.param( "full", diff --git a/vime/backends/megatron_utils/update_weight/__init__.py b/vime/backends/megatron_utils/update_weight/__init__.py index dc40d630..106738f8 100644 --- a/vime/backends/megatron_utils/update_weight/__init__.py +++ b/vime/backends/megatron_utils/update_weight/__init__.py @@ -32,6 +32,10 @@ def create_weight_updater( from .update_weight_from_disk import UpdateWeightFromDisk update_weight_cls = UpdateWeightFromDisk + elif update_weight_transport == "modelexpress": + from .update_weight_from_modelexpress import UpdateWeightFromModelExpress + + update_weight_cls = UpdateWeightFromModelExpress elif args.colocate: from .update_weight_from_tensor import UpdateWeightFromTensor diff --git a/vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py b/vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py new file mode 100644 index 00000000..429a8207 --- /dev/null +++ b/vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py @@ -0,0 +1,315 @@ +from __future__ import annotations + +from argparse import Namespace +from collections.abc import Callable, Mapping, Sequence +from time import perf_counter +from typing import Any + +import ray +import torch +import torch.distributed as dist +from modelexpress_rl import ( + ModelExpressControlClient, + ModelExpressTrainerClient, + ModelExpressTrainerConfig, + ObjectStorageConfig, + ObjectStorageSource, + ObjectStorageType, + TrainerStagingMode, + WeightPayloadFormat, + WeightVersionRef, + WeightVersionState, +) +from ray.actor import ActorHandle + +from vime.utils.distributed_utils import get_gloo_group + +from .hf_weight_iterator_direct import HfWeightIteratorDirect + + +class UpdateWeightFromModelExpress: + """Publish Megatron weights through ModelExpress and install them in vLLM. + + The current path publishes canonical S3 XOR deltas with optional periodic full + checkpoints. Future integrations may add P2P NIXL/RDMA transfer formats behind + the same updater lifecycle. + """ + + def __init__( + self, + args: Namespace, + model: Sequence[torch.nn.Module], + weights_getter: Callable[[], Mapping[str, torch.Tensor]], + model_name: str, + quantization_config: dict[str, int | str | list[str]] | None, + ) -> None: + self._config = dict(args.modelexpress_config) + self._weights_getter = weights_getter + self._weight_iterator = HfWeightIteratorDirect( + args=args, + model=model, + model_name=model_name, + quantization_config=quantization_config, + ) + self.weight_version = 0 + full_checkpoint_interval = self._config.get("full_hf_checkpoint_interval") + if full_checkpoint_interval is not None and ( + isinstance(full_checkpoint_interval, bool) + or not isinstance(full_checkpoint_interval, int) + or full_checkpoint_interval <= 0 + ): + raise ValueError("full_hf_checkpoint_interval must be a positive integer") + self._full_hf_checkpoint_interval = full_checkpoint_interval + self.rollout_engines: Sequence[ActorHandle] | None = None + self._baseline_captured = False + self._update_engine_weights_time = 0.0 + self.update_weight_metrics: dict[str, int | float] = {} + self._initialize() + + def _initialize(self) -> None: + """Initialize object storage and rank-local ModelExpress clients.""" + self._rpc_timeout_seconds = self._config.get("rpc_timeout_seconds", 30.0) + self._object_storage_config = ObjectStorageConfig( + storage_type=ObjectStorageType.S3, + uri_prefix=self._config.get("s3_uri_prefix") or "", + initial_base_version_id=self._config.get("initial_base_version_id") or "", + seed_checkpoint_path=self._config.get("seed_checkpoint_path") or "", + endpoint_url=self._config.get("s3_endpoint_url"), + region_name=self._config.get("s3_region_name"), + ) + self._current_version_id = self._object_storage_config.initial_base_version_id + + registration_ttl = self._config.get("registration_ttl_seconds") + self._trainer = ModelExpressTrainerClient.initialize( + ModelExpressTrainerConfig( + model_name=self._config.get("model_name"), + staging_mode=TrainerStagingMode.WRITE_TO_STORAGE, + payload_format=WeightPayloadFormat.XOR_DELTA, + server_url=self._config.get("server_url"), + registration_ttl_seconds=registration_ttl, + rpc_timeout_seconds=self._rpc_timeout_seconds, + process_group=get_gloo_group(), + object_storage=self._object_storage_config, + ) + ) + + if dist.get_rank() == 0: + self._control = ModelExpressControlClient.connect( + server_url=self._trainer.server_url, + rpc_timeout_seconds=self._rpc_timeout_seconds, + ) + else: + self._control = None + + def connect_rollout_engines( + self, + rollout_engines: Sequence[ActorHandle], + rollout_engine_lock: ActorHandle, + engine_gpu_counts: Sequence[int] | None = None, + engine_gpu_offsets: Sequence[int] | None = None, + engine_parallel_configs: Sequence[Mapping[str, object]] | None = None, + ) -> None: + """Initialize ModelExpress on a newly connected rollout-engine cohort.""" + del rollout_engine_lock, engine_gpu_counts, engine_gpu_offsets, engine_parallel_configs + connected = tuple(rollout_engines) + if self.rollout_engines == connected: + return + + if not self._baseline_captured: + base_uri = f"{self._object_storage_config.uri_prefix.rstrip('/')}" "/v0/model.safetensors.index.json" + self._rank_zero_call( + lambda: self._control.create_weight_version( + uid=self._current_version_id, + model_name=self._trainer.model_name, + idempotency_key=f"vime:{base_uri}", + payload_format=WeightPayloadFormat.FULL_TENSOR, + object_storage=ObjectStorageSource( + storage_type=self._object_storage_config.storage_type, + uri=base_uri, + ), + state=WeightVersionState.READY, + ), + f"ModelExpress baseline version {self._current_version_id} registration failed", + ) + + registration_ttl = self._config.get("registration_ttl_seconds") + lease_ttl = self._config.get("lease_ttl_seconds") + + init_info = { + "model_name": self._trainer.model_name, + "server_url": self._trainer.server_url, + "initial_base_version_id": self._current_version_id, + "seed_checkpoint_path": self._object_storage_config.seed_checkpoint_path, + "refit_checkpoint_dir": self._config.get("refit_checkpoint_dir"), + "object_storage_type": self._object_storage_config.storage_type.value, + "object_storage_endpoint_url": self._config.get("s3_endpoint_url"), + "object_storage_region_name": self._config.get("s3_region_name"), + "registration_ttl_seconds": registration_ttl, + "lease_ttl_seconds": lease_ttl, + "max_transfer_attempts": self._config.get("max_transfer_attempts", 3), + "rpc_timeout_seconds": self._rpc_timeout_seconds, + } + + self._rank_zero_call( + lambda: ray.get( + [engine.init_weight_transfer_engine.remote({"init_info": init_info}) for engine in connected] + ), + "vLLM ModelExpress initialization failed", + ) + self.rollout_engines = connected + + def disconnect_rollout_engines(self) -> None: + """Forget the current rollout-engine cohort.""" + self.rollout_engines = None + + def pop_metrics(self) -> dict[str, int | float]: + """Return and clear metrics from the latest weight update.""" + metrics, self.update_weight_metrics = self.update_weight_metrics, {} + return metrics + + def _rank_zero_call(self, action: Callable[[], Any], description: str) -> Any: + """Run an action on rank zero and broadcast its result or failure.""" + result = [None, None] + if dist.get_rank() == 0: + try: + result[0] = action() + except Exception as error: + result[1] = str(error) + dist.broadcast_object_list(result, src=0, group=get_gloo_group()) + if result[1] is not None: + raise RuntimeError(f"{description}: {result[1]}") + return result[0] + + @torch.no_grad() + def update_weights(self) -> None: + """Capture the initial base or publish and install one policy update.""" + if not self._baseline_captured: + self._trainer.prepare_delta_base(hf_tensor_iter=self._iter_hf_buckets()) + self._baseline_captured = True + return + + self._update_engine_weights_time = 0.0 + update_number = self.weight_version + 1 + + version = self._create_weight_version(update_number) + self._stage_and_publish(version) + + self._rank_zero_call( + lambda: self._update_engine_weights(version.version_id), + f"ModelExpress version {version.version_id} install failed", + ) + + self._current_version_id = version.version_id + self.weight_version += 1 + self.update_weight_metrics = self._gather_metrics( + group=get_gloo_group(), + ) + + def _create_weight_version(self, update_number: int) -> WeightVersionRef: + """Create one version and return its opaque ModelExpress identity.""" + publish_full_checkpoint = ( + self._full_hf_checkpoint_interval is not None and update_number % self._full_hf_checkpoint_interval == 0 + ) + object_storage_uri = ( + f"{self._object_storage_config.uri_prefix.rstrip('/')}" f"/v{update_number}/model.safetensors.index.json" + ) + create_kwargs = { + "model_name": self._trainer.model_name, + "idempotency_key": f"vime:{object_storage_uri}", + "payload_format": ( + WeightPayloadFormat.FULL_HF_CHECKPOINT if publish_full_checkpoint else WeightPayloadFormat.XOR_DELTA + ), + "object_storage": ObjectStorageSource( + storage_type=self._object_storage_config.storage_type, + uri=object_storage_uri, + ), + "state": WeightVersionState.STAGING, + } + if not publish_full_checkpoint: + create_kwargs["base_version_id"] = self._current_version_id + target_version_id = str( + self._rank_zero_call( + lambda: self._control.create_weight_version(**create_kwargs).version_id, + f"ModelExpress update {update_number} creation failed", + ) + ) + return WeightVersionRef(target_version_id) + + def _stage_and_publish(self, version: WeightVersionRef) -> None: + """Stage a version, publish its artifacts, and mark it READY.""" + staged = self._trainer.stage_shard( + version=version, + hf_tensor_iter=self._iter_hf_buckets(), + ) + + staged.publish() + + dist.barrier(group=get_gloo_group()) + self._rank_zero_call( + lambda: self._control.update_weight_version_state( + version.version_id, + WeightVersionState.READY, + ), + f"ModelExpress version {version.version_id} activation failed", + ) + + def _iter_hf_buckets(self): + """Yield canonical non-expert and expert Hugging Face weight buckets.""" + rank = dist.get_rank() + world_size = dist.get_world_size() + yield from self._weight_iterator.get_hf_weight_chunks( + self._weights_getter(), + progress_desc="Stage ModelExpress weights", + should_convert_chunk=lambda chunk_idx: chunk_idx % world_size == rank, + ) + + def _gather_metrics( + self, + *, + group: Any, + ) -> dict[str, int | float]: + """Aggregate ModelExpress publication and serving-cutover metrics.""" + local_metrics = self._trainer.pop_metrics() + counts = torch.tensor( + [ + local_metrics.get("changed_bytes", 0), + local_metrics.get("total_bytes", 0), + local_metrics.get("wire_bytes", 0), + ], + dtype=torch.int64, + ) + dist.all_reduce(counts, op=dist.ReduceOp.SUM, group=group) + + timings = torch.tensor( + [ + local_metrics.get("stage_delta_time", 0.0), + local_metrics.get("publish_object_storage_time", 0.0), + self._update_engine_weights_time, + ], + dtype=torch.float64, + ) + dist.all_reduce(timings, op=dist.ReduceOp.MAX, group=group) + + changed_bytes, total_bytes, wire_bytes = counts.tolist() + stage_delta_time, publish_object_storage_time, update_engine_weights_time = timings.tolist() + return { + "perf/update_weights_density": changed_bytes / max(total_bytes, 1), + "perf/update_weights_wire_bytes": wire_bytes, + "perf/mx_stage_delta_time": stage_delta_time, + "perf/mx_publish_object_storage_time": publish_object_storage_time, + "perf/mx_update_engine_weights_time": update_engine_weights_time, + } + + def _update_engine_weights(self, target_version_id: str) -> None: + """Pause rollout engines, install one version, and resume generation.""" + engines = tuple(self.rollout_engines or ()) + if not engines: + raise RuntimeError("ModelExpress requires rollout engines") + phase_started = perf_counter() + ray.get([engine.pause_generation.remote() for engine in engines]) + ray.get([engine.flush_cache.remote() for engine in engines]) + ray.get([engine.start_weight_update.remote() for engine in engines]) + ray.get([engine.update_weights.remote({"version_id": target_version_id}) for engine in engines]) + ray.get([engine.finish_weight_update.remote() for engine in engines]) + ray.get([engine.continue_generation.remote() for engine in engines]) + self._update_engine_weights_time = perf_counter() - phase_started diff --git a/vime/backends/vllm_utils/vllm_engine.py b/vime/backends/vllm_utils/vllm_engine.py index 82c07e6f..b1c4336f 100644 --- a/vime/backends/vllm_utils/vllm_engine.py +++ b/vime/backends/vllm_utils/vllm_engine.py @@ -686,7 +686,9 @@ def _compute_server_args( ): kwargs["max_model_len"] = args.rollout_max_context_len - if args.colocate: + if getattr(args, "update_weight_transport", "nccl") == "modelexpress": + kwargs["weight_transfer_config"] = {"backend": "modelexpress"} + elif args.colocate: kwargs["weight_transfer_config"] = {"backend": "ipc"} else: kwargs["weight_transfer_config"] = {"backend": "nccl"} diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index 4eefc41f..c9f232a8 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -142,13 +142,26 @@ def add_train_arguments(parser): ) parser.add_argument( "--update-weight-transport", - choices=["nccl", "disk"], + choices=["nccl", "disk", "modelexpress"], default="nccl", help=( - "Carrier for weight sync. In full mode, 'nccl' broadcasts chunks and " - "'disk' writes a complete HF checkpoint under --update-weight-disk-dir " - "before engines reload it. Delta mode is 'disk' only: each host applies the " - "published deltas into its local checkpoint and reloads via update_weights_from_disk." + "Carrier for weight sync. In full mode, 'nccl' broadcasts chunks, " + "'disk' writes a complete HF checkpoint under --update-weight-disk-dir, and " + "'modelexpress' publishes canonical S3 checkpoints through the " + "version lifecycle " + "configured by --modelexpress-config. " + "Delta mode is 'disk' only." + ), + ) + parser.add_argument( + "--modelexpress-config", + type=json.loads, + default={}, + help=( + "ModelExpress configuration as a JSON object, including " + "seed_checkpoint_path and refit_checkpoint_dir. Set " + "full_hf_checkpoint_interval to a positive integer to publish " + "periodic full HF checkpoints." ), ) parser.add_argument( @@ -2183,6 +2196,11 @@ def vime_validate_args(args): if args.only_train_params_name_list and args.freeze_params_name_list: raise ValueError("You can only specify ONE of: --only-train-params-name-list, or --freeze-params-name-list.") + if not isinstance(args.modelexpress_config, dict): + raise ValueError("--modelexpress-config must be a JSON object") + if args.modelexpress_config and args.update_weight_transport != "modelexpress": + raise ValueError("--modelexpress-config requires --update-weight-transport=modelexpress") + # disk-backed sync (full or delta) writes on the trainer and reads on the engines: needs a shared dir if args.update_weight_transport == "disk" and not args.update_weight_disk_dir: raise ValueError( From 22aabda20b8de38f2787754012821552f3e639f4 Mon Sep 17 00:00:00 2001 From: Hyunjae Woo Date: Tue, 15 Sep 2026 22:58:24 -0700 Subject: [PATCH 2/2] feat(modelexpress): forward refit checkpoint cache quota Forward an explicitly configured refit_checkpoint_max_size_gb to rollout init_info. Preserve the ModelExpress default when omitted and forward null to disable the quota. Document the setting and cover explicit limits and null in updater tests. Signed-off-by: Hyunjae Woo (cherry picked from commit ee6364f7a2cc268b36b0651120cfb709a4a1b5b8) Signed-off-by: Hyunjae Woo --- tests/test_modelexpress_weight_update.py | 14 ++++++++++++++ .../update_weight_from_modelexpress.py | 2 ++ vime/utils/arguments.py | 3 +++ 3 files changed, 19 insertions(+) diff --git a/tests/test_modelexpress_weight_update.py b/tests/test_modelexpress_weight_update.py index 4071af0d..924764bb 100644 --- a/tests/test_modelexpress_weight_update.py +++ b/tests/test_modelexpress_weight_update.py @@ -445,6 +445,20 @@ def reduce_metrics(value, op=None, group=None): ] +@pytest.mark.parametrize("max_size_gb", [16, None]) +def test_vime_forwards_refit_checkpoint_cache_limit(monkeypatch, max_size_gb): + instance = updater( + monkeypatch, + FakeControl(), + FakeTrainer(), + refit_checkpoint_max_size_gb=max_size_gb, + ) + events = [] + instance.connect_rollout_engines([FakeEngine(events)], object()) + + assert events[0][1]["init_info"]["refit_checkpoint_max_size_gb"] == max_size_gb + + def test_vime_publishes_periodic_full_hf_checkpoints(monkeypatch): control = FakeControl() instance = updater( diff --git a/vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py b/vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py index 429a8207..034b0b9f 100644 --- a/vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py +++ b/vime/backends/megatron_utils/update_weight/update_weight_from_modelexpress.py @@ -149,6 +149,8 @@ def connect_rollout_engines( "max_transfer_attempts": self._config.get("max_transfer_attempts", 3), "rpc_timeout_seconds": self._rpc_timeout_seconds, } + if "refit_checkpoint_max_size_gb" in self._config: + init_info["refit_checkpoint_max_size_gb"] = self._config["refit_checkpoint_max_size_gb"] self._rank_zero_call( lambda: ray.get( diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index c9f232a8..6633b2b6 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -160,6 +160,9 @@ def add_train_arguments(parser): help=( "ModelExpress configuration as a JSON object, including " "seed_checkpoint_path and refit_checkpoint_dir. Set " + "refit_checkpoint_max_size_gb to a positive integer for the rollout " + "checkpoint-cache quota in decimal GB; omission keeps the ModelExpress " + "default and null disables the quota. Set " "full_hf_checkpoint_interval to a positive integer to publish " "periodic full HF checkpoints." ),