Skip to content

Commit 27c7ee2

Browse files
committed
feat(subsystembenchmarks): add Ray Data checkpoint load benchmark
1 parent aa92267 commit 27c7ee2

7 files changed

Lines changed: 399 additions & 79 deletions

File tree

gcsfs/tests/perf/subsystembenchmarks/checkpointing/ray_data/common.py

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22

33
import logging
44
import os
5+
import shutil
56
import socket
67
import tempfile
78

@@ -10,6 +11,8 @@
1011
import ray
1112
import torch
1213
import torch.distributed as dist
14+
import torch.distributed.checkpoint as dcp
15+
from torch.distributed.checkpoint.state_dict import StateDictOptions, get_state_dict
1316

1417
from gcsfs.tests.perf.subsystembenchmarks.dataloading.driver import assert_fsspec_gcsfs
1518

@@ -163,6 +166,15 @@ def parallelize_model(model, params):
163166
raise ValueError(f"Unknown strategy: {strategy}")
164167

165168

169+
def setup_model_and_optimizer(params):
170+
"""Loads, parallelizes the benchmark model and materializes AdamW optimizer states."""
171+
model = load_benchmark_model(params)
172+
model = parallelize_model(model, params)
173+
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
174+
materialize_adamw_states(optimizer)
175+
return model, optimizer
176+
177+
166178
def resolve_storage(prefix: str):
167179
"""Resolves fsspec and PyArrow filesystems for checkpoint storage."""
168180
fs, base_path = fsspec.core.url_to_fs(prefix)
@@ -201,3 +213,64 @@ def _get_staging_dir(prefix: str, min_free_gb: float = 50.0) -> str:
201213
except Exception:
202214
pass
203215
return tempfile.mkdtemp(prefix=prefix)
216+
217+
218+
def save_checkpoint_step(
219+
model,
220+
optimizer,
221+
params,
222+
rank: int,
223+
arrow_fs,
224+
fs,
225+
destination_ckpt: str,
226+
staging_prefix: str = "ray-ckpt",
227+
):
228+
"""Performs a single distributed or single-node checkpoint save to storage."""
229+
is_sharded = params.strategy in (
230+
"fsdp_sharded",
231+
"model_parallel_sharded",
232+
)
233+
options = StateDictOptions(
234+
full_state_dict=not is_sharded,
235+
cpu_offload=not is_sharded,
236+
)
237+
238+
local_dir = None
239+
try:
240+
model_state, opt_state = get_state_dict(
241+
model, optimizer, options=options
242+
)
243+
app_state = {"model": model_state, "optimizer": opt_state}
244+
245+
if is_sharded:
246+
local_dir = _get_staging_dir(f"{staging_prefix}-rank{rank}-")
247+
dcp.save(
248+
{"app": app_state},
249+
storage_writer=dcp.FileSystemWriter(local_dir),
250+
)
251+
arrow_fs.create_dir(destination_ckpt)
252+
_pyarrow_fs_copy_files(
253+
local_dir,
254+
destination_ckpt,
255+
destination_filesystem=arrow_fs,
256+
)
257+
else:
258+
if rank == 0:
259+
local_dir = _get_staging_dir(f"{staging_prefix}-rank0-")
260+
ckpt_file = os.path.join(local_dir, "checkpoint.pt")
261+
torch.save(app_state, ckpt_file)
262+
arrow_fs.create_dir(destination_ckpt)
263+
_pyarrow_fs_copy_files(
264+
local_dir,
265+
destination_ckpt,
266+
destination_filesystem=arrow_fs,
267+
)
268+
269+
del app_state, model_state, opt_state
270+
dist.barrier()
271+
if is_sharded and rank == 0:
272+
fs.touch(f"{destination_ckpt.rstrip('/')}/_SUCCESS")
273+
finally:
274+
if local_dir and os.path.exists(local_dir):
275+
shutil.rmtree(local_dir, ignore_errors=True)
276+

