Skip to content

Commit 5ddbf0a

Browse files
ddmm2020liuweijian刘伟健
authored
feat: expose ParameterServer metas via HTTP for cross-process join (MoonshotAI#90)
* feat: expose ParameterServer metas via HTTP for cross-process join Add two HTTP endpoints (GET/POST /v1/checkpoints/{name}/{metas,load-metas}) and a standalone `python -m checkpoint_engine.join_cli` entrypoint, so a new ParameterServer instance can join an existing P2P weight world over mooncake RDMA without re-reading checkpoints from disk. Motivation: in elastic-rollout setups (e.g. mshrl), a long-running training job already holds pinned CPU weight buffers registered with the mooncake P2PStore. Newly-started inference replicas should be able to pull these weights over RDMA instead of re-converting the checkpoint from disk. Changes: * api.py: GET /v1/checkpoints/{name}/metas returns pickle.dumps(ps.get_metas()) as application/octet-stream; POST /v1/checkpoints/{name}/load-metas accepts the same bytes and feeds them into ps.load_metas(). Bad pickle is rejected with 400; PS errors are surfaced as 500. * join_cli.py: `python -m checkpoint_engine.join_cli` -- the join() flow from examples/update.py, packaged as a first-class CLI under the published package so consumers can invoke it without checking out the source tree. Accepts metas from either a local pickle file or a remote HTTP URL. * tests/test_api.py: 6 CPU-only tests covering pickle round-trip, ps-error propagation, bad-input rejection, and a GET-then-POST chain that validates the new endpoints are mutually consistent. Verified end-to-end on a 2-node launchpad job: 14.5 GiB Qwen2.5-7B weights transferred from main to elastic in 1.49s over real RDMA (4 mlx5_bond HCAs) vs 6.49s over TCP fallback in environments without RDMA passthrough. * refactor: serialize metas as JSON instead of pickle Replace pickle with pydantic TypeAdapter(dict[int, MemoryBufferMetaList]) for the metas wire format across the HTTP endpoints, join_cli, and examples/update.py. This reuses the existing pydantic schema (torch.dtype / torch.Size already have serializers in data_types.py), removes the arbitrary-code-execution risk of pickle.loads on request bodies, and makes the metas self-describing for cross-language consumers. - api.py: GET /metas returns application/json; POST /load-metas validates via validate_json and returns 400 on ValidationError (was broad except). - join_cli.py / examples/update.py: read/write metas as JSON; document the --metas-url HTTP path alongside --load-metas-file. - tests/test_api.py: use real MemoryBufferMetaList fixtures; add a schema-mismatch case (valid JSON, wrong shape -> 400). * refactor: merge join_cli into examples/update.py join_cli was a duplicate of examples/update.py:join() with an extra --metas-url flag. Add the flag to update.py, remove join_cli. * refactor: simplify metas endpoints to /v1/metas The checkpoint_name path param was never read — ps.get_metas() and ps.load_metas() act on a single global field. Drop it from the URL. * refactor: drop _METAS_ADAPTER from api.py * test: set Content-Type on POST /v1/metas requests --------- Co-authored-by: liuweijian <liuweijian@msh.team> Co-authored-by: 刘伟健 <liuweijian@moonshot.ai>
1 parent 59a3c4f commit 5ddbf0a

3 files changed

Lines changed: 185 additions & 9 deletions

File tree

‎checkpoint_engine/api.py‎

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,11 +3,12 @@
33

44
import fastapi
55
import httpx
6-
from fastapi import Request
6+
from fastapi import HTTPException, Request
77
from fastapi.responses import JSONResponse, Response
88
from loguru import logger
99
from pydantic import BaseModel
1010

11+
from checkpoint_engine.data_types import MemoryBufferMetaList
1112
from checkpoint_engine.ps import ParameterServer
1213

1314

@@ -79,6 +80,18 @@ async def healthz() -> Response:
7980
async def gather_metas(checkpoint_name: str) -> Response:
8081
return wrap_exception(lambda: ps.gather_metas(checkpoint_name))
8182

83+
@app.get("/v1/metas")
84+
async def get_metas() -> dict[int, MemoryBufferMetaList]:
85+
try:
86+
return ps.get_metas()
87+
except Exception as e:
88+
logger.exception("get_metas failed")
89+
raise HTTPException(status_code=500, detail=str(e)) from e
90+
91+
@app.post("/v1/metas")
92+
async def load_metas(metas: dict[int, MemoryBufferMetaList]) -> Response:
93+
return wrap_exception(lambda: ps.load_metas(metas))
94+
8295
@app.post("/v1/checkpoints/{checkpoint_name}/update")
8396
async def update(checkpoint_name: str, req: UpdateRequest) -> Response:
8497
def update_func(socket_paths: list[tuple[str, str]]):

‎examples/update.py‎

Lines changed: 32 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
import argparse
22
import json
33
import os
4-
import pickle
54
import time
65
from collections import defaultdict
76
from collections.abc import Callable
@@ -11,13 +10,18 @@
1110
import httpx
1211
import torch
1312
from loguru import logger
13+
from pydantic import TypeAdapter
1414
from safetensors import safe_open
1515

1616
import checkpoint_engine.distributed as dist
1717
from checkpoint_engine import request_inference_to_update
18+
from checkpoint_engine.data_types import MemoryBufferMetaList
1819
from checkpoint_engine.ps import ParameterServer
1920

2021

22+
_METAS_ADAPTER = TypeAdapter(dict[int, MemoryBufferMetaList])
23+
24+
2125
@contextmanager
2226
def timer(msg: str):
2327
start = time.perf_counter()
@@ -110,7 +114,7 @@ def update_weights(
110114
ps.gather_metas(checkpoint_name)
111115
if save_metas_file and int(os.getenv("RANK")) == 0:
112116
with open(save_metas_file, "wb") as f:
113-
pickle.dump(ps.get_metas(), f)
117+
f.write(_METAS_ADAPTER.dump_json(ps.get_metas()))
114118

115119
if update_method == "broadcast" or update_method == "all":
116120
with timer("Update weights without setting ranks"):
@@ -127,15 +131,22 @@ def update_weights(
127131
def join(
128132
ps: ParameterServer,
129133
checkpoint_name: str,
130-
load_metas_file: str,
134+
load_metas_file: str | None,
135+
metas_url: str | None,
131136
req_func: Callable[[list[tuple[str, str]]], None],
132137
inference_parallel_size: int,
133138
endpoint: str,
134139
uds: str | None = None,
135140
):
136-
assert load_metas_file, "load_metas_file is required"
137-
with open(load_metas_file, "rb") as f:
138-
metas = pickle.load(f)
141+
if load_metas_file:
142+
with open(load_metas_file, "rb") as f:
143+
metas = _METAS_ADAPTER.validate_json(f.read())
144+
elif metas_url:
145+
resp = httpx.get(metas_url, timeout=300.0)
146+
resp.raise_for_status()
147+
metas = _METAS_ADAPTER.validate_json(resp.content)
148+
else:
149+
raise ValueError("either load_metas_file or metas_url is required")
139150
ps.init_process_group()
140151
check_vllm_ready(endpoint, inference_parallel_size, uds)
141152
dist.barrier()
@@ -152,7 +163,19 @@ def join(
152163
parser = argparse.ArgumentParser(description="Update weights example")
153164
parser.add_argument("--checkpoint-path", type=str, default=None)
154165
parser.add_argument("--save-metas-file", type=str, default=None)
155-
parser.add_argument("--load-metas-file", type=str, default=None)
166+
metas_src = parser.add_mutually_exclusive_group()
167+
metas_src.add_argument(
168+
"--load-metas-file",
169+
type=str,
170+
default=None,
171+
help="Path to a metas JSON file (triggers join mode)",
172+
)
173+
metas_src.add_argument(
174+
"--metas-url",
175+
type=str,
176+
default=None,
177+
help="HTTP URL returning a metas JSON (triggers join mode)",
178+
)
156179
parser.add_argument("--sleep-time", type=int, default=0)
157180
parser.add_argument("--endpoint", type=str, default="http://localhost:19730")
158181
parser.add_argument("--inference-parallel-size", type=int, default=8)
@@ -167,11 +190,12 @@ def join(
167190
req_func = req_inference(args.endpoint, args.inference_parallel_size, args.uds)
168191
dist.use_backend(args.custom_dist)
169192
ps = ParameterServer(auto_pg=True)
170-
if args.load_metas_file:
193+
if args.load_metas_file or args.metas_url:
171194
join(
172195
ps,
173196
args.checkpoint_name,
174197
args.load_metas_file,
198+
args.metas_url,
175199
req_func,
176200
args.inference_parallel_size,
177201
args.endpoint,

‎tests/test_api.py‎

Lines changed: 139 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,139 @@
1+
"""CPU-only tests for the metas endpoints in api.py."""
2+
3+
from unittest.mock import MagicMock
4+
5+
import pytest
6+
import torch
7+
from fastapi.testclient import TestClient
8+
from pydantic import TypeAdapter
9+
10+
from checkpoint_engine.api import _init_api
11+
from checkpoint_engine.data_types import (
12+
MemoryBufferMetaList,
13+
MemoryBufferMetas,
14+
ParameterMeta,
15+
)
16+
17+
18+
_METAS_ADAPTER = TypeAdapter(dict[int, MemoryBufferMetaList])
19+
20+
21+
def _make_meta(rdma_device: str, ip: str) -> MemoryBufferMetaList:
22+
return MemoryBufferMetaList(
23+
p2p_store_addr=f"{ip}:12345",
24+
rdma_device=rdma_device,
25+
memory_buffer_metas_list=[
26+
MemoryBufferMetas(
27+
metas=[
28+
ParameterMeta(
29+
name="w",
30+
dtype=torch.float16,
31+
shape=torch.Size([2, 3]),
32+
aligned_size=12,
33+
)
34+
],
35+
ptr=0x12345678,
36+
size=1024,
37+
)
38+
],
39+
)
40+
41+
42+
@pytest.fixture
43+
def fake_metas() -> dict[int, MemoryBufferMetaList]:
44+
return {
45+
0: _make_meta("mlx5_0", "192.168.1.1"),
46+
1: _make_meta("mlx5_1", "192.168.1.1"),
47+
}
48+
49+
50+
@pytest.fixture
51+
def ps_mock(fake_metas: dict[int, MemoryBufferMetaList]) -> MagicMock:
52+
ps = MagicMock()
53+
ps.get_metas.return_value = fake_metas
54+
return ps
55+
56+
57+
def test_get_metas_returns_json(
58+
ps_mock: MagicMock, fake_metas: dict[int, MemoryBufferMetaList]
59+
) -> None:
60+
client = TestClient(_init_api(ps_mock))
61+
resp = client.get("/v1/metas")
62+
assert resp.status_code == 200
63+
assert resp.headers["content-type"] == "application/json"
64+
assert _METAS_ADAPTER.validate_json(resp.content) == fake_metas
65+
ps_mock.get_metas.assert_called_once_with()
66+
67+
68+
def test_get_metas_propagates_ps_error(ps_mock: MagicMock) -> None:
69+
ps_mock.get_metas.side_effect = RuntimeError("metas not gathered yet")
70+
client = TestClient(_init_api(ps_mock))
71+
resp = client.get("/v1/metas")
72+
assert resp.status_code == 500
73+
assert "metas not gathered yet" in resp.text
74+
75+
76+
def test_load_metas_decodes_and_calls_ps(
77+
ps_mock: MagicMock, fake_metas: dict[int, MemoryBufferMetaList]
78+
) -> None:
79+
client = TestClient(_init_api(ps_mock))
80+
resp = client.post(
81+
"/v1/metas",
82+
content=_METAS_ADAPTER.dump_json(fake_metas),
83+
headers={"content-type": "application/json"},
84+
)
85+
assert resp.status_code == 200
86+
ps_mock.load_metas.assert_called_once_with(fake_metas)
87+
88+
89+
def test_load_metas_rejects_bad_json(ps_mock: MagicMock) -> None:
90+
client = TestClient(_init_api(ps_mock))
91+
resp = client.post(
92+
"/v1/metas",
93+
content=b"not a valid json",
94+
headers={"content-type": "application/json"},
95+
)
96+
assert resp.status_code == 422
97+
ps_mock.load_metas.assert_not_called()
98+
99+
100+
def test_load_metas_rejects_schema_mismatch(ps_mock: MagicMock) -> None:
101+
"""JSON that parses but doesn't match MemoryBufferMetaList shape -> 422."""
102+
client = TestClient(_init_api(ps_mock))
103+
resp = client.post(
104+
"/v1/metas",
105+
content=b'{"0": {"foo": "bar"}}',
106+
headers={"content-type": "application/json"},
107+
)
108+
assert resp.status_code == 422
109+
ps_mock.load_metas.assert_not_called()
110+
111+
112+
def test_load_metas_propagates_ps_error(
113+
ps_mock: MagicMock, fake_metas: dict[int, MemoryBufferMetaList]
114+
) -> None:
115+
ps_mock.load_metas.side_effect = RuntimeError("rdma device mismatch")
116+
client = TestClient(_init_api(ps_mock))
117+
resp = client.post(
118+
"/v1/metas",
119+
content=_METAS_ADAPTER.dump_json(fake_metas),
120+
headers={"content-type": "application/json"},
121+
)
122+
assert resp.status_code == 500
123+
assert "rdma device mismatch" in resp.text
124+
125+
126+
def test_round_trip_get_then_load(
127+
ps_mock: MagicMock, fake_metas: dict[int, MemoryBufferMetaList]
128+
) -> None:
129+
"""JSON bytes returned by GET /v1/metas must be accepted by POST /v1/metas."""
130+
client = TestClient(_init_api(ps_mock))
131+
get_resp = client.get("/v1/metas")
132+
assert get_resp.status_code == 200
133+
load_resp = client.post(
134+
"/v1/metas",
135+
content=get_resp.content,
136+
headers={"content-type": "application/json"},
137+
)
138+
assert load_resp.status_code == 200
139+
ps_mock.load_metas.assert_called_once_with(fake_metas)

0 commit comments

Comments
 (0)