ravex 0.0.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ravex/__init__.py +142 -0
- ravex/_backends.py +458 -0
- ravex/_bootstrap.py +221 -0
- ravex/_cli.py +117 -0
- ravex/_config.py +392 -0
- ravex/_distributed.py +393 -0
- ravex/_frameworks.py +85 -0
- ravex/_patches.py +387 -0
- ravex/_registry.py +791 -0
- ravex/_resume.py +121 -0
- ravex/_runtime.py +483 -0
- ravex/_sampler.py +305 -0
- ravex/py.typed +0 -0
- ravex-0.0.2.dist-info/METADATA +253 -0
- ravex-0.0.2.dist-info/RECORD +20 -0
- ravex-0.0.2.dist-info/WHEEL +5 -0
- ravex-0.0.2.dist-info/entry_points.txt +2 -0
- ravex-0.0.2.dist-info/licenses/LICENSE +201 -0
- ravex-0.0.2.dist-info/licenses/NOTICE +7 -0
- ravex-0.0.2.dist-info/top_level.txt +1 -0
ravex/__init__.py
ADDED
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
"""Ravex — transparent checkpointing for PyTorch training.
|
|
2
|
+
|
|
3
|
+
Two ways in.
|
|
4
|
+
|
|
5
|
+
**Zero code changes** (what the project is for). Install the autoloader once,
|
|
6
|
+
drop a ``ravex.yaml`` next to your code, and run your script unchanged::
|
|
7
|
+
|
|
8
|
+
ravex enable
|
|
9
|
+
python train.py
|
|
10
|
+
|
|
11
|
+
**One line**, when you would rather be explicit::
|
|
12
|
+
|
|
13
|
+
import ravex
|
|
14
|
+
ravex.activate()
|
|
15
|
+
|
|
16
|
+
Either way, from that point on every ``optimizer.step()`` is counted, a
|
|
17
|
+
checkpoint is written every ``checkpoint_every`` steps, and rerunning the same
|
|
18
|
+
script picks up where the last one stopped — model, optimizer, LR scheduler,
|
|
19
|
+
AMP scaler, RNG state and dataset position included.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
from __future__ import annotations
|
|
23
|
+
|
|
24
|
+
from typing import Any, Dict, Optional
|
|
25
|
+
|
|
26
|
+
#: The one place the version is written. `pyproject.toml` reads it from here
|
|
27
|
+
#: (`[tool.setuptools.dynamic]`), because declaring it in both is a pair that
|
|
28
|
+
#: drifts at the first bump — and `ravex status` reports this one, so the drift
|
|
29
|
+
#: shows up as the CLI stating a version nobody installed.
|
|
30
|
+
__version__ = "0.0.2"
|
|
31
|
+
|
|
32
|
+
__all__ = [
|
|
33
|
+
"activate",
|
|
34
|
+
"checkpoint",
|
|
35
|
+
"deactivate",
|
|
36
|
+
"flush",
|
|
37
|
+
"is_active",
|
|
38
|
+
"status",
|
|
39
|
+
"step",
|
|
40
|
+
"__version__",
|
|
41
|
+
]
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def activate(**overrides: Any):
|
|
45
|
+
"""Install the PyTorch patches and start checkpointing.
|
|
46
|
+
|
|
47
|
+
Keyword arguments override the resolved configuration, e.g.
|
|
48
|
+
``ravex.activate(checkpoint_every=100, backend="torch_save")``.
|
|
49
|
+
Calling this twice is harmless.
|
|
50
|
+
"""
|
|
51
|
+
from ravex._config import RavexConfig
|
|
52
|
+
from ravex._runtime import get_runtime
|
|
53
|
+
|
|
54
|
+
runtime = get_runtime(create=False)
|
|
55
|
+
if runtime is None:
|
|
56
|
+
config = RavexConfig.load()
|
|
57
|
+
for key, value in overrides.items():
|
|
58
|
+
if not hasattr(config, key):
|
|
59
|
+
raise TypeError(f"unknown configuration option {key!r}")
|
|
60
|
+
setattr(config, key, value)
|
|
61
|
+
config._normalize()
|
|
62
|
+
|
|
63
|
+
from ravex._runtime import RavexRuntime
|
|
64
|
+
import ravex._runtime as runtime_module
|
|
65
|
+
|
|
66
|
+
runtime = RavexRuntime(config)
|
|
67
|
+
runtime_module._runtime = runtime
|
|
68
|
+
runtime.activate()
|
|
69
|
+
|
|
70
|
+
import atexit
|
|
71
|
+
|
|
72
|
+
atexit.register(runtime.shutdown)
|
|
73
|
+
elif overrides:
|
|
74
|
+
raise RuntimeError(
|
|
75
|
+
"Ravex is already active; configuration cannot be changed after "
|
|
76
|
+
"activation. Set the options in ravex.yaml or RAVEX_* env vars."
|
|
77
|
+
)
|
|
78
|
+
return runtime
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def deactivate() -> None:
|
|
82
|
+
"""Stop checkpointing and remove the patches. Mostly for tests."""
|
|
83
|
+
from ravex._runtime import get_runtime, reset_runtime
|
|
84
|
+
|
|
85
|
+
runtime = get_runtime(create=False)
|
|
86
|
+
if runtime is not None:
|
|
87
|
+
runtime.shutdown()
|
|
88
|
+
reset_runtime()
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def is_active() -> bool:
|
|
92
|
+
from ravex._runtime import get_runtime
|
|
93
|
+
|
|
94
|
+
runtime = get_runtime(create=False)
|
|
95
|
+
return runtime is not None and runtime.enabled
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def checkpoint() -> bool:
|
|
99
|
+
"""Force a checkpoint now, outside the normal cadence."""
|
|
100
|
+
from ravex._runtime import get_runtime
|
|
101
|
+
|
|
102
|
+
runtime = get_runtime(create=False)
|
|
103
|
+
if runtime is None:
|
|
104
|
+
return False
|
|
105
|
+
return runtime.checkpoint()
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def flush() -> None:
|
|
109
|
+
"""Block until every pending checkpoint write has landed."""
|
|
110
|
+
from ravex._runtime import get_runtime
|
|
111
|
+
|
|
112
|
+
runtime = get_runtime(create=False)
|
|
113
|
+
if runtime is not None:
|
|
114
|
+
runtime.flush()
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def step() -> int:
|
|
118
|
+
"""Current optimizer-step count as Ravex sees it."""
|
|
119
|
+
from ravex._runtime import get_runtime
|
|
120
|
+
|
|
121
|
+
runtime = get_runtime(create=False)
|
|
122
|
+
return runtime.step if runtime is not None else 0
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def status() -> Dict[str, Optional[Any]]:
|
|
126
|
+
"""Snapshot of what the runtime is doing, for debugging."""
|
|
127
|
+
from ravex._runtime import get_runtime
|
|
128
|
+
|
|
129
|
+
runtime = get_runtime(create=False)
|
|
130
|
+
if runtime is None:
|
|
131
|
+
return {"active": False}
|
|
132
|
+
|
|
133
|
+
return {
|
|
134
|
+
"active": runtime.enabled,
|
|
135
|
+
"step": runtime.step,
|
|
136
|
+
"models": len(runtime.registry.models),
|
|
137
|
+
"optimizers": len(runtime.registry.optimizers),
|
|
138
|
+
"dataloaders": len(runtime.registry.dataloaders),
|
|
139
|
+
"resumed": runtime.registry.resumed,
|
|
140
|
+
"backend": type(runtime.backend).__name__ if runtime.backend else None,
|
|
141
|
+
"config": runtime.config.describe(),
|
|
142
|
+
}
|
ravex/_backends.py
ADDED
|
@@ -0,0 +1,458 @@
|
|
|
1
|
+
"""Checkpoint backends.
|
|
2
|
+
|
|
3
|
+
A backend takes the state dict the registry collected and gets it durably
|
|
4
|
+
stored, without blocking the training loop for longer than the copy takes.
|
|
5
|
+
|
|
6
|
+
Two implementations ship with Ravex:
|
|
7
|
+
|
|
8
|
+
``moonclip``
|
|
9
|
+
The default. Per-tensor delta tracking, zstd, S3/R2 sync, background
|
|
10
|
+
writer. State is flattened into individual tensors so unchanged weights
|
|
11
|
+
cost zero I/O on the next checkpoint.
|
|
12
|
+
|
|
13
|
+
``torch_save``
|
|
14
|
+
The fallback for environments without Moonclip. Plain ``torch.save`` on a
|
|
15
|
+
single background thread. Correct, just slower and much larger on disk.
|
|
16
|
+
|
|
17
|
+
Both copy the state to CPU memory *before* returning from ``save``. That copy
|
|
18
|
+
is the whole point: it lets the training loop mutate the live tensors while the
|
|
19
|
+
writer is still working, without the writer reading half-updated weights.
|
|
20
|
+
|
|
21
|
+
``save`` reports how long each phase of that handoff took, and the runtime puts
|
|
22
|
+
those numbers in its per-checkpoint log line. On 8× RTX 5060 Ti the handoff came
|
|
23
|
+
out at 10.6 s against 1.5 s of state collection, and three A/B runs against the
|
|
24
|
+
suspects — thread count, compression level, checkpoint cadence — each moved it
|
|
25
|
+
by under a second. Guessing has a poor record here; the breakdown is cheap
|
|
26
|
+
enough to always be on.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
from __future__ import annotations
|
|
30
|
+
|
|
31
|
+
import glob
|
|
32
|
+
import inspect
|
|
33
|
+
import logging
|
|
34
|
+
import os
|
|
35
|
+
import re
|
|
36
|
+
import time
|
|
37
|
+
from abc import ABC, abstractmethod
|
|
38
|
+
from concurrent.futures import Future, ThreadPoolExecutor
|
|
39
|
+
from typing import Any, Dict, Optional
|
|
40
|
+
|
|
41
|
+
logger = logging.getLogger("ravex")
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class CheckpointBackend(ABC):
|
|
45
|
+
"""Storage interface used by the runtime."""
|
|
46
|
+
|
|
47
|
+
@abstractmethod
|
|
48
|
+
def save(
|
|
49
|
+
self, step: int, state: Dict[str, Any], metadata: Dict[str, str]
|
|
50
|
+
) -> Dict[str, float]:
|
|
51
|
+
"""Persist ``state``. Returns once the state has been copied, not
|
|
52
|
+
once it has been written.
|
|
53
|
+
|
|
54
|
+
The return value is how long each phase of that copy took, in seconds,
|
|
55
|
+
in the order the phases ran. It exists so the runtime's log line can say
|
|
56
|
+
*where* a slow handoff went instead of only how long it took — the
|
|
57
|
+
difference between a number that ends an investigation and one that
|
|
58
|
+
starts another.
|
|
59
|
+
"""
|
|
60
|
+
|
|
61
|
+
@abstractmethod
|
|
62
|
+
def load_latest(self) -> Optional[Dict[str, Any]]:
|
|
63
|
+
"""Return the most recent checkpoint, or None if there is none."""
|
|
64
|
+
|
|
65
|
+
def latest_step(self) -> Optional[int]:
|
|
66
|
+
"""Step of the newest stored checkpoint, without loading its tensors."""
|
|
67
|
+
return None
|
|
68
|
+
|
|
69
|
+
def load_step(self, step: int) -> Optional[Dict[str, Any]]:
|
|
70
|
+
"""Return the checkpoint written at ``step``, or None.
|
|
71
|
+
|
|
72
|
+
Only per-rank checkpointing needs this: the ranks have to agree on a
|
|
73
|
+
step every one of them holds, and "the latest" is not that step when a
|
|
74
|
+
kill landed between two ranks' writes.
|
|
75
|
+
"""
|
|
76
|
+
return None
|
|
77
|
+
|
|
78
|
+
@abstractmethod
|
|
79
|
+
def has_checkpoint(self) -> bool:
|
|
80
|
+
"""Whether a resumable checkpoint exists. Cheap; no tensor loading."""
|
|
81
|
+
|
|
82
|
+
@abstractmethod
|
|
83
|
+
def flush(self) -> None:
|
|
84
|
+
"""Block until every pending write has completed."""
|
|
85
|
+
|
|
86
|
+
def close(self) -> None:
|
|
87
|
+
self.flush()
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
# ─── moonclip ───────────────────────────────────────────────────────
|
|
91
|
+
|
|
92
|
+
#: Single prefix for the whole state tree. Tensor names come out as
|
|
93
|
+
#: ``ravex/models/<key>/<param>``, stable across steps, which is exactly
|
|
94
|
+
#: what Moonclip's per-tensor delta tracking keys on.
|
|
95
|
+
_PREFIX = "ravex"
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
class MoonclipBackend(CheckpointBackend):
|
|
99
|
+
"""Default backend, built on the Moonclip checkpoint engine."""
|
|
100
|
+
|
|
101
|
+
def __init__(self, config):
|
|
102
|
+
import moonclip
|
|
103
|
+
|
|
104
|
+
self._moonclip = moonclip
|
|
105
|
+
storage = config.storage
|
|
106
|
+
|
|
107
|
+
kwargs: Dict[str, Any] = {
|
|
108
|
+
"storage_root": storage.path,
|
|
109
|
+
"compression_level": config.compression_level,
|
|
110
|
+
"max_total_snapshots": config.keep_last,
|
|
111
|
+
"async_save": True,
|
|
112
|
+
# Always single-rank, stated explicitly. Moonclip otherwise infers
|
|
113
|
+
# world_size from RANK/WORLD_SIZE in the environment, and under
|
|
114
|
+
# torchrun it then rejects the single-rank save API outright:
|
|
115
|
+
# "Multi-rank save requires explicit create_snapshot/save_rank/
|
|
116
|
+
# finalize flow". Ravex does not need that flow — sharded state is
|
|
117
|
+
# gathered before it gets here and exactly one rank writes — but
|
|
118
|
+
# the mismatch is silent apart from a log line, so every
|
|
119
|
+
# distributed run would lose checkpointing altogether.
|
|
120
|
+
"world_size": 1,
|
|
121
|
+
"rank": 0,
|
|
122
|
+
}
|
|
123
|
+
|
|
124
|
+
if not config.delta:
|
|
125
|
+
# No dedicated switch in Moonclip: a full snapshot on every step is
|
|
126
|
+
# exactly "delta disabled".
|
|
127
|
+
kwargs["full_every_steps"] = 1
|
|
128
|
+
|
|
129
|
+
if storage.is_remote:
|
|
130
|
+
if not (storage.access_key and storage.secret_key):
|
|
131
|
+
raise ValueError(
|
|
132
|
+
"storage.type is %r but no credentials were found. Set "
|
|
133
|
+
"RAVEX_S3_ACCESS_KEY / RAVEX_S3_SECRET_KEY (or "
|
|
134
|
+
"AWS_ACCESS_KEY_ID / AWS_SECRET_ACCESS_KEY)." % storage.type
|
|
135
|
+
)
|
|
136
|
+
kwargs.update(
|
|
137
|
+
s3_bucket=storage.bucket,
|
|
138
|
+
s3_region=storage.region,
|
|
139
|
+
s3_prefix=storage.prefix,
|
|
140
|
+
s3_endpoint=storage.endpoint,
|
|
141
|
+
s3_access_key=storage.access_key,
|
|
142
|
+
s3_secret_key=storage.secret_key,
|
|
143
|
+
s3_path_style=storage.path_style,
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
self._manager = moonclip.CheckpointManager(**kwargs)
|
|
147
|
+
|
|
148
|
+
# `as_tensors` arrived after the first released Moonclip. Passing it to
|
|
149
|
+
# a build that predates it is a TypeError in the middle of a training
|
|
150
|
+
# run, which is exactly the failure mode this backend exists to avoid,
|
|
151
|
+
# so ask once here rather than guess later.
|
|
152
|
+
try:
|
|
153
|
+
self._as_tensors = "as_tensors" in inspect.signature(
|
|
154
|
+
moonclip.flatten_state_dict
|
|
155
|
+
).parameters
|
|
156
|
+
except (TypeError, ValueError): # pragma: no cover - defensive
|
|
157
|
+
self._as_tensors = False
|
|
158
|
+
if not self._as_tensors:
|
|
159
|
+
logger.info(
|
|
160
|
+
"Moonclip predates flatten_state_dict(as_tensors=); checkpoints "
|
|
161
|
+
"will block the training loop for roughly 5x longer per save"
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
def save(
|
|
165
|
+
self, step: int, state: Dict[str, Any], metadata: Dict[str, str]
|
|
166
|
+
) -> Dict[str, float]:
|
|
167
|
+
# `as_tensors` hands Moonclip the tensors and lets it take the shadow
|
|
168
|
+
# copy itself, on all cores with the GIL released. Flattening to bytes
|
|
169
|
+
# here instead cost two serial copies of the whole state — one in
|
|
170
|
+
# `.tobytes()`, one on the way into Rust — and the training loop was
|
|
171
|
+
# blocked for both: 1393 ms against 271 ms on 3.8 GiB of weights.
|
|
172
|
+
#
|
|
173
|
+
# The copy is still taken before this method returns, which is the
|
|
174
|
+
# contract this class documents. It just happens one line later, inside
|
|
175
|
+
# `save_raw`. That matters because for a model already on CPU the dict
|
|
176
|
+
# below aliases live parameter memory rather than owning a copy of it.
|
|
177
|
+
started = time.perf_counter()
|
|
178
|
+
if self._as_tensors:
|
|
179
|
+
tensors, _ = self._moonclip.flatten_state_dict(
|
|
180
|
+
state, _PREFIX, as_tensors=True
|
|
181
|
+
)
|
|
182
|
+
else:
|
|
183
|
+
tensors, _ = self._moonclip.flatten_state_dict(state, _PREFIX)
|
|
184
|
+
flattened = time.perf_counter()
|
|
185
|
+
self._manager.save_raw(step=step, tensors=tensors, metadata=metadata)
|
|
186
|
+
# `store` is not the write: that runs in the background. It is the
|
|
187
|
+
# shadow copy, plus however long the previous checkpoint's writer still
|
|
188
|
+
# needed — Moonclip allows one save in flight, so a writer that has not
|
|
189
|
+
# drained is backpressure landing on this line. `MOONCLIP_PROFILE=1`
|
|
190
|
+
# separates the two.
|
|
191
|
+
return {
|
|
192
|
+
"flatten": flattened - started,
|
|
193
|
+
"store": time.perf_counter() - flattened,
|
|
194
|
+
}
|
|
195
|
+
|
|
196
|
+
def load_latest(self) -> Optional[Dict[str, Any]]:
|
|
197
|
+
snapshots = self._manager.list_snapshots()
|
|
198
|
+
if not snapshots:
|
|
199
|
+
return None
|
|
200
|
+
_, loaded = self._manager.load_latest()
|
|
201
|
+
state = loaded.get(_PREFIX)
|
|
202
|
+
if state is None:
|
|
203
|
+
logger.warning(
|
|
204
|
+
"Latest snapshot has no %r payload - it was probably written by "
|
|
205
|
+
"something other than Ravex. Ignoring it.",
|
|
206
|
+
_PREFIX,
|
|
207
|
+
)
|
|
208
|
+
return None
|
|
209
|
+
return state
|
|
210
|
+
|
|
211
|
+
def latest_step(self) -> Optional[int]:
|
|
212
|
+
try:
|
|
213
|
+
snapshots = self._manager.list_snapshots()
|
|
214
|
+
except Exception as exc: # pragma: no cover - defensive
|
|
215
|
+
logger.warning("Could not list snapshots: %s", exc)
|
|
216
|
+
return None
|
|
217
|
+
steps = [int(s["step"]) for s in snapshots if s.get("step") is not None]
|
|
218
|
+
return max(steps) if steps else None
|
|
219
|
+
|
|
220
|
+
def load_step(self, step: int) -> Optional[Dict[str, Any]]:
|
|
221
|
+
try:
|
|
222
|
+
snapshots = self._manager.list_snapshots()
|
|
223
|
+
except Exception as exc: # pragma: no cover - defensive
|
|
224
|
+
logger.warning("Could not list snapshots: %s", exc)
|
|
225
|
+
return None
|
|
226
|
+
# Newest first: a step can appear more than once if a run was restarted
|
|
227
|
+
# and rewrote it, and the last write is the one that counts.
|
|
228
|
+
for snapshot in reversed(snapshots):
|
|
229
|
+
if int(snapshot.get("step", -1)) != step:
|
|
230
|
+
continue
|
|
231
|
+
loaded = self._manager.load(snapshot["id"])
|
|
232
|
+
return loaded.get(_PREFIX)
|
|
233
|
+
return None
|
|
234
|
+
|
|
235
|
+
def has_checkpoint(self) -> bool:
|
|
236
|
+
try:
|
|
237
|
+
return bool(self._manager.list_snapshots())
|
|
238
|
+
except Exception as exc: # pragma: no cover - defensive
|
|
239
|
+
logger.warning("Could not list snapshots: %s", exc)
|
|
240
|
+
return False
|
|
241
|
+
|
|
242
|
+
def flush(self) -> None:
|
|
243
|
+
self._manager.flush()
|
|
244
|
+
|
|
245
|
+
def close(self) -> None:
|
|
246
|
+
self.flush()
|
|
247
|
+
if getattr(self, "_manager", None) is not None:
|
|
248
|
+
try:
|
|
249
|
+
self._manager.sync_now()
|
|
250
|
+
except Exception as exc:
|
|
251
|
+
logger.warning("Final remote sync failed: %s", exc)
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
# ─── torch.save fallback ────────────────────────────────────────────
|
|
255
|
+
|
|
256
|
+
_STEP_FILE = re.compile(r"step_(\d+)\.pt$")
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
def _step_of(path: str) -> int:
|
|
260
|
+
"""The step a checkpoint filename encodes."""
|
|
261
|
+
match = _STEP_FILE.search(path)
|
|
262
|
+
if match is None: # pragma: no cover - _files() only yields names that match
|
|
263
|
+
raise ValueError("not a checkpoint filename: %s" % path)
|
|
264
|
+
return int(match.group(1))
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def _cpu_copy(value: Any) -> Any:
|
|
268
|
+
"""Deep-copy a state tree onto CPU memory.
|
|
269
|
+
|
|
270
|
+
``to("cpu", copy=True)`` rather than ``.cpu()``: for a model already on CPU
|
|
271
|
+
the latter is a no-op and the writer would race the next training step.
|
|
272
|
+
"""
|
|
273
|
+
import torch
|
|
274
|
+
|
|
275
|
+
if torch.is_tensor(value):
|
|
276
|
+
return value.detach().to("cpu", copy=True)
|
|
277
|
+
if isinstance(value, dict):
|
|
278
|
+
return {k: _cpu_copy(v) for k, v in value.items()}
|
|
279
|
+
if isinstance(value, list):
|
|
280
|
+
return [_cpu_copy(v) for v in value]
|
|
281
|
+
if isinstance(value, tuple):
|
|
282
|
+
return tuple(_cpu_copy(v) for v in value)
|
|
283
|
+
return value
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
class TorchSaveBackend(CheckpointBackend):
|
|
287
|
+
"""Fallback backend: one ``.pt`` file per checkpoint."""
|
|
288
|
+
|
|
289
|
+
def __init__(self, config):
|
|
290
|
+
self.directory = config.storage.path
|
|
291
|
+
self.keep_last = config.keep_last
|
|
292
|
+
os.makedirs(self.directory, exist_ok=True)
|
|
293
|
+
self._executor = ThreadPoolExecutor(
|
|
294
|
+
max_workers=1, thread_name_prefix="ravex-writer"
|
|
295
|
+
)
|
|
296
|
+
self._pending: Optional[Future] = None
|
|
297
|
+
|
|
298
|
+
def save(
|
|
299
|
+
self, step: int, state: Dict[str, Any], metadata: Dict[str, str]
|
|
300
|
+
) -> Dict[str, float]:
|
|
301
|
+
import torch
|
|
302
|
+
|
|
303
|
+
started = time.perf_counter()
|
|
304
|
+
snapshot = _cpu_copy(state)
|
|
305
|
+
snapshot["_ravex_metadata"] = dict(metadata)
|
|
306
|
+
path = os.path.join(self.directory, f"step_{step:012d}.pt")
|
|
307
|
+
copied = time.perf_counter()
|
|
308
|
+
|
|
309
|
+
# Serialize writes: two concurrent torch.save calls on one disk are
|
|
310
|
+
# slower than one, and ordering matters for pruning.
|
|
311
|
+
self.flush()
|
|
312
|
+
try:
|
|
313
|
+
self._pending = self._executor.submit(self._write, torch, snapshot, path)
|
|
314
|
+
except RuntimeError:
|
|
315
|
+
# "cannot schedule new futures after shutdown": concurrent.futures
|
|
316
|
+
# registers its own atexit hook, and by the time ours runs the pool
|
|
317
|
+
# is closed. This is the checkpoint_on_exit path - the last one a
|
|
318
|
+
# normally-finishing run takes - so write it here instead of losing
|
|
319
|
+
# it. Nothing is racing us: the training loop is over.
|
|
320
|
+
self._write(torch, snapshot, path)
|
|
321
|
+
|
|
322
|
+
# `queue` is the previous checkpoint's `torch.save` finishing: one
|
|
323
|
+
# writer thread, and `flush()` above waits for it.
|
|
324
|
+
return {
|
|
325
|
+
"copy": copied - started,
|
|
326
|
+
"queue": time.perf_counter() - copied,
|
|
327
|
+
}
|
|
328
|
+
|
|
329
|
+
def _write(self, torch, snapshot: Dict[str, Any], path: str) -> None:
|
|
330
|
+
temporary = path + ".tmp"
|
|
331
|
+
try:
|
|
332
|
+
torch.save(snapshot, temporary)
|
|
333
|
+
os.replace(temporary, path) # atomic: never leave a half file
|
|
334
|
+
self._prune()
|
|
335
|
+
except Exception as exc:
|
|
336
|
+
logger.error("Checkpoint write failed for %s: %s", path, exc)
|
|
337
|
+
if os.path.exists(temporary):
|
|
338
|
+
try:
|
|
339
|
+
os.remove(temporary)
|
|
340
|
+
except OSError:
|
|
341
|
+
pass
|
|
342
|
+
|
|
343
|
+
def _files(self):
|
|
344
|
+
files = glob.glob(os.path.join(self.directory, "step_*.pt"))
|
|
345
|
+
return sorted(files, key=_step_of)
|
|
346
|
+
|
|
347
|
+
def _prune(self) -> None:
|
|
348
|
+
files = self._files()
|
|
349
|
+
for stale in files[: max(0, len(files) - self.keep_last)]:
|
|
350
|
+
try:
|
|
351
|
+
os.remove(stale)
|
|
352
|
+
except OSError as exc: # pragma: no cover - defensive
|
|
353
|
+
logger.warning("Could not delete old checkpoint %s: %s", stale, exc)
|
|
354
|
+
|
|
355
|
+
def load_latest(self) -> Optional[Dict[str, Any]]:
|
|
356
|
+
import torch
|
|
357
|
+
|
|
358
|
+
files = self._files()
|
|
359
|
+
if not files:
|
|
360
|
+
return None
|
|
361
|
+
state = torch.load(files[-1], map_location="cpu", weights_only=False)
|
|
362
|
+
state.pop("_ravex_metadata", None)
|
|
363
|
+
return state
|
|
364
|
+
|
|
365
|
+
def latest_step(self) -> Optional[int]:
|
|
366
|
+
files = self._files()
|
|
367
|
+
if not files:
|
|
368
|
+
return None
|
|
369
|
+
return _step_of(files[-1])
|
|
370
|
+
|
|
371
|
+
def load_step(self, step: int) -> Optional[Dict[str, Any]]:
|
|
372
|
+
import torch
|
|
373
|
+
|
|
374
|
+
path = os.path.join(self.directory, f"step_{step:012d}.pt")
|
|
375
|
+
if not os.path.exists(path):
|
|
376
|
+
return None
|
|
377
|
+
state = torch.load(path, map_location="cpu", weights_only=False)
|
|
378
|
+
state.pop("_ravex_metadata", None)
|
|
379
|
+
return state
|
|
380
|
+
|
|
381
|
+
def has_checkpoint(self) -> bool:
|
|
382
|
+
return bool(self._files())
|
|
383
|
+
|
|
384
|
+
def flush(self) -> None:
|
|
385
|
+
pending, self._pending = self._pending, None
|
|
386
|
+
if pending is not None:
|
|
387
|
+
pending.result()
|
|
388
|
+
|
|
389
|
+
def close(self) -> None:
|
|
390
|
+
self.flush()
|
|
391
|
+
self._executor.shutdown(wait=True)
|
|
392
|
+
|
|
393
|
+
|
|
394
|
+
# ─── selection ──────────────────────────────────────────────────────
|
|
395
|
+
|
|
396
|
+
|
|
397
|
+
def _per_rank_config(config, per_rank: bool):
|
|
398
|
+
"""Give this rank a store of its own, under ``rank_<n>``.
|
|
399
|
+
|
|
400
|
+
Per-rank checkpointing has every rank writing at once. Pointed at one
|
|
401
|
+
store they would be N writers against one manifest, each reading it,
|
|
402
|
+
adding itself and writing it back with no lock between them — a lost
|
|
403
|
+
update every time two land together. Separate stores need no coordination
|
|
404
|
+
at all, and each rank keeps the background writer, the delta chain and the
|
|
405
|
+
retention it would have had on its own.
|
|
406
|
+
|
|
407
|
+
The cost is N manifests, and a resume that has to agree on a step (see
|
|
408
|
+
``agree_on_step``).
|
|
409
|
+
"""
|
|
410
|
+
from dataclasses import replace
|
|
411
|
+
|
|
412
|
+
from ravex._distributed import get_rank
|
|
413
|
+
|
|
414
|
+
if not per_rank:
|
|
415
|
+
return config
|
|
416
|
+
|
|
417
|
+
suffix = "rank_%d" % get_rank()
|
|
418
|
+
storage = replace(
|
|
419
|
+
config.storage,
|
|
420
|
+
path=os.path.join(config.storage.path, suffix),
|
|
421
|
+
prefix=(
|
|
422
|
+
"%s/%s" % (config.storage.prefix.rstrip("/"), suffix)
|
|
423
|
+
if config.storage.prefix
|
|
424
|
+
else suffix
|
|
425
|
+
),
|
|
426
|
+
)
|
|
427
|
+
return replace(config, storage=storage)
|
|
428
|
+
|
|
429
|
+
|
|
430
|
+
def get_backend(config, per_rank: bool = False) -> CheckpointBackend:
|
|
431
|
+
"""Build the configured backend, falling back to ``torch_save``.
|
|
432
|
+
|
|
433
|
+
A missing or broken Moonclip must never stop a training run: the whole
|
|
434
|
+
proposition is that Ravex is invisible when it works and harmless when
|
|
435
|
+
it does not.
|
|
436
|
+
|
|
437
|
+
``per_rank`` gives this rank a store of its own. The caller decides, not
|
|
438
|
+
the config: whether per-rank checkpointing is actually in force depends on
|
|
439
|
+
the model as well as the setting, and the store has to match what gets
|
|
440
|
+
written into it.
|
|
441
|
+
"""
|
|
442
|
+
config = _per_rank_config(config, per_rank)
|
|
443
|
+
name = config.backend
|
|
444
|
+
|
|
445
|
+
if name == "moonclip":
|
|
446
|
+
try:
|
|
447
|
+
return MoonclipBackend(config)
|
|
448
|
+
except ImportError:
|
|
449
|
+
logger.warning("Moonclip is not installed - falling back to torch.save")
|
|
450
|
+
except Exception as exc:
|
|
451
|
+
logger.warning("Moonclip backend unavailable (%s) - falling back", exc)
|
|
452
|
+
return TorchSaveBackend(config)
|
|
453
|
+
|
|
454
|
+
if name in ("torch_save", "torch"):
|
|
455
|
+
return TorchSaveBackend(config)
|
|
456
|
+
|
|
457
|
+
logger.warning("Unknown backend %r - using torch.save", name)
|
|
458
|
+
return TorchSaveBackend(config)
|