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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
66 changes: 66 additions & 0 deletions tests/test_megatron_argument_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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__]))
70 changes: 70 additions & 0 deletions tests/test_modelexpress_vllm.py
Original file line number Diff line number Diff line change
@@ -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"}
Loading