blut-core 0.1.0__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.
- blut_core/__init__.py +54 -0
- blut_core/async_io.py +246 -0
- blut_core/checkpoint.py +172 -0
- blut_core/evaluator.py +98 -0
- blut_core/ingredients/__init__.py +20 -0
- blut_core/ingredients/_specs.py +2 -0
- blut_core/ingredients/checkpoint/__init__.py +0 -0
- blut_core/ingredients/checkpoint/_specs.py +162 -0
- blut_core/ingredients/data/__init__.py +0 -0
- blut_core/ingredients/data/_specs.py +87 -0
- blut_core/ingredients/ema/__init__.py +0 -0
- blut_core/ingredients/ema/_specs.py +62 -0
- blut_core/ingredients/eval/__init__.py +0 -0
- blut_core/ingredients/eval/_specs.py +150 -0
- blut_core/ingredients/forward/__init__.py +0 -0
- blut_core/ingredients/forward/_specs.py +35 -0
- blut_core/ingredients/logging/__init__.py +0 -0
- blut_core/ingredients/logging/_specs.py +49 -0
- blut_core/ingredients/loss/__init__.py +0 -0
- blut_core/ingredients/loss/_specs.py +142 -0
- blut_core/ingredients/model/__init__.py +0 -0
- blut_core/ingredients/model/_specs.py +138 -0
- blut_core/ingredients/optimizer/__init__.py +0 -0
- blut_core/ingredients/optimizer/_specs.py +164 -0
- blut_core/ingredients/optimizer/cautious_wd.py +23 -0
- blut_core/ingredients/optimizer/muon_optimizer.py +131 -0
- blut_core/ingredients/optimizer/soap_optimizer.py +301 -0
- blut_core/ingredients/sampler/__init__.py +0 -0
- blut_core/ingredients/sampler/_specs.py +92 -0
- blut_core/ingredients/scheduler/__init__.py +0 -0
- blut_core/ingredients/scheduler/_specs.py +195 -0
- blut_core/ingredients/scheduler/wsd.py +147 -0
- blut_core/ingredients/step/__init__.py +0 -0
- blut_core/ingredients/step/_specs.py +123 -0
- blut_core/load_dataset.py +68 -0
- blut_core/metric_log.py +136 -0
- blut_core/read_metric.py +296 -0
- blut_core/registry.py +109 -0
- blut_core/run_ledger.py +346 -0
- blut_core/run_manifest.py +204 -0
- blut_core/runctx.py +68 -0
- blut_core/spec.py +102 -0
- blut_core/status.py +90 -0
- blut_core/sysgauge.py +82 -0
- blut_core/tests/test_async_compute.py +161 -0
- blut_core/tests/test_async_io_contract.py +466 -0
- blut_core/tests/test_import_contract.py +32 -0
- blut_core/tests/test_load_dataset.py +39 -0
- blut_core/tests/test_parallel_strategy.py +64 -0
- blut_core/tests/test_primitives.py +111 -0
- blut_core/tests/test_read_metric.py +129 -0
- blut_core/tests/test_run_ledger.py +73 -0
- blut_core/trainer.py +631 -0
- blut_core-0.1.0.dist-info/METADATA +38 -0
- blut_core-0.1.0.dist-info/RECORD +57 -0
- blut_core-0.1.0.dist-info/WHEEL +4 -0
- blut_core-0.1.0.dist-info/licenses/LICENSE +202 -0
blut_core/__init__.py
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
"""Shared Python runtime and ingredient registry for BLUT cookbooks.
|
|
2
|
+
|
|
3
|
+
The package contains domain-agnostic building blocks reusable by any cookbook:
|
|
4
|
+
|
|
5
|
+
- ``async_io`` — bounded, backpressured delivery for persistence work
|
|
6
|
+
- ``runctx`` — run identity + filesystem anchors from the BLUT env (P10)
|
|
7
|
+
- ``MetricLog`` — atomic per-epoch CSV/Parquet metric writer
|
|
8
|
+
- ``read_metric`` — verbatim metric/log reader (no LLM in the path; ADR 0038)
|
|
9
|
+
- ``status`` — emit/read the StatusUpdate wire protocol (P4)
|
|
10
|
+
- ``RunManifest`` — run provenance, written even on crash (P8)
|
|
11
|
+
- ``checkpoint`` — corruption-safe save/load + resume payload contract (P7)
|
|
12
|
+
- ``sysgauge`` — best-effort GPU/host snapshot for the metric stream (P10)
|
|
13
|
+
- ingredient specs and builders used by generic training recipes
|
|
14
|
+
|
|
15
|
+
``torch`` is lazy-imported only inside ``checkpoint``/``sysgauge`` functions, so
|
|
16
|
+
``import blut_core`` stays cheap and dependency-light. See ``blut/docs/metrics.md``
|
|
17
|
+
and decisions/0037, 0044.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
from blut_core.registry import (
|
|
21
|
+
build_ingredient,
|
|
22
|
+
get_spec,
|
|
23
|
+
list_ingredients,
|
|
24
|
+
register_ingredient,
|
|
25
|
+
)
|
|
26
|
+
from blut_core.spec import KINDS, IngredientSpec
|
|
27
|
+
|
|
28
|
+
# Register generic ingredient specs on package import.
|
|
29
|
+
from blut_core.ingredients import _specs as _ingredient_specs # noqa: F401
|
|
30
|
+
|
|
31
|
+
from . import async_io, checkpoint, runctx, status, sysgauge
|
|
32
|
+
from .metric_log import MetricLog
|
|
33
|
+
from .run_manifest import RunManifest
|
|
34
|
+
|
|
35
|
+
# read_metric is a CLI module (`python -m blut_core.read_metric`); it is NOT
|
|
36
|
+
# eager-imported here (that triggers a runpy double-import warning under -m).
|
|
37
|
+
# `from blut_core import read_metric` still works (submodule import).
|
|
38
|
+
|
|
39
|
+
__all__ = [
|
|
40
|
+
"IngredientSpec",
|
|
41
|
+
"KINDS",
|
|
42
|
+
"MetricLog",
|
|
43
|
+
"RunManifest",
|
|
44
|
+
"async_io",
|
|
45
|
+
"build_ingredient",
|
|
46
|
+
"checkpoint",
|
|
47
|
+
"get_spec",
|
|
48
|
+
"list_ingredients",
|
|
49
|
+
"read_metric",
|
|
50
|
+
"register_ingredient",
|
|
51
|
+
"runctx",
|
|
52
|
+
"status",
|
|
53
|
+
"sysgauge",
|
|
54
|
+
]
|
blut_core/async_io.py
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
1
|
+
"""Bounded, backpressured asynchronous I/O for cookbook runtimes.
|
|
2
|
+
|
|
3
|
+
Bounded workers are daemons so an exceptional interpreter shutdown cannot hang
|
|
4
|
+
on an idle worker. Successful callers must still call :meth:`AsyncSink.close`;
|
|
5
|
+
that is the only boundary that drains accepted work and surfaces worker errors.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from queue import Queue
|
|
12
|
+
from threading import BoundedSemaphore, Condition, Thread, local
|
|
13
|
+
from typing import Callable, Generic, TypeVar, cast
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
__all__ = [
|
|
17
|
+
"AsyncSink",
|
|
18
|
+
"AsyncSinkWorkerError",
|
|
19
|
+
"Bounded",
|
|
20
|
+
"Inline",
|
|
21
|
+
"ItemTooLargeError",
|
|
22
|
+
"SinkClosedError",
|
|
23
|
+
]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
ItemT = TypeVar("ItemT")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _identity(item: ItemT) -> ItemT:
|
|
30
|
+
return item
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class ItemTooLargeError(ValueError):
|
|
34
|
+
"""A submitted item exceeds the sink's declared retention limit."""
|
|
35
|
+
|
|
36
|
+
def __init__(self, size_bytes: int, max_item_bytes: int) -> None:
|
|
37
|
+
self.size_bytes = size_bytes
|
|
38
|
+
self.max_item_bytes = max_item_bytes
|
|
39
|
+
super().__init__(f"item is {size_bytes} bytes; limit is {max_item_bytes}")
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class SinkClosedError(RuntimeError):
|
|
43
|
+
"""Work was submitted after sink shutdown began."""
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class AsyncSinkWorkerError(RuntimeError):
|
|
47
|
+
"""A sink worker failed while delivering an accepted item."""
|
|
48
|
+
|
|
49
|
+
def __init__(self, worker_error_type: str, message: str) -> None:
|
|
50
|
+
self.worker_error_type = worker_error_type
|
|
51
|
+
super().__init__(message)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@dataclass(frozen=True)
|
|
55
|
+
class _WorkerFailure:
|
|
56
|
+
worker_error_type: str
|
|
57
|
+
message: str
|
|
58
|
+
|
|
59
|
+
@classmethod
|
|
60
|
+
def capture(cls, error: BaseException) -> _WorkerFailure:
|
|
61
|
+
error_type = type(error)
|
|
62
|
+
worker_error_type = f"{error_type.__module__}.{error_type.__qualname__}"
|
|
63
|
+
try:
|
|
64
|
+
message = str(error)
|
|
65
|
+
except BaseException:
|
|
66
|
+
message = "worker failure with an unprintable message"
|
|
67
|
+
return cls(worker_error_type=worker_error_type, message=message)
|
|
68
|
+
|
|
69
|
+
def to_exception(self) -> AsyncSinkWorkerError:
|
|
70
|
+
return AsyncSinkWorkerError(self.worker_error_type, self.message)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
@dataclass(frozen=True)
|
|
74
|
+
class Inline:
|
|
75
|
+
"""Run the installed handler synchronously on the submitting thread."""
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
@dataclass(frozen=True)
|
|
79
|
+
class Bounded:
|
|
80
|
+
"""Run work on one worker with bounded retained residency."""
|
|
81
|
+
|
|
82
|
+
capacity: int
|
|
83
|
+
max_item_bytes: int
|
|
84
|
+
|
|
85
|
+
def __post_init__(self) -> None:
|
|
86
|
+
if isinstance(self.capacity, bool) or not isinstance(self.capacity, int):
|
|
87
|
+
raise TypeError("capacity must be an integer")
|
|
88
|
+
if isinstance(self.max_item_bytes, bool) or not isinstance(
|
|
89
|
+
self.max_item_bytes, int
|
|
90
|
+
):
|
|
91
|
+
raise TypeError("max_item_bytes must be an integer")
|
|
92
|
+
if self.capacity <= 0:
|
|
93
|
+
raise ValueError("capacity must be greater than zero")
|
|
94
|
+
if self.max_item_bytes <= 0:
|
|
95
|
+
raise ValueError("max_item_bytes must be greater than zero")
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
_STOP = object()
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
class AsyncSink(Generic[ItemT]):
|
|
102
|
+
"""Deliver items to one handler under an explicit, close-required mode."""
|
|
103
|
+
|
|
104
|
+
def __init__(
|
|
105
|
+
self,
|
|
106
|
+
handler: Callable[[ItemT], None],
|
|
107
|
+
mode: Inline | Bounded,
|
|
108
|
+
*,
|
|
109
|
+
item_size: Callable[[ItemT], int] | None = None,
|
|
110
|
+
retain: Callable[[ItemT], ItemT] = _identity,
|
|
111
|
+
retained_size: Callable[[ItemT], int] | None = None,
|
|
112
|
+
thread_name: str = "blut-async-io",
|
|
113
|
+
) -> None:
|
|
114
|
+
self._handler = handler
|
|
115
|
+
self._mode = mode
|
|
116
|
+
self._item_size = item_size
|
|
117
|
+
self._retain = retain
|
|
118
|
+
self._retained_size = retained_size
|
|
119
|
+
self._closed = False
|
|
120
|
+
self._failure: _WorkerFailure | None = None
|
|
121
|
+
self._state = Condition()
|
|
122
|
+
self._active_submits = 0
|
|
123
|
+
self._stop_enqueued = False
|
|
124
|
+
self._callback_state = local()
|
|
125
|
+
self._queue: Queue[ItemT | object] | None = None
|
|
126
|
+
self._slots: BoundedSemaphore | None = None
|
|
127
|
+
self._worker: Thread | None = None
|
|
128
|
+
if isinstance(mode, Bounded):
|
|
129
|
+
self._queue = Queue(maxsize=mode.capacity)
|
|
130
|
+
self._slots = BoundedSemaphore(mode.capacity)
|
|
131
|
+
self._worker = Thread(target=self._run, name=thread_name, daemon=True)
|
|
132
|
+
self._worker.start()
|
|
133
|
+
|
|
134
|
+
def submit(self, item: ItemT, *, size_bytes: int | None = None) -> None:
|
|
135
|
+
"""Deliver one item according to the configured mode."""
|
|
136
|
+
|
|
137
|
+
if getattr(self._callback_state, "active", False):
|
|
138
|
+
raise RuntimeError("submit cannot be called from a sink callback")
|
|
139
|
+
|
|
140
|
+
with self._state:
|
|
141
|
+
self._raise_if_failed_locked()
|
|
142
|
+
if self._closed:
|
|
143
|
+
raise SinkClosedError("sink is closed")
|
|
144
|
+
self._active_submits += 1
|
|
145
|
+
|
|
146
|
+
acquired_slot = False
|
|
147
|
+
try:
|
|
148
|
+
if isinstance(self._mode, Inline):
|
|
149
|
+
self._invoke_handler(self._invoke_callback(self._retain, item))
|
|
150
|
+
return
|
|
151
|
+
|
|
152
|
+
if size_bytes is None and self._item_size is not None:
|
|
153
|
+
size_bytes = self._invoke_callback(self._item_size, item)
|
|
154
|
+
if size_bytes is None and self._retained_size is None:
|
|
155
|
+
raise TypeError(
|
|
156
|
+
"bounded submit requires size_bytes, item_size, or retained_size"
|
|
157
|
+
)
|
|
158
|
+
if size_bytes is not None:
|
|
159
|
+
if isinstance(size_bytes, bool) or not isinstance(size_bytes, int):
|
|
160
|
+
raise TypeError("item size must be an integer byte count")
|
|
161
|
+
if size_bytes < 0:
|
|
162
|
+
raise ValueError("item size must not be negative")
|
|
163
|
+
if size_bytes > self._mode.max_item_bytes:
|
|
164
|
+
raise ItemTooLargeError(size_bytes, self._mode.max_item_bytes)
|
|
165
|
+
|
|
166
|
+
slots = cast(BoundedSemaphore, self._slots)
|
|
167
|
+
queue = cast(Queue[ItemT | object], self._queue)
|
|
168
|
+
slots.acquire()
|
|
169
|
+
acquired_slot = True
|
|
170
|
+
with self._state:
|
|
171
|
+
self._raise_if_failed_locked()
|
|
172
|
+
if self._closed:
|
|
173
|
+
raise SinkClosedError("sink is closed")
|
|
174
|
+
retained = self._invoke_callback(self._retain, item)
|
|
175
|
+
if self._retained_size is not None:
|
|
176
|
+
retained_bytes = self._invoke_callback(self._retained_size, retained)
|
|
177
|
+
if isinstance(retained_bytes, bool) or not isinstance(retained_bytes, int):
|
|
178
|
+
raise TypeError("retained item size must be an integer byte count")
|
|
179
|
+
if retained_bytes < 0:
|
|
180
|
+
raise ValueError("retained item size must not be negative")
|
|
181
|
+
if retained_bytes > self._mode.max_item_bytes:
|
|
182
|
+
raise ItemTooLargeError(retained_bytes, self._mode.max_item_bytes)
|
|
183
|
+
queue.put(retained)
|
|
184
|
+
acquired_slot = False
|
|
185
|
+
finally:
|
|
186
|
+
if acquired_slot:
|
|
187
|
+
cast(BoundedSemaphore, self._slots).release()
|
|
188
|
+
with self._state:
|
|
189
|
+
self._active_submits -= 1
|
|
190
|
+
self._state.notify_all()
|
|
191
|
+
|
|
192
|
+
def close(self) -> None:
|
|
193
|
+
"""Drain accepted work and stop the sink."""
|
|
194
|
+
|
|
195
|
+
if getattr(self._callback_state, "active", False):
|
|
196
|
+
raise RuntimeError("close cannot be called from a sink callback")
|
|
197
|
+
|
|
198
|
+
with self._state:
|
|
199
|
+
self._closed = True
|
|
200
|
+
while self._active_submits:
|
|
201
|
+
self._state.wait()
|
|
202
|
+
if self._worker is not None and not self._stop_enqueued:
|
|
203
|
+
cast(Queue[ItemT | object], self._queue).put(_STOP)
|
|
204
|
+
self._stop_enqueued = True
|
|
205
|
+
|
|
206
|
+
if self._worker is not None:
|
|
207
|
+
self._worker.join()
|
|
208
|
+
self._raise_if_failed()
|
|
209
|
+
|
|
210
|
+
def _run(self) -> None:
|
|
211
|
+
queue = cast(Queue[ItemT | object], self._queue)
|
|
212
|
+
slots = cast(BoundedSemaphore, self._slots)
|
|
213
|
+
while True:
|
|
214
|
+
item = queue.get()
|
|
215
|
+
if item is _STOP:
|
|
216
|
+
return
|
|
217
|
+
try:
|
|
218
|
+
self._invoke_handler(cast(ItemT, item))
|
|
219
|
+
except BaseException as error:
|
|
220
|
+
with self._state:
|
|
221
|
+
if self._failure is None:
|
|
222
|
+
self._failure = _WorkerFailure.capture(error)
|
|
223
|
+
finally:
|
|
224
|
+
del item
|
|
225
|
+
slots.release()
|
|
226
|
+
|
|
227
|
+
def _invoke_handler(self, item: ItemT) -> None:
|
|
228
|
+
self._invoke_callback(self._handler, item)
|
|
229
|
+
|
|
230
|
+
def _invoke_callback(self, callback, item):
|
|
231
|
+
previous = getattr(self._callback_state, "active", False)
|
|
232
|
+
self._callback_state.active = True
|
|
233
|
+
try:
|
|
234
|
+
return callback(item)
|
|
235
|
+
finally:
|
|
236
|
+
self._callback_state.active = previous
|
|
237
|
+
|
|
238
|
+
def _raise_if_failed_locked(self) -> None:
|
|
239
|
+
if self._failure is not None:
|
|
240
|
+
raise self._failure.to_exception()
|
|
241
|
+
|
|
242
|
+
def _raise_if_failed(self) -> None:
|
|
243
|
+
with self._state:
|
|
244
|
+
failure = self._failure
|
|
245
|
+
if failure is not None:
|
|
246
|
+
raise failure.to_exception()
|
blut_core/checkpoint.py
ADDED
|
@@ -0,0 +1,172 @@
|
|
|
1
|
+
"""checkpoint — corruption-safe torch checkpoint I/O + a resume payload contract
|
|
2
|
+
(ADR 0044 P7, "durability"; the single most important reliability gap).
|
|
3
|
+
|
|
4
|
+
- ``save(payload, path)`` — free-space preflight → write to a unique tmp in the
|
|
5
|
+
same dir → ``fsync`` → atomic ``os.replace`` → write a sidecar ``<path>.sha256``.
|
|
6
|
+
A mid-write kill leaves the PRIOR checkpoint intact, never a truncated file
|
|
7
|
+
(the failure mode that silently corrupts on a disk-full ``torch.save``).
|
|
8
|
+
Optional ``keep_last_n`` rotation of ``<stem>_ep*.pt`` siblings.
|
|
9
|
+
- ``load(path, validate=True)`` — verify the sidecar SHA (or, absent a sidecar,
|
|
10
|
+
a dry ``torch.load`` of the bytes) BEFORE returning → a corrupt/truncated
|
|
11
|
+
checkpoint raises ``CheckpointError``, never silent garbage.
|
|
12
|
+
- ``validate_payload(payload, required)`` — assert the resume contract keys are
|
|
13
|
+
present so a resume fails fast at SAVE, not 3 epochs into a relaunch (prior
|
|
14
|
+
bites: optimizer-state, QAT-alpha). Default contract = the training-resume set.
|
|
15
|
+
|
|
16
|
+
torch is lazy-imported (only on save/load), so importing ``blut_core`` stays cheap.
|
|
17
|
+
"""
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import hashlib
|
|
21
|
+
import os
|
|
22
|
+
import shutil
|
|
23
|
+
import tempfile
|
|
24
|
+
from pathlib import Path
|
|
25
|
+
from typing import Any, Dict, Iterable, Optional
|
|
26
|
+
|
|
27
|
+
__all__ = ["save", "load", "validate_payload", "RESUME_CONTRACT", "CheckpointError"]
|
|
28
|
+
|
|
29
|
+
# ADR 0044 P7 resume payload contract. A trainer that wants resumability MUST
|
|
30
|
+
# put these in the checkpoint; validate_payload(strict=True) enforces it.
|
|
31
|
+
RESUME_CONTRACT = (
|
|
32
|
+
"model", "optimizer", "lr_scheduler", "rng_state", "step",
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class CheckpointError(RuntimeError):
|
|
37
|
+
"""A checkpoint is missing, corrupt, or violates the resume contract."""
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def _sha256(path: Path) -> str:
|
|
41
|
+
h = hashlib.sha256()
|
|
42
|
+
with path.open("rb") as f:
|
|
43
|
+
for chunk in iter(lambda: f.read(1 << 20), b""):
|
|
44
|
+
h.update(chunk)
|
|
45
|
+
return h.hexdigest()
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def validate_payload(payload: Dict[str, Any], required: Iterable[str] = RESUME_CONTRACT,
|
|
49
|
+
*, strict: bool = True) -> list:
|
|
50
|
+
"""Return the list of missing contract keys. ``strict`` → raise on any missing.
|
|
51
|
+
Use before save so a resume can't silently lack optimizer/RNG/step."""
|
|
52
|
+
if not isinstance(payload, dict):
|
|
53
|
+
raise CheckpointError(f"payload must be a dict, got {type(payload).__name__}")
|
|
54
|
+
missing = [k for k in required if k not in payload]
|
|
55
|
+
if missing and strict:
|
|
56
|
+
raise CheckpointError(
|
|
57
|
+
f"checkpoint payload missing resume-contract keys {missing}; "
|
|
58
|
+
f"present: {sorted(payload)}")
|
|
59
|
+
return missing
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _free_bytes(directory: Path) -> int:
|
|
63
|
+
return shutil.disk_usage(directory).free
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def save(payload: Dict[str, Any], path, *, contract: Optional[Iterable[str]] = None,
|
|
67
|
+
keep_last_n: Optional[int] = None, min_free_mb: int = 256) -> Path:
|
|
68
|
+
"""Atomically save ``payload`` to ``path`` with a sidecar SHA.
|
|
69
|
+
|
|
70
|
+
``contract`` (default None = skip) → validate_payload(strict) first.
|
|
71
|
+
``keep_last_n`` → after writing, keep only the newest N ``<stem>_ep*.pt``.
|
|
72
|
+
Raises ``CheckpointError`` on insufficient free space or a contract miss."""
|
|
73
|
+
import torch # lazy
|
|
74
|
+
path = Path(path)
|
|
75
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
76
|
+
if contract is not None:
|
|
77
|
+
validate_payload(payload, contract, strict=True)
|
|
78
|
+
|
|
79
|
+
# Free-space preflight: a disk-full torch.save silently yields garbage.
|
|
80
|
+
if _free_bytes(path.parent) < min_free_mb * (1 << 20):
|
|
81
|
+
raise CheckpointError(
|
|
82
|
+
f"insufficient free space in {path.parent} (<{min_free_mb} MiB) — refusing to save")
|
|
83
|
+
|
|
84
|
+
fd, tmp_name = tempfile.mkstemp(dir=str(path.parent), prefix=path.stem + ".",
|
|
85
|
+
suffix=path.suffix + ".tmp")
|
|
86
|
+
tmp = Path(tmp_name)
|
|
87
|
+
try:
|
|
88
|
+
os.close(fd) # torch.save writes by path
|
|
89
|
+
torch.save(payload, tmp)
|
|
90
|
+
with tmp.open("rb") as f:
|
|
91
|
+
os.fsync(f.fileno())
|
|
92
|
+
sha = _sha256(tmp)
|
|
93
|
+
os.replace(tmp, path)
|
|
94
|
+
# Sidecar: write to a tmp, fsync, then atomic rename — write_text() left
|
|
95
|
+
# the SHA in the page cache (no fsync), and a non-atomic overwrite could
|
|
96
|
+
# leave a .sha256 pointing at the PREVIOUS checkpoint after a crash.
|
|
97
|
+
sha_path = path.with_suffix(path.suffix + ".sha256")
|
|
98
|
+
sha_tmp = sha_path.with_suffix(sha_path.suffix + ".tmp")
|
|
99
|
+
payload_bytes = (sha + "\n").encode("utf-8")
|
|
100
|
+
try:
|
|
101
|
+
sfd = os.open(str(sha_tmp), os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o644)
|
|
102
|
+
try:
|
|
103
|
+
n = os.write(sfd, payload_bytes)
|
|
104
|
+
if n != len(payload_bytes): # never short on a local FS at this size
|
|
105
|
+
raise OSError(f"short sidecar write: {n}/{len(payload_bytes)}")
|
|
106
|
+
os.fsync(sfd)
|
|
107
|
+
finally:
|
|
108
|
+
os.close(sfd)
|
|
109
|
+
os.replace(sha_tmp, sha_path)
|
|
110
|
+
except Exception:
|
|
111
|
+
try:
|
|
112
|
+
sha_tmp.unlink()
|
|
113
|
+
except OSError:
|
|
114
|
+
pass
|
|
115
|
+
raise
|
|
116
|
+
except Exception:
|
|
117
|
+
try:
|
|
118
|
+
tmp.unlink()
|
|
119
|
+
except Exception:
|
|
120
|
+
pass
|
|
121
|
+
raise
|
|
122
|
+
|
|
123
|
+
if keep_last_n is not None and keep_last_n > 0:
|
|
124
|
+
_rotate(path, keep_last_n)
|
|
125
|
+
return path
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def _rotate(path: Path, keep_last_n: int) -> None:
|
|
129
|
+
# Rotates the EPOCH-SNAPSHOT siblings `{stem}_ep<N>{suffix}` (the trainer
|
|
130
|
+
# convention, e.g. snn_4state_best_ep10.pt), keeping the newest N. It does
|
|
131
|
+
# NOT touch the base `{stem}{suffix}` (the canonical/best checkpoint) — that
|
|
132
|
+
# is intentional: you keep best + the last N epoch snapshots.
|
|
133
|
+
sibs = sorted(path.parent.glob(f"{path.stem}_ep*{path.suffix}"),
|
|
134
|
+
key=lambda p: p.stat().st_mtime, reverse=True)
|
|
135
|
+
for old in sibs[keep_last_n:]:
|
|
136
|
+
for victim in (old, old.with_suffix(old.suffix + ".sha256")):
|
|
137
|
+
try:
|
|
138
|
+
victim.unlink()
|
|
139
|
+
except OSError:
|
|
140
|
+
pass
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
def load(path, *, validate: bool = True, map_location: Any = "cpu") -> Dict[str, Any]:
|
|
144
|
+
"""Load a checkpoint, verifying integrity first. Raises ``CheckpointError``
|
|
145
|
+
on a missing/corrupt file rather than returning silent garbage."""
|
|
146
|
+
import torch # lazy
|
|
147
|
+
path = Path(path)
|
|
148
|
+
if not path.exists():
|
|
149
|
+
raise CheckpointError(f"no checkpoint: {path}")
|
|
150
|
+
if validate:
|
|
151
|
+
sidecar = path.with_suffix(path.suffix + ".sha256")
|
|
152
|
+
if sidecar.exists():
|
|
153
|
+
want = sidecar.read_text().strip()
|
|
154
|
+
got = _sha256(path)
|
|
155
|
+
if got != want:
|
|
156
|
+
raise CheckpointError(
|
|
157
|
+
f"checkpoint SHA mismatch for {path}: sidecar {want[:12]}… got {got[:12]}…")
|
|
158
|
+
else:
|
|
159
|
+
# No sidecar → integrity is UNVERIFIED (e.g. a crash between the
|
|
160
|
+
# checkpoint rename and the sidecar write; the sidecar is not part
|
|
161
|
+
# of the atomic unit). Surface it; don't silently accept-as-clean.
|
|
162
|
+
import sys as _sys
|
|
163
|
+
_sys.stderr.write(
|
|
164
|
+
f"[checkpoint] WARNING: no .sha256 sidecar for {path}; integrity unverified\n")
|
|
165
|
+
try:
|
|
166
|
+
# NOT weights_only=True: the resume payload carries optimizer/LR/RNG/
|
|
167
|
+
# scheduler objects (not plain tensors), which weights_only rejects.
|
|
168
|
+
# Trust model: checkpoints are own-produced under the job dir, not loaded
|
|
169
|
+
# from untrusted sources — the sidecar SHA guards corruption, not supply chain.
|
|
170
|
+
return torch.load(path, map_location=map_location)
|
|
171
|
+
except Exception as e:
|
|
172
|
+
raise CheckpointError(f"failed to load {path}: {type(e).__name__}: {e}") from e
|
blut_core/evaluator.py
ADDED
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
"""Generic ingredient-based evaluator for BLUT core cookbook.
|
|
2
|
+
|
|
3
|
+
Reads an eval config JSON (written by the EvaluateModel Rust stage),
|
|
4
|
+
builds the eval ingredient, and runs evaluation.
|
|
5
|
+
|
|
6
|
+
Usage:
|
|
7
|
+
python -m blut_core.evaluator --config eval_config.json --output eval_report.json
|
|
8
|
+
"""
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import argparse
|
|
12
|
+
import json
|
|
13
|
+
import sys
|
|
14
|
+
from pathlib import Path
|
|
15
|
+
|
|
16
|
+
sys.path.insert(0, str(Path(__file__).parent.parent))
|
|
17
|
+
|
|
18
|
+
from blut_core import build_ingredient
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def load_config(config_path: str) -> dict:
|
|
22
|
+
with open(config_path) as f:
|
|
23
|
+
return json.load(f)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def main():
|
|
27
|
+
parser = argparse.ArgumentParser(description="BLUT core generic evaluator")
|
|
28
|
+
parser.add_argument("--config", required=True, help="Path to eval config JSON")
|
|
29
|
+
parser.add_argument("--output", required=True, help="Path to output eval report JSON")
|
|
30
|
+
args = parser.parse_args()
|
|
31
|
+
|
|
32
|
+
config = load_config(args.config)
|
|
33
|
+
|
|
34
|
+
checkpoint_path = config["checkpoint_path"]
|
|
35
|
+
dataset_path = config["dataset_path"]
|
|
36
|
+
eval_cfg = config["eval"]
|
|
37
|
+
batch_size = config.get("batch_size", 64)
|
|
38
|
+
device = config.get("device", "cuda")
|
|
39
|
+
|
|
40
|
+
# Build eval ingredient
|
|
41
|
+
print(f"[evaluator] building eval: {eval_cfg['kind']}:{eval_cfg['name']}")
|
|
42
|
+
evaluate_fn = build_ingredient(eval_cfg["kind"], eval_cfg["name"], eval_cfg.get("config"))
|
|
43
|
+
|
|
44
|
+
# Load checkpoint
|
|
45
|
+
import torch
|
|
46
|
+
print(f"[evaluator] loading checkpoint from {checkpoint_path}")
|
|
47
|
+
checkpoint_path = Path(checkpoint_path)
|
|
48
|
+
|
|
49
|
+
# Try to load as HuggingFace model
|
|
50
|
+
model = None
|
|
51
|
+
hf_dir = checkpoint_path / "hf"
|
|
52
|
+
if hf_dir.exists():
|
|
53
|
+
try:
|
|
54
|
+
from transformers import AutoModelForCausalLM
|
|
55
|
+
model = AutoModelForCausalLM.from_pretrained(str(hf_dir))
|
|
56
|
+
print(f"[evaluator] loaded HF model from {hf_dir}")
|
|
57
|
+
except Exception as e:
|
|
58
|
+
print(f"[evaluator] warning: could not load HF model: {e}")
|
|
59
|
+
|
|
60
|
+
if model is None:
|
|
61
|
+
# Try loading as torch checkpoint
|
|
62
|
+
pt_path = checkpoint_path / "model.pt"
|
|
63
|
+
if pt_path.exists():
|
|
64
|
+
print(f"[evaluator] found torch checkpoint at {pt_path}")
|
|
65
|
+
# Need a model architecture to load into — this is a limitation
|
|
66
|
+
# of the generic evaluator. Real usage should provide the model.
|
|
67
|
+
print("[evaluator] error: no model architecture available for torch checkpoint")
|
|
68
|
+
sys.exit(1)
|
|
69
|
+
else:
|
|
70
|
+
print(f"[evaluator] error: no checkpoint found at {checkpoint_path}")
|
|
71
|
+
sys.exit(1)
|
|
72
|
+
|
|
73
|
+
# Load dataset
|
|
74
|
+
print(f"[evaluator] loading dataset from {dataset_path}")
|
|
75
|
+
try:
|
|
76
|
+
from datasets import load_dataset
|
|
77
|
+
ds = load_dataset("json", data_files=str(dataset_path), split="train")
|
|
78
|
+
print(f"[evaluator] loaded {len(ds)} examples")
|
|
79
|
+
except Exception as e:
|
|
80
|
+
print(f"[evaluator] error: could not load dataset: {e}")
|
|
81
|
+
sys.exit(1)
|
|
82
|
+
|
|
83
|
+
# Run evaluation
|
|
84
|
+
print(f"[evaluator] running evaluation...")
|
|
85
|
+
if device != "cpu" and torch.cuda.is_available():
|
|
86
|
+
model = model.to(device)
|
|
87
|
+
|
|
88
|
+
metrics = evaluate_fn(model, ds, device=device)
|
|
89
|
+
print(f"[evaluator] metrics: {metrics}")
|
|
90
|
+
|
|
91
|
+
# Write output
|
|
92
|
+
with open(args.output, "w") as f:
|
|
93
|
+
json.dump(metrics, f, indent=2)
|
|
94
|
+
print(f"[evaluator] report written to {args.output}")
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
if __name__ == "__main__":
|
|
98
|
+
main()
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""Generic ingredient specs — the core training primitives.
|
|
2
|
+
|
|
3
|
+
Importing this module registers all generic specs into the global registry.
|
|
4
|
+
Domain cookbooks import ``blut_core`` first (registering these), then register
|
|
5
|
+
their own domain-specific specs on top.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
# Import all _specs modules to trigger their @register_ingredient decorators.
|
|
9
|
+
from blut_core.ingredients.data import _specs as _data_specs # noqa: F401
|
|
10
|
+
from blut_core.ingredients.model import _specs as _model_specs # noqa: F401
|
|
11
|
+
from blut_core.ingredients.optimizer import _specs as _optimizer_specs # noqa: F401
|
|
12
|
+
from blut_core.ingredients.scheduler import _specs as _scheduler_specs # noqa: F401
|
|
13
|
+
from blut_core.ingredients.loss import _specs as _loss_specs # noqa: F401
|
|
14
|
+
from blut_core.ingredients.step import _specs as _step_specs # noqa: F401
|
|
15
|
+
from blut_core.ingredients.ema import _specs as _ema_specs # noqa: F401
|
|
16
|
+
from blut_core.ingredients.checkpoint import _specs as _checkpoint_specs # noqa: F401
|
|
17
|
+
from blut_core.ingredients.eval import _specs as _eval_specs # noqa: F401
|
|
18
|
+
from blut_core.ingredients.sampler import _specs as _sampler_specs # noqa: F401
|
|
19
|
+
from blut_core.ingredients.logging import _specs as _logging_specs # noqa: F401
|
|
20
|
+
from blut_core.ingredients.forward import _specs as _forward_specs # noqa: F401
|
|
File without changes
|