gcsfs/tests/perf/subsystembenchmarks/checkpointing/ray_data/configs.yaml

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,3 +15,16 @@ scenarios:
1515
- {axis: "strategy", strategy: "fsdp_full", world_size: 8}
1616
- {axis: "strategy", strategy: "model_parallel_full", tensor_parallel_size: 4, data_parallel_size: 2, world_size: 8}
1717
- {axis: "strategy", strategy: "model_parallel_sharded", tensor_parallel_size: 4, data_parallel_size: 2, world_size: 8}
18+
- name: "checkpoint_read"
19+
scenario: "checkpoint_read"
20+
variants:
21+
# Strategy variants
22+
- {axis: "strategy", strategy: "ddp", world_size: 8}
23+
- {axis: "strategy", strategy: "fsdp_sharded", world_size: 8}
24+
- {axis: "strategy", strategy: "fsdp_full", world_size: 8}
25+
- {axis: "strategy", strategy: "model_parallel_full", tensor_parallel_size: 4, data_parallel_size: 2, world_size: 8}
26+
- {axis: "strategy", strategy: "model_parallel_sharded", tensor_parallel_size: 4, data_parallel_size: 2, world_size: 8}
27+
# Cross-topology restore variants
28+
- {axis: "cross_size", strategy: "fsdp_sharded", setup_world_size: 4, world_size: 8}
29+
- {axis: "cross_size", strategy: "fsdp_sharded", setup_world_size: 8, world_size: 4}
30+
- {axis: "cross_size", strategy: "model_parallel_sharded", setup_tensor_parallel_size: 4, setup_data_parallel_size: 2, setup_world_size: 8, tensor_parallel_size: 2, data_parallel_size: 2, world_size: 4}
Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
"""Ray checkpoint load benchmark."""
Lines changed: 237 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,237 @@
1+
"""Driver for Ray checkpoint load benchmark."""
2+
3+
import contextlib
4+
import dataclasses
5+
import os
6+
import time
7+
import uuid
8+
9+
import ray
10+
import ray.train
11+
import torch
12+
import torch.distributed as dist
13+
import torch.distributed.checkpoint as dcp
14+
from torch.distributed.checkpoint.state_dict import (
15+
StateDictOptions,
16+
get_state_dict,
17+
set_model_state_dict,
18+
set_optimizer_state_dict,
19+
set_state_dict,
20+
)
21+
22+
from gcsfs.tests.perf.subsystembenchmarks.checkpointing.driver import (
23+
CheckpointDriver,
24+
CheckpointResult,
25+
)
26+
from gcsfs.tests.perf.subsystembenchmarks.checkpointing.ray_data.common import (
27+
ensure_ray_initialized,
28+
find_free_port,
29+
resolve_storage,
30+
save_checkpoint_step,
31+
setup_distributed_env,
32+
setup_model_and_optimizer,
33+
)
34+
35+
36+
@ray.remote
37+
class RayCheckpointSetupWorker:
38+
"""Ray Actor that creates the initial checkpoint on storage for load benchmarks."""
39+
40+
def __init__(self, rank, world_size, port, prefix, params):
41+
self.rank = rank
42+
self.world_size = world_size
43+
self.port = port
44+
self.prefix = prefix
45+
self.params = params
46+
47+
def setup_and_save(self):
48+
try:
49+
setup_distributed_env(self.rank, self.world_size, self.port)
50+
model, optimizer = setup_model_and_optimizer(self.params)
51+
fs, arrow_fs, base_path = resolve_storage(self.prefix)
52+
destination_ckpt = f"{base_path.rstrip('/')}/model.ckpt"
53+
save_checkpoint_step(
54+
model,
55+
optimizer,
56+
self.params,
57+
self.rank,
58+
arrow_fs,
59+
fs,
60+
destination_ckpt,
61+
staging_prefix="ray-setup-ckpt",
62+
)
63+
finally:
64+
dist.destroy_process_group()
65+
66+
67+
@ray.remote
68+
class RayCheckpointLoadWorker:
69+
"""Ray Actor executing distributed checkpoint load operations on CPU."""
70+
71+
def __init__(self, rank, world_size, port, prefix, params):
72+
self.rank = rank
73+
self.world_size = world_size
74+
self.port = port
75+
self.prefix = prefix
76+
self.params = params
77+
78+
def setup(self):
79+
setup_distributed_env(self.rank, self.world_size, self.port)
80+
self.model, self.optimizer = setup_model_and_optimizer(self.params)
81+
self.fs, self.arrow_fs, self.base_path = resolve_storage(self.prefix)
82+
self.destination_ckpt = f"{self.base_path.rstrip('/')}/model.ckpt"
83+
84+
def load_rounds(self):
85+
try:
86+
durations = []
87+
is_sharded = self.params.strategy in (
88+
"fsdp_sharded",
89+
"model_parallel_sharded",
90+
)
91+
for round_idx in range(self.params.rounds):
92+
dist.barrier()
93+
t_start = time.perf_counter()
94+
95+
checkpoint = ray.train.Checkpoint(
96+
path=self.destination_ckpt, filesystem=self.arrow_fs
97+
)
98+
# Ensure all workers on the host share the same UUID for this round
99+
# so Ray's built-in file locking deduplicates the download across workers
100+
# and cleans up the shared temporary directory after all workers exit.
101+
checkpoint._uuid = uuid.uuid5(
102+
uuid.NAMESPACE_URL, f"{self.destination_ckpt}-round-{round_idx}"
103+
)
104+
105+
directory_context = (
106+
checkpoint.as_directory()
107+
if is_sharded or self.rank == 0
108+
else contextlib.nullcontext(None)
109+
)
110+
with directory_context as local_dir:
111+
if is_sharded:
112+
options = StateDictOptions(
113+
full_state_dict=False,
114+
cpu_offload=False,
115+
)
116+
model_state, opt_state = get_state_dict(
117+
self.model, self.optimizer, options=options
118+
)
119+
app_state = {"model": model_state, "optimizer": opt_state}
120+
dcp.load(
121+
{"app": app_state},
122+
storage_reader=dcp.FileSystemReader(local_dir),
123+
)
124+
set_state_dict(
125+
self.model,
126+
self.optimizer,
127+
model_state_dict=app_state["model"],
128+
optim_state_dict=app_state["optimizer"],
129+
options=options,
130+
)
131+
del app_state, model_state, opt_state
132+
else:
133+
options = StateDictOptions(
134+
full_state_dict=True,
135+
cpu_offload=False,
136+
broadcast_from_rank0=True,
137+
)
138+
if self.rank == 0:
139+
ckpt_file = os.path.join(local_dir, "checkpoint.pt")
140+
state = torch.load(
141+
ckpt_file, map_location="cpu", weights_only=False
142+
)
143+
model_state = state["model"]
144+
opt_state = state["optimizer"]
145+
else:
146+
model_state = {}
147+
opt_state = {}
148+
149+
set_model_state_dict(
150+
self.model,
151+
model_state,
152+
options=options,
153+
)
154+
set_optimizer_state_dict(
155+
self.model,
156+
self.optimizer,
157+
opt_state,
158+
options=options,
159+
)
160+
if self.rank == 0:
161+
del state, model_state, opt_state
162+
163+
dist.barrier()
164+
t_end = time.perf_counter()
165+
durations.append((t_start, t_end))
166+
167+
return durations
168+
finally:
169+
dist.destroy_process_group()
170+
171+
172+
def run_ray_load(prefix, params):
173+
"""Runs single or distributed checkpoint load benchmark across Ray actor workers."""
174+
ensure_ray_initialized()
175+
world_size = params.world_size
176+
port = find_free_port()
177+
178+
workers = [
179+
RayCheckpointLoadWorker.remote(rank, world_size, port, prefix, params)
180+
for rank in range(world_size)
181+
]
182+
ray.get([w.setup.remote() for w in workers])
183+
results = ray.get([w.load_rounds.remote() for w in workers])
184+
185+
durations = []
186+
for r in range(params.rounds):
187+
begins = [results[rank][r][0] for rank in range(world_size)]
188+
ends = [results[rank][r][1] for rank in range(world_size)]
189+
durations.append(max(ends) - min(begins))
190+
return durations
191+
192+
193+
class RayCheckpointReadDriver(CheckpointDriver):
194+
"""Driver for Ray checkpoint load benchmarks."""
195+
196+
def setup(self, prefix: str, params):
197+
"""Generates the source checkpoint on storage using setup topology."""
198+
ensure_ray_initialized()
199+
setup_world_size = (
200+
getattr(params, "setup_world_size", None) or params.world_size
201+
)
202+
setup_tp = (
203+
getattr(params, "setup_tensor_parallel_size", None)
204+
or params.tensor_parallel_size
205+
)
206+
setup_dp = (
207+
getattr(params, "setup_data_parallel_size", None)
208+
or params.data_parallel_size
209+
)
210+
211+
setup_params = dataclasses.replace(
212+
params,
213+
world_size=setup_world_size,
214+
tensor_parallel_size=setup_tp,
215+
data_parallel_size=setup_dp,
216+
)
217+
port = find_free_port()
218+
try:
219+
workers = [
220+
RayCheckpointSetupWorker.remote(
221+
rank, setup_world_size, port, prefix, setup_params
222+
)
223+
for rank in range(setup_world_size)
224+
]
225+
ray.get([w.setup_and_save.remote() for w in workers])
226+
finally:
227+
if ray.is_initialized():
228+
ray.shutdown()
229+
230+
def run(self, prefix: str, params) -> CheckpointResult:
231+
"""Executes the checkpoint load benchmark across Ray actor workers."""
232+
try:
233+
durations = run_ray_load(prefix, params)
234+
return CheckpointResult(durations=durations)
235+
finally:
236+
if ray.is_initialized():
237+
ray.shutdown()
Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,35 @@
1+
import os
2+
3+
import pytest
4+
5+
from gcsfs.tests.perf.subsystembenchmarks.checkpointing.ray_data import configs
6+
from gcsfs.tests.perf.subsystembenchmarks.checkpointing.ray_data.configs import (
7+
RayCheckpointConfigurator,
8+
)
9+
10+
pytest.importorskip("ray")
11+
pytest.importorskip("torch")
12+
13+
pytestmark = pytest.mark.skipif(
14+
not os.environ.get("GCSFS_SUBSYSTEM_BUCKET_PREFIX"),
15+
reason="the checkpoint benchmarks create a bucket per case; CI-only (run.py exports the prefix)",
16+
)
17+
18+
CASES = [
19+
c
20+
for c in RayCheckpointConfigurator(configs.__file__).generate_cases()
21+
if c.scenario == "checkpoint_read"
22+
]
23+
24+
25+
@pytest.mark.timeout(7200)
26+
@pytest.mark.parametrize("params", CASES, ids=lambda p: p.name)
27+
def test_checkpoint_load(benchmark, params, monitor):
28+
from gcsfs.tests.perf.subsystembenchmarks.checkpointing.checkpoint_case import (
29+
run_checkpoint_case,
30+
)
31+
from gcsfs.tests.perf.subsystembenchmarks.checkpointing.ray_data.read.driver import (
32+
RayCheckpointReadDriver,
33+
)
34+
35+
run_checkpoint_case(benchmark, monitor, params, RayCheckpointReadDriver())

0 commit comments

Comments
 (0)