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 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